已合并
refactor: OpDispatcher 重构 #508
refactor: OpDispatcher 重构 #508
已合并
hedongdong创建于 4月8日
hedongdong成员
4月8日

相关的 Issue

无相关 Issue,框架扩展性增强。

原因(目的、解决的问题等)

当前 OpDispatcher 的实现存在以下问题,影响代码的可维护性和正确性:

1. extra_args 丢失参数名导致缓存冲突(核心问题)

问题:PyTorch 透传模式下 extra_args 只保留 value,丢失参数名和顺序。

  • torch.sum(x, 1, True) → extra_args = [1, True]
  • torch.sum(x, True, 1) → extra_args = [True, 1](相同!)
  • 但实际执行结果不同(dim=1 vs dim=True),导致缓存命中错误

根本原因:透传模式下参数绑定发生在 C++ 层,Python 框架无法获知参数名。

2. 副作用设计

问题:_process_args_and_kwargs 原地修改 cache_key.layout_ids(可变 list),违反纯函数设计原则。

# 旧代码 - 原地修改
def _process_args_and_kwargs(args, kwargs, cache_key):
    ...
    cache_key.layout_ids.append(id_str)  # 副作用!
    ...

3. LayoutCacheKey 可变性与性能问题

问题:

  • 使用可变 list 存储 cache key,违反不可变设计原则
  • 每次 __hash__ 重新计算(循环 XOR + 位运算),性能差
  • 无 __slots__,内存占用高

4. 旧流程 LayoutCacheKey缓存内容过多,性能差

问题:原流程将args/kwargs中所有的layouts以及关键字参数的value进行缓存,其中可能包含与切分无关的参数。无关参数的修改不影响layout的推导,但是导致缓存不命中。需要最小化缓存内容。

5. 旧流程 suffix 逻辑硬编码

问题:dispatch 中的 5 个 if 分支硬编码(WithShape/Reshape/WithTupleExpand/Slice 等),新增算子需要修改 OpDispatcher 核心逻辑。

详细设计

新流程架构

┌─────────────────────────────────────────────────────────────────────────┐
│                         新流程架构                                       │
├─────────────────────────────────────────────────────────────────────────┤
│                                                                          │
│  OpDispatcher.dispatch(op_call, args, kwargs)                           │
│       │                                                                  │
│       ├── 白名单 → 直接 to_local 执行                                    │
│       ├── 随机算子 → _dispatch_random_op                                │
│       └── 分布式算子 ↓                                                   │
│                                                                          │
│  distribute_op.preprocess(args, kwargs)                                 │
│       │                                                                  │
│       ├── 返回 None → 旧流程(兼容)                                     │
│       └── 返回 (local_args, local_kwargs, cache_values) ↓               │
│                                                                          │
│  构建 cache_key:                                                         │
│       cache_key_values = []                                              │
│       for v in cache_values:                                             │
│           if hasattr(v, 'compact_str'):                                  │
│               cache_key_values.append(str(v.compact_str))  # Layout      │
│           else:                                                          │
│               cache_key_values.append(str(v))              # 原始值        │
│       cache_key = LayoutCacheKey(cache_key_values)                       │
│       │                                                                  │
│       ↓                                                                  │
│  缓存查询                                                                │
│       ├── 命中 → (infer_result, op_impl)                                │
│       └── 未命中 →                                                       │
│            infer_result = infer_layout(cache_values)                     │
│            op_impl = get_expand_impl(func, infer_result, cache_values)   │
│       │                                                                  │
│       ↓                                                                  │
│  解析 infer_result:                                                      │
│       output_layouts, extra_info = infer_result                          │
│       # output_layouts: tuple of Layout                                  │
│       # extra_info: 额外信息,默认为 None                                │
│       │                                                                  │
│       ↓                                                                  │
│  op_impl(*local_args, **local_kwargs) 或                                │
│  op_impl(*local_args, *extra_info, **local_kwargs)                     │
│       │                                                                  │
│       ↓                                                                  │
│  DTensor.from_local()                                                   │
│                                                                          │
└─────────────────────────────────────────────────────────────────────────┘

1. LayoutCacheKey 不可变设计

class LayoutCacheKey:
    """Immutable layout cache key."""
    __slots__ = ('_tuple', '_hash')  # 内存优化

    def __init__(self, layout_ids: List[str]):
        self._tuple = tuple(layout_ids)  # 不可变 tuple
        self._hash = hash(self._tuple)   # 预计算 hash

    @classmethod
    def from_cache_values(cls, cache_values):
        """从新流程 cache_values 构建缓存键"""
        key_values = []
        for v in cache_values:
            if hasattr(v, 'compact_str'):
                key_values.append(str(v.compact_str))  # Layout
            else:
                key_values.append(str(v))             # 原始值
        return cls(key_values)

    def __hash__(self):
        return self._hash  # 直接返回预计算值

改进点:

  • 使用 __slots__ 优化内存
  • 使用不可变 tuple 替代可变 list
  • 预计算 hash 值,避免重复计算

2. 消除副作用

# 旧代码 - 原地修改
def _process_args_and_kwargs(args, kwargs, cache_key):
    ...
    cache_key.layout_ids.append(id_str)  # 副作用!

# 新代码 - 返回列表
def _process_args_and_kwargs(args, kwargs):
    ...
    cache_key_values.append(id_str)  # 无副作用
    return input_layouts, extra_args, input_args, input_kwargs, cache_key_values

改进点:

  • _process_args_and_kwargs 不再接收 cache_key 参数
  • 返回 cache_key_values 列表,由调用方构建 LayoutCacheKey
  • 符合纯函数设计原则

3. 新增 preprocess 流程入口

为后续扩展预留接口,支持新流程的预处理机制:

def dispatch(self, op_call, args, kwargs):
    ...
    distribute_op = cache_manager.distributed_op(op_name)
    result = distribute_op.preprocess(args, kwargs)
    if result is not None:
        return self._dispatch_new(op_call, distribute_op, result)
    ...

4. 统一输出包装方法

新增 _wrap_output 方法,统一 DTensor 包装逻辑:

def _wrap_output(self, py_output, output_layouts) -> Tensor:
    if isinstance(py_output, (tuple, list)):
        return tuple(
            DTensor.from_local(item, layout.mesh, layout.alias_placements)
            for item, layout in zip(py_output, output_layouts))
    return DTensor.from_local(
        py_output, output_layouts[0].mesh, output_layouts[0].alias_placements)

变更文件列表

修改的现有文件:

  1. hyper_parallel/core/shard/_op_dispatch.py - LayoutCacheKey 不可变设计,消除副作用,新增 preprocess 流程
  2. hyper_parallel/core/shard/ops/parallel_matmul.py - 适配测试
  3. hyper_parallel/core/shard/ops/parallel_ops.py - 适配测试
  4. hyper_parallel/core/shard/ops/parallel_sort.py - 适配测试
  5. tests/ut/core/shard/ops/test_parallel_linear.py - 适配测试
  6. tests/ut/core/shard/ops/test_parallel_sort.py - 适配测试
  7. tests/ut/core/shard/test_op_dispatch.py - 适配测试

收益

直接收益

  1. 消除副作用:_process_args_and_kwargs 不再修改外部状态,符合纯函数设计
  2. 缓存键不可变:LayoutCacheKey 使用不可变 tuple,避免意外修改
  3. 性能提升:预计算 hash + __slots__ 优化,减少内存和计算开销
  4. 最小化缓存内容:preprocess将定义返回的缓存内容规范化/最小化,提高缓存命中率。

扩展收益

  1. 统一输出包装:新增 _wrap_output 方法,消除代码重复
  2. preprocess 入口:为后续新流程(参数规范化)预留接口,解决 extra_args 丢失参数名问题
  3. 为后续重构奠基:preprocess 方法可让各算子自行处理参数规范化,从根本上解决透传模式下的缓存冲突问题

测试验证

  • 21 个相关 UT 测试全部通过
  • 无功能回归
likedislike
Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 hedongdong 的贡献)
Hhedongdong成员
4月8日 创建了 pull request,commit 0f69ba4f
Hhedongdong成员
4月8日 virtual merging failed, update merge request[project_id: 8714117, iid: 508, target_commit_sha: 0f69ba4f5490ce066452d28fd6d4c2552c423008], message: Conflict detected
Hhedongdong成员
4月8日 审查状态已重置,审查人: fengyixing,yangzhenzhang
Hhedongdong成员
4月8日 virtual merging failed, update merge request[project_id: 8714117, iid: 508, target_commit_sha: 0f69ba4f5490ce066452d28fd6d4c2552c423008], message: Conflict detected
Hhedongdong成员
4月8日 修改标题为 “refactor: OpDispatcher 重构 - 消除副作用,统一缓存键设计”,原标题为“refactor: OpDispatcher 重构 - 统一预处理流程”
此处折叠了211条消息 查看更多
fengyixingfengyixing成员
4月9日 通过审查
MindSpore-Bot
MindSpore-Bot成员
4月9日 评论:

Notice

The PR needs 3 assignees to review. 1 does not review. if all are passed, please comment /check-pr to try merge the PR. 😄

likedislike
MengXY107MengXY107成员
4月9日 通过审查
MindSpore-BotMindSpore-Bot成员
4月9日 删除了label:no-pass-all-review
MindSpore-BotMindSpore-Bot成员
4月9日 合入了pull request,合并节点 SHA:353b0ab199b7fd015df9ba9e2e01bae420ec084f