mindspore-ai / mindspore-ai/hyper-parallel

feat: 接入 AllGatherMatmul 和 MatmulReduceScatter MC2 算子分布式调度逻辑

Open
#720 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

ISSUE: 接入 MindSpore MC2 通算融合算子分布式调度(AllGatherMatmul / MatmulReduceScatter)

背景

在分布式 Tensor Parallelism(TP)训练场景中,CANN MC2 算子(AllGatherMatmul / MatmulReduceScatter)在内核内部自行完成 AllGather / ReduceScatter 通信,性能优于分离式通信+计算方案。然而 HyperParallel 当前没有这两个算子的分布式调度逻辑,用 DTensor 对象调用时会直接失败(CANN 内核无法处理 DTensor),即使绕过也无法推断输出 DTensor 的分布状态,下游算子无法正确组合分布式计算图。

目标

为以下两个 MindSpore MC2 算子提供分布式调度支持:

AllGatherMatmul
  • CANN 内核内部完成 AllGather(消耗 x1 m 维分片),HyperParallel 负责:
    • 通过 preprocess()to_local() 提取本地 tensor 传给 CANN
    • 通过 infer_layout() 推断输出 layout:
      • m 维始终 Replicate(AllGather 消耗了 m 维分片)
      • n 维继承 x2 的 n 维 placement
      • k 被切分时输出带 Partial(sum),调用方需 AllReduce(与 LinearDistributedOp contract_dim 分片语义一致)
      • gather_output=False 时 CANN 返回空 tensor,gather_out layout 强制全 Replicate,避免 Shard dim 越界

支持的分片模式:

  • x1 Shard(0) on tp,x2 Replicate
  • x1 Shard(0) on tp,x2 Shard(1) on n
  • trans_x2=True
  • 2D mesh (dp, tp)
  • k 切分:4 卡 (mp=2, tp=2),x1 Shard(0)+Shard(1),x2 Replicate+Shard(0),输出带 Partial(sum) on tp
  • k 切分:8 卡 (mp=2, np=2, tp=2),x1 Shard(0)+Replicate+Shard(1),x2 Replicate+Shard(1)+Shard(0),输出带 Partial(sum) on tp
MatmulReduceScatter
  • CANN 内核内部完成 ReduceScatter(k 方向 sum reduce + m 方向 scatter),HyperParallel 负责:
    • 通过 preprocess()to_local() 提取本地 tensor 传给 CANN
    • 通过 infer_layout() 推断输出 layout:
      • output dim 0 (m):ReduceScatter 将 k 的 TP 分片转化为 m 分片;若 x1 m 维有 DP 分片则联合分片(tuple tensor_map)
      • output dim 1 (n):继承 x2 的 n 维 placement
      • 无 Partial 状态:CANN ReduceScatter 已在内部完成 k 方向 sum + m 方向 scatter

支持的分片模式:

  • x1 Shard(1) on tp(k 维),x2 Shard(0) on tp(k 维)
  • trans_x2=True
  • 2D mesh (dp, tp),x1 Shard(1) on tp
  • 8 卡 (mp=2, np=2, tp=2),m/n/k 三轴全切

两个算子均继承 DistributedOp 基类,实现 preprocessinfer_layout,通过 YAML 注册。
不需要 get_expand_impl 覆盖(基类默认返回 None)。

Partial 说明:AllGatherMatmul 在 k 维切分时输出带 Partial(sum)full_tensor() 自动触发 AllReduce;MatmulReduceScatter 不需要 Partial(CANN 内核已内部完成 ReduceScatter,结果完整)。

约束:

算子 约束
AllGatherMatmul x1 k 维与 x2 k 维 placement 必须一致(均 Replicate 或均 Shard 在同一 mesh axis);k 被切分时输出带 Partial(sum);x1 m 维不支持多 mesh 联合分片(tuple tensor_map)
MatmulReduceScatter x1 k 维必须 Shard(TP);x2 k 维 placement 必须与 x1 k 维一致
两者共用 gather_index=0、trans_input=False、bias=None(当前 MindSpore 约束);不允许 Partial 输入

UT:test_parallel_all_gather_matmul.py(15 cases)、test_parallel_matmul_reduce_scatter.py(12 cases)
ST:4 个 MindSpore 分布式测试文件(AGM 3 组,MRS 3 组,含 8 卡 mnk 全切场景)

相关 PR

#656

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 152
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/152

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reviewing related PR #656, then read the DistributedOp base class and the YAML registration described in the issue. Run test_parallel_all_gather_matmul.py and test_parallel_matmul_reduce_scatter.py to inspect the 27 stated cases. Done means both operators implement preprocess and infer_layout and the listed distributed tests pass, including the 8-card scenarios.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems, machine-learning, testing-qa
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Clearly specified
Newbie friendliness
30/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.