mindspore-ai / mindspore-ai/hyper-parallel

【RFC】HyperParallel Symmetric Memory — 子 Group、完整集合通信、反向传播及 Hccl 互操作

Open
#658 2 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

【RFC】HyperParallel Symmetric Memory — 子 Group、完整集合通信、反向传播及 Hccl 互操作

对称内存基础原语、shmem_allgathershmem_alltoall(Torch)及 MC2 融合算子的设计已通过 #59 完成评审。本 ISSUE 在此基础上规划子 Group 支持、高级集合通信封装、反向传播、Hccl 互操作及 DFX 能力。


一、对外特性描述

1.1 背景

单边通信与集合通信内存拷贝链路对比

对称内存(Symmetric Memory)基于昇腾芯片内 Device 间共享内存,以 Push 模式实现单边通信——发起方可直接通过 Python API 读写远端对称内存,无需目标 rank 显式参与,消除了传统集合通信的同步屏障和中间数据拷贝。基础原语(put/get/signal/wait/put_with_signal)及 shmem_allgather 等接口设计详见 #59

Push模式主流程

当前待补齐的能力:

能力 PyTorch MindSpore
对称内存基础原语(put/get/signal/wait/put_with_signal)
shmem_allgather
shmem_alltoall
MC2 融合算子(AG+MatMul / MatMul+RS)
内存分配(empty / Manager.malloc / MemPool)
1.2 现状与问题

基于单边通信的MoE模块示意

当前实现存在三大能力缺口:

  1. 不支持子 Group——所有对称内存操作基于全局 world_size,无法用于 FSDP/HSDP 中不同 DP group、TP group 等场景。典型需求:DP group 内用对称内存做 AllReduce,TP group 内用 Hccl 做 AllGather,两组互不干扰。

  2. 缺少高级集合通信封装——用户需自行基于 put/get/signal 原语组装 AllReduce / ReduceScatter / Send / Recv,开发门槛高、容易出错。

  3. 无自动微分支持——shmem_allgather / shmem_alltoall 无法在训练中直接使用,因为缺少反向传播实现。

1.3 设计目标
目标 说明
对称内存子 Group 所有单边通信操作支持 group 参数,在指定子通信域内执行
高级集合通信接口 封装 shmem_allreduce / shmem_reduce_scatter / shmem_send / shmem_recv,对标 Hccl 语义
自动微分(反向传播) AllGather / AllReduce / ReduceScatter / AllToAll 的正反向配套实现
Hccl 混用互操作 同一训练任务中 Hccl 集合通信与对称内存通信可安全混用
低精度通信高精累加 AllReduce / ReduceScatter 支持低精度数据传输 + 高精度本地累加
DFX 可观测性 内存用量监控、日志分级、signal 超时检测、性能 Profiling 集成
跨平台补齐 MindSpore 侧 shmem_alltoall、MC2 融合算子补全
1.4 典型场景
场景 说明
MoE 通算掩盖 TP/EP group 内对称内存替代 Hccl,减少同步开销,与多核并行协同
FSDP/HSDP 梯度同步 DP group 内用 shmem_reduce_scatter / shmem_allreduce,降低延迟
序列并行 / Context Parallel CP group 内用对称内存通信,与计算流水线重叠
张量并行 MC2 子 Group MC2,通算掩盖效率进一步提升
参数 Broadcast shmem_send / shmem_recv 点对点通信,无 barrier
1.5 功能与规格详述
1.5.1 对称内存子 Group 支持

核心设计:所有对称内存通信接口新增可选参数 group,指定通信域。group=None 时使用全局通信组。

# 新增 group 参数(所有接口统一)
symm.shmem_put(target, target_offset, src, src_offset, size, target_rank, group=None)
symm.shmem_get(target, target_offset, src, src_offset, size, target_rank, group=None)
symm.shmem_signal_op(signal, signal_offset, signal_value, signal_op, target_rank, group=None)
symm.shmem_put_with_signal(target, target_offset, src, src_offset,
                            size, signal, signal_offset, signal_value, signal_op, target_rank,
                            group=None)
symm.shmem_allgather(output_tensor, input_tensor, group=None)
symm.shmem_alltoall(send_tensor_list, receive_tensor, receive_list, group=None)
symm.shmem_allreduce(tensor, op='sum', group=None)
symm.shmem_reduce_scatter(output_tensor, input_tensor, op='sum', group=None)
symm.shmem_send(tensor, dst_rank, group=None)
symm.shmem_recv(tensor, src_rank, group=None)

target_rank 语义变更:当指定 group 时,target_rank 变为 group 内局部 rank(0 ~ group.size-1),框架内部负责映射到全局 rank。

约束

  • group 内所有 rank 必须处于同一物理节点(对称内存仅在节点内 Device 间共享)。
  • 对已创建的 group 缓存其 rank 映射关系,避免重复初始化开销。
  • group 通过 hyper_parallel.collectives.split_group()platform.create_group() 创建,类型与平台无关(Torch: ProcessGroup, MindSpore: str)。
1.5.2 高级集合通信接口

在底层单边原语之上,封装对标 Hccl 的高级通信接口:

接口 说明 对标 Hccl 实现方式
shmem_allreduce(tensor, op, group) In-place AllReduce hcclAllReduce ReduceScatter → AllGather 两步,或直接 tree-based reduce
shmem_allgather(output, input, group) AllGather(已有,加 group) hcclAllGather Push 模式:各 rank 向所有 rank 写入数据 + signal 累加
shmem_alltoall(send_list, recv_tensor, recv_list, group) AllToAll(已有,加 group) hcclAlltoAllV Push 模式:先交换偏移元数据,再 push 各自 segment
shmem_reduce_scatter(output, input, op, group) ReduceScatter hcclReduceScatter Push 模式 + 本地 reduce 后 scatter 到各 rank
shmem_send(tensor, dst, group) 点对点 Send hcclSend shmem_put_with_signal 封装
shmem_recv(tensor, src, group) 点对点 Recv hcclRecv shmem_wait_for_signal 封装

op 参数支持'sum' / 'avg' / 'max' / 'min',与 Hccl hcclRedOp 对齐。'avg' 操作内部先 sum 再除以 group size。

低精度通信高精累加

  • AllReduce / ReduceScatter 支持 low_precision_dtype 参数,指定通信时使用的低精度数据类型(如 fp8 / fp16)。
  • 通信前本地 cast 到低精度 → 对称内存传输 → 通信后 cast 回高精度 → 本地累加。
  • 通过 SYMMETRIC_MEMORY_LOW_PRECISION_DTYPE 环境变量全局配置,或接口级覆盖。
1.5.3 反向传播(自动微分)

遵循经典通信-计算对偶关系:

Forward:  shmem_allgather(x)       →  y
Backward: shmem_reduce_scatter(grad_y) → grad_x

Forward:  shmem_reduce_scatter(x)  →  y
Backward: shmem_allgather(grad_y)  → grad_x

Forward:  shmem_allreduce(x)       →  y
Backward: Identity (grad_y 各 rank 已一致,直接回传)

Forward:  shmem_alltoall(x)        →  y
Backward: shmem_alltoall(grad_y, transpose=True) → grad_x

Forward:  shmem_send(x, dst)       →  (no output)
Backward: Identity (grad from recv side)

Forward:  shmem_recv(x, src)       →  y
Backward: shmem_send(grad_y, src)

实现方式:每个通信接口注册对应的 autograd.Function(Torch)或 custom bprop(MindSpore),group 参数通过 ctx 传递到反向。

1.5.4 与 Hccl 集合通信的互操作性

核心结论:可以混用,但需遵循约束。

维度 Hccl 集合通信 对称内存(SHMEM) 混用约束
底层资源 HCCS 链路 + HCCL 协议栈 Device 间共享内存 + P2P DMA 无资源冲突,可同时存在
同步方式 Barrier 隐式同步 Signal/Wait 显式同步 跨界同步需自行管理
Group 全局组 / 子 Group 新增 Group 支持 同一 Group 对象可被两种通信方式共用
自动微分 Hccl 提供可导封装 新增反向实现 梯度路径上可混用
Stream 依赖 框架管理 用户 / 框架管理 同一 Stream 上的两种通信操作顺序由 Stream 保证

推荐混用模式

训练循环:
  # DP group 内对称内存 AllReduce(低延迟、无 barrier)
  symm.shmem_allreduce(grads, op='sum', group=dp_group)

  # TP group 内 Hccl AllGather(利用 HCCS 高带宽)
  platform.all_gather_into_tensor(activations, tp_group_info)

  # 同步点:两类操作的 Stream 顺序自动保证正确性

禁止:在同一数据的同一通信阶段混用 Hccl 和对称内存(如用 Hccl AllReduce 聚合一部分梯度、用 SHMEM AllReduce 聚合另一部分),这会导致结果不一致。

1.5.5 跨平台补齐计划
功能 Torch MindSpore 优先级
子 Group 支持(全部接口) 待实现 待实现 P0
shmem_allreduce 待实现 待实现 P0
shmem_reduce_scatter 待实现 待实现 P0
shmem_send / shmem_recv 待实现 待实现 P1
反向传播(全部接口) 待实现 待实现 P0
shmem_alltoall (MindSpore) ✅ 已有 待实现 P0
MC2 融合算子 (MindSpore) ✅ 已有 待实现 P1
低精度通信高精累加 待实现 待实现 P1
对称内存用量监控 待实现 待实现 P1
1.5.6 性能与精度规格

性能指标(以 910B 8 卡节点为基准,nccl-tests 等价测试):

指标 目标 验证方式
SHMEM AllGather 带宽 ≥ Hccl AllGather 的 85%(消息 ≥ 1MB) shmem_allgather vs hcclAllGather perftest
SHMEM AllReduce 带宽 ≥ Hccl AllReduce 的 80%(消息 ≥ 1MB) shmem_allreduce vs hcclAllReduce perftest
SHMEM AllToAll 延迟 ≤ Hccl AllToAll 的 90%(MoE 典型 128KB~4MB) shmem_alltoall vs hcclAlltoAllV perftest
SHMEM ReduceScatter 带宽 ≥ Hccl ReduceScatter 的 80%(消息 ≥ 1MB) shmem_reduce_scatter vs hcclReduceScatter perftest
MC2 融合算子吞吐 ≥ 分离式通信+计算的 1.1x 典型 MoE 矩阵形状测试
反向通信耗时 ≤ 正向通信耗时的 1.2x 与正向对比
SHMEM+低精度累加带宽 ≥ 同消息大小 fp16 Hccl 带宽的 90% fp8 SHMEM + fp16 acc vs fp16 Hccl

精度对齐规格

场景 精度要求 备注
allreduce(input, 'sum') vs Hccl AllReduce fp32: bit-exact 或 allclose(rtol=1e-7, atol=1e-7) 若 reduce order 不同,按树形拓扑保证 bit-exact
allgather → reduce_scatter roundtrip 与原始输入一致 梯度路径正确性基础保证
fp16/bf16 通信 allclose(rtol=1e-3, atol=1e-5)
低精度通信高精累加 与全高精度 Hccl 结果 allclose(rtol=1e-3, atol=1e-5) fp8 通信 + fp16 累加 vs fp16 Hccl
MC2 融合算子 ≤ 分离式通信+计算的 1e-5(fp32)/ 1e-3(fp16/bf16)
1.5.7 DFX 能力
能力 说明 接口 / 环境变量
日志分级 对称内存运行时日志级别控制 HYPER_PARALLEL_SHMEM_LOG_LEVEL(DEBUG/INFO/WARNING/ERROR)
日志输出控制 日志输出目标(stdout / 文件) SHMEM_LOG_TO_STDOUT=1
内存用量监控 查询对称内存已分配/峰值/剩余量 symm_memory_stats(){"allocated": ..., "peak": ..., "free": ...}
Signal 超时检测 可配置超时,超时时输出 stuck rank 信息后退出 SYMMETRIC_MEMORY_SIGNAL_TIMEOUT_MS(默认 300000ms)
参数校验 非法参数前置校验,抛出明确错误信息 所有接口入参处
Profiling 集成 关键操作标记 CANN 事件,可被 msprof / torch_npu.profiler 采集 自动注入
可用性检测 运行时检测库是否可用,不可用时给出明确提示 symm.is_shmem_available()
版本兼容检查 初始化时校验 CANN 版本 ≥ 最低要求 自动检查

参考#59 Symmetric Memory 特性设计 RFC

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

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 reading the prior symmetric-memory design in issue #59 and the current Torch and MindSpore implementations of the listed communication interfaces. Break the RFC into sub-group support, higher-level collectives, automatic differentiation, Hccl interoperability, and DFX requirements before locating entry points. Done requires implementing and validating the specified APIs, gradient behavior, interoperability constraints, performance targets, precision criteria, and diagnostics.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems, machine-learning
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.