已开启
[基础能力3] 张量并行 TP 对标 PyTorch:差距梳理与补齐路线图 #6
changzherui创建于  6月30日
changzherui
changzherui成员
6月30日 创建

一、目标与背景

📎 姊妹专项: #2 DTensor / DeviceMesh 对标 PyTorch · #5 HSDP / FSDP 对标 PyTorch

📎 子专项(上游 mindspore/hyper-parallel): #269 loss_parallel / 分布式 CE · #270 TP redistribute 异步(ACT)

本 Issue 目标: 在保持 Hyper 现有架构不变的前提下,完善 张量并行(TP) 能力,使其在语义与 API 上尽可能对齐 PyTorch torch.distributed.tensor.parallel 及 TorchTitan 训练编排习惯,降低 PyTorch 生态迁移成本。

上级文档: 本 Issue 是 HyperParallel / PyTorch / TorchTitan 用户接口对标 在 TP 模块的子专项,聚焦「声明式 TP 策略 + 边界 layout 重分布 + loss parallel + 与 FSDP/AC 组合」层的差距分析与补齐计划。

分析依据(源码):

路径
PyTorch torch/distributed/tensor/parallel/tensor/parallel/loss.pyinput_reshard.pyddp.pyfsdp.py
HyperParallel hyper_parallel/core/tensor_parallel/platform/*/loss_parallel_ops.pycore/shard/
TorchTitan torchtitan/distributed/tensor_parallel.pymodels/common/decoder_sharding.pymodels/*/parallelize.pycomponents/loss.py

一致性图例(附录中使用): ✅ 一致 · ⚠️ 部分一致 · ❌ 仅一侧有 · 🔶 Hyper 扩展


二、范围与架构约束

2.1 模块范围

本 Issue 覆盖以下子模块的对标与补齐:

子模块 PyTorch Hyper TorchTitan
入口 API parallel/api.py core/tensor_parallel/api.py models/*/parallelize.py
并行策略 parallel/style.py core/tensor_parallel/style.py models/common/decoder_sharding.py
分布式 CE parallel/loss.py core/tensor_parallel/loss_parallel.py + platform/*/loss_parallel_ops.py components/loss.py_LossParallelCrossEntropy
AC + TP 输入重切 parallel/input_reshard.py 声明式 sharding 路径内处理
DP/FSDP + TP 衔接 parallel/ddp.pyparallel/fsdp.py fully_shard 与 DTensor 参数共存 apply_fsdp_to_decoder 在 TP 之后
底层分片(非 TP 专属) core/shard/api.pyshard_module)、DFunction

2.2 架构约束(补齐前提)

以下设计保持不变,补齐工作在此基础上演进:

  • 公开入口仍为 parallelize_module() + ParallelStyle.apply()改为 PyTorch 私有 _apply() 命名,也不引入 TorchTitan 新一代 ShardingConfig + Module.parallelize() 作为唯一路径
  • 双栈保留:同一套 ParallelStyleget_platform() 适配 PyTorch nn.Module 与 MindSpore Cell
  • 参数/激活切分仍经 distribute_module + forward hook 完成,改为全量 __torch_dispatch__
  • shard_module / DFunction 作为更低层扩展保留,与 declarative TP 并存

2.3 与 PyTorch 的三条结构性差异

理解以下差异,是阅读后文差距与路线图的前提:

  1. I/O 边界通信模型:PyTorch Col/Row/Seq 在 redistribute 时默认 async_op=True,collective 提交后由 PyTorch AsyncCollectiveTensor(ACT)__torch_dispatch__ 延迟 wait_tensor();Hyper 当前 全同步 redistribute,TP 边界尚未接入 ACT(PP/FSDP/MoE 路径已有 ACT 先例,见 #270)。架构要点:Hyper DTensor 走 __torch_function__ + _OP_DISPATCHER在 DTensor 上实现全套 __torch_dispatch__;异步 wait 由 ACT 子类承担,use_local_output=Trueto_local() 透传 ACT 给 nn.Linear 即可重叠。
  2. 参数初始化语义:TP styles 通过 distribute_tensor(..., src_data_rank=...) 从逻辑全局张量 scatter;Hyper distribute_tensor 当前 无 scatter/broadcast 通信#2 P0),导致 parallelize_module 接口虽已对齐,但「rank0 单份权重一次分片」路径仍不等价。
  3. 训练框架编排:TorchTitan 主流模型已迁移到声明式 ShardingConfig + Module.parallelize(),生产 loss 用 _LossParallelCrossEntropy 而非 loss_parallel() 上下文;Hyper 提供 examples/torch/llama3/parallelize.py直接调用 parallelize_module 的参考实现,尚无 Titan 级 parallelize_qwen3 全模型 TP 编排(如 Qwen3.5 线性注意力层)。

2.4 并行模式速览

策略 权重切分 激活布局 通信 典型模块
ColwiseParallel Linear Shard(0);Embedding Shard(1) Replicate → 出 Shard(-1) 列切分 matmul Q/K/V、gate/up
RowwiseParallel Linear Shard(1);Embedding Shard(0) Shard(-1) → 出 Replicate/Partial all-reduce / partial O、down
SequenceParallel Replicate Shard(seq_dim) 序列维切分 LayerNorm、Dropout
PrepareModuleInput/Output 不改参 边界 annotate + redistribute 布局转换 Attention 入口、FFN 入口
NoParallel 🔶 Replicate 可配 I/O layout 无切分 MoE Router、未开 SP 的 Norm
典型 Transformer Block(SP 开启,Titan / Hyper llama3 示例)

  tok_embeddings: Rowwise → Shard(1) on seq
  attention_norm: SequenceParallel
  attention:      PrepareModuleInput(Shard(1) → Replicate)
  wq/wk/wv:       ColwiseParallel
  wo:             RowwiseParallel → Shard(1)
  ffn_norm:       SequenceParallel
  feed_forward:   PrepareModuleInput(Shard(1) → Replicate)
  w1/w3:          ColwiseParallel
  w2:             RowwiseParallel → Shard(1)
  output/lm_head: ColwiseParallel → Shard(-1) [loss parallel]

2.5 算子分发与异步 wait 架构

层次 Hyper PyTorch DTensor
DTensor 算子拦截 DTensorBase.__torch_function___OP_DISPATCHER(YAML 注册表) DTensor.__torch_dispatch__(全量 ATen)
TP 边界 pending 张量 redistribute(async_op=True) → ACT(#270 待补) AsyncCollectiveTensor
延迟 wait 触发 ACT __torch_dispatch__(下游 Linear/mm 等) 同左
CE 拦截 _OP_DISPATCHER + _ce_op_registry#269 DTensor custom handler
MindSpore __ms_dispatch__(MS AsyncCollectiveTensor
推荐异步路径(#270,Torch 后端):

  Colwise post-hook: redistribute(async_op=True) → DTensor._local_tensor = ACT
  use_local_output=True: to_local() → ACT(plain tensor)
  nn.Linear(ACT): ACT.__torch_dispatch__ → wait → matmul(与 collective 重叠)

不推荐:给 Hyper DTensorBase 加全套 __torch_dispatch__(与现有 OpDispatcher 架构冲突)

三、对标结论(速览)

3.1 已对齐的能力(可直接复用思路)

能力 说明
parallelize_module(module, mesh, parallelize_plan) 1-D mesh、单 style 或 dict + fnmatch glob、原地修改
src_data_rank 写入 style 后传入 distribute_tensorNone 表示仅用本地数据
ColwiseParallel / RowwiseParallel / SequenceParallel 参数切分维、默认 I/O layout、use_local_output 语义与 PyTorch 一致
PrepareModuleInput / PrepareModuleOutput / PrepareModuleInputOutput 仅 hook、不改参;支持 kwarg layout
distribute_module 三阶段 partition_fn + input_fn + output_fn;防重入 _distribute_module_applied
loss_parallel 上下文 拦截 cross_entropy 等 CE 入口,分布式 vocab-shard softmax
ParallelDims.tp + get_mesh("tp") Hyper trainer.parallel_dims 移植自 Titan,mesh 轴名一致
Llama3 风格 TP 编排 examples/torch/llama3/parallelize.py 对标 Titan apply_tp 路径

3.2 Hyper 独有能力(保留,不作为补齐对象)

  • NoParallel:参数 replicate + DTensor 边界;Titan 在 torchtitan/distributed/tensor_parallel.py 有同名实现,但 PyTorch 官方 tensor.parallel 包未导出
  • loss_parallel(mesh=..., strict=...) + is_loss_parallel_active():可显式指定 mesh、非严格模式
  • MindSpore 双栈Dense/Embedding/Cell 与 PyTorch Linear/Embedding/Module 同一套 style
  • shard_module + ShardingPlan:多维 mesh 别名("dp", "tp")的低层分片,与 declarative TP 互补
  • DFunction:用户自定义分布式 autograd,经 _OP_DISPATCHER 与 TP 算子生态衔接
  • 推理侧 gather_tensor_parallel_logitshyper_parallel/infer/ 用于生成阶段 vocab gather

3.3 按场景一致性总表

用户场景 一致性 说明
单层 Linear Col/Row + 1D tp mesh 与 PyTorch 用法一致
parallelize_plan dict + layers.* glob fnmatch 规则一致
Sequence Parallel + Col/Row 链式 use_local_output=False 保持 DTensor 传递
loss_parallel() + vocab-shard logits ⚠️ 内核已有;多维 mesh / 框架组件差异见 #269
Col/Row/Seq 边界通信-计算重叠 全同步 redistribute;见 #270
rank0 加载全局权重 → parallelize_module ⚠️ 依赖 #2 P0 distribute_tensor 通信
TP + fully_shard(先 TP 后 FSDP) ⚠️ 有 e2e 测试;缺 PyTorch DTensorExtensions 级 state_dict/flatten 辅助
AC + TP(非 SP 模块输入重切) PyTorch input_reshard;Hyper 无
Async TP(compile micro-pipeline) Titan maybe_enable_async_tp;Hyper 无
DDP + TP PyTorch _pre_dp_module_transform;Hyper 训练主路径走 FSDP
Titan 声明式 Module.parallelize() Hyper 无 ShardingConfig / SpmdLayout 体系
Qwen3.5 全模型 TP(含线性注意力) parallelize_qwen3_5 显式 NotImplementedError
MoE Router + TP ⚠️ Hyper/Titan 均用 NoParallel;需与 EP 编排联调
MindSpore 后端 TP 训练 🔶 Hyper 独有;pre-hook 须返回 tuple

3.4 迁移三原则

  1. 类名与 import 基本一致from hyper_parallel import parallelize_module, ColwiseParallel, ...loss_parallelfrom hyper_parallel.core.tensor_parallel import loss_parallel(顶层 __init__ 暂未导出)。
  2. 多维 mesh 先切片parallelize_module(model, mesh["tp"], plan);Hyper 对 ndim > 1 显式抛错(与 PyTorch 文档约束一致,但 PyTorch 运行时未必强制)。
  3. 编排顺序与 Titan 相同CP(可选)→ TP → AC → compile(可选)→ FSDP;勿在 TP 之前 fully_shard

四、能力差距(PyTorch 可达 · Hyper 未达)

在架构约束下,以下差距需通过功能补齐弥合;❌ 不支持 · ⚠️ 有替代但不等价。

4.1 阻塞级差距(对应 P0)

差距 PyTorch Hyper 影响
distribute_tensor scatter 语义 TP 参数从 rank0 切分 ❌ 仅本地 slice(#2 P0 parallelize_module 权重初始化
TP I/O redistribute(async_op=True) Col/Row/Seq 默认异步 ❌ 全同步 通信-计算重叠、AsyncTP 前置;→ #270
loss_parallel 顶层导出 tensor.parallel.loss_parallel ⚠️ 子包有、未进 __all__ #269 P0#1
多维 mesh loss_parallel 支持 batch/CP 非 TP 维 ⚠️ 实现偏 1D TP mesh #269 P0#2~#4

4.2 高级并行差距(对应 P1)

差距 PyTorch Hyper 影响
input_reshard AC + 非 SP 时 backward 恢复 replicate 输入 AC 下 TP 显存/正确性
DTensorExtensions(FSDP+TP) tensor/parallel/fsdp.py flatten param、DCP 元数据
_pre_dp_module_transform DDP 自动 localize DTensor 参数 DDP+TP(Hyper 主路径不用 DDP)
maybe_enable_async_tp Inductor micro-pipeline compile 场景 TP 性能
Attention 内 local_map / head 维 SP Titan decoder_sharding ⚠️ llama3 示例手写 PrepareModuleInput 复杂 attention 变体
CE 特性 parity label_smoothing、多维 reduction ⚠️ Hyper 限制更多 训练配置迁移
Titan _LossParallelCrossEntropy 无上下文、ChunkedLoss 集成 ⚠️ 可用 loss_parallel 替代 #269 P1~P2

4.3 完备性与生态差距(对应 P2)

差距 PyTorch Hyper
全模型 parallelize_<model> Titan llama/qwen/deepseek/gpt_oss llama3 示例;qwen3_5 无 TP
声明式 ShardingConfig Titan Module.parallelize()
gather_tensor_parallel_logits 与训练 CE 统一 Titan 在 loss 组件内 ⚠️ 训练/推理分离实现
TP + EP + FSDP 参考编排 Titan MoE decoder_sharding examples/moe 部分覆盖
动转静 TP compile + AsyncTP 推进中 ❌(README 标注待实现)
parallelize_plan=None auto-plan 双方均 warn + no-op 文档需强调

4.4 明确不在本 Issue 补齐范围

原因
改为 PyTorch _apply + MRO 风格 仅命名差异,无功能收益
完整移植 TorchTitan ShardingConfig / SpmdLayout 属训练框架层,非 TP 原语
DDP + TP 一等支持 Hyper 训练主路径为 FSDP/HSDP
XLA / torch_xla TP 文档已声明不支持

五、补齐路线图(P0 → P1 → P2)

原则: 不改 ParallelStyle + distribute_module + 双栈架构,在现有 parallelize_module 上增量补齐。

5.1 优先级总览

优先级 目标 项数
P0 阻塞迁移 / TP 初始化正确性 4
P1 AC/FSDP 组合与性能 6
P2 生态完备、模型编排 5

5.2 P0 — 必须补

# 任务 现状 补齐方向 工作量
1 distribute_tensor 通信 仅本地 slice #2 P0#1 联动;TP partition_fn 验证 中(依赖 DTensor)
2 loss_parallel 顶层导出 子包 only #269 P0#1
3 TP 边界同步 redistribute 全同步 #270(ACT + async_op 中~大
4 多维 mesh CE 偏 1D #269 P0#2~#4

5.3 P1 — 应补

# 任务 现状 补齐方向 工作量
5 input_reshard 移植 saved_tensors_hooks pack/unpack 逻辑
6 FSDP+TP DTensorExtensions 对等 参数 flatten/chunk、与 fully_shard load 桥接 中~大
7 maybe_enable_async_tp 对等 NPU compile 路径评估 + 文档
8 CE 特性对齐 缺 smoothing 等 #269 P1~P2 小~中
9 Attention TP 辅助 手写 Prepare 抽取 PrepareAttentionTP 或文档化 Titan 映射
10 MoE NoParallel + EP 联调 分散在 examples 与 expert_parallel 指南统一

5.4 P2 — 建议补

# 任务 补齐方向 工作量
11 parallelize_qwen3_5 TP 线性注意力层 Col/Row 策略
12 更多模型 parallelize_* deepseek_v3、gpt_oss 等对标 Titan 中~大
13 ChunkedLoss + loss_parallel 集成 #269 P2#9~#10
14 推理/训练 logits gather 统一 gather_tensor_parallel_logits 与 CE 共享语义
15 TP 动转静局部验证 与 README 路线图联动

5.5 里程碑

M1(TP 可迁移)     P0 #1 + [#269](https://gitcode.com/mindspore/hyper-parallel/issues/269) P0#1
M2(组合训练稳定)  [#269](https://gitcode.com/mindspore/hyper-parallel/issues/269) P0#2~#4 + P1 #5 #6
M2b(通信重叠)     [#270](https://gitcode.com/mindspore/hyper-parallel/issues/270) M1~M2
M3(性能与 MoE)    P1 #7 #9 #10 + [#270](https://gitcode.com/mindspore/hyper-parallel/issues/270) M3
M4(生态完备)      P2 #11~#15 + [#269](https://gitcode.com/mindspore/hyper-parallel/issues/269) P2

5.6 执行优先级结论

  1. distribute_tensor 通信(#2 P0)+ #269 P0 — 不改架构下解锁「加载权重 → parallelize → 训练」主路径。
  2. #270 ACT 异步 redistribute + input_reshard — 解锁通信重叠与 AC+TP;复用 PyTorch ACT,扩展 Hyper DTensor 为 __torch_dispatch__
  3. P2 模型编排 + #269 框架 loss — M2 后与 FSDP/EP 专项并行推进。

附录 A:逐类逐接口对照

详细技术参考;日常决策优先看第三~五章。

A.1 入口 API:parallelize_module

参数 / 行为 PyTorch Hyper 一致性 功能描述
module nn.Module Module(Torch/MS) 根模块
device_mesh 1D;None → 当前 mesh 同左;ndim>1 抛错 TP 拓扑
parallelize_plan ParallelStyle | dict 同左 声明式计划
fnmatch glob layers.*.mlp
src_data_rank ⚠️ 接口有;底层通信待 #2 P0
parallelize_plan=None warn + no-op 同左 无 auto-plan
style 方法名 _apply apply ⚠️ 纯命名
in-place style.src_data_rank 调用方需注意突变

A.2 并行策略类

ColwiseParallel

PyTorch Hyper 一致性
Linear weight Shard(0) Shard(0)
Embedding weight Shard(1) Shard(1)
默认 input Replicate() Replicate()
默认 output Shard(-1) Shard(-1)
I/O redistribute async_op=True 同步 ⚠️
支持模块 Linear, Embedding + MS Dense/Embedding ⚠️

RowwiseParallel

PyTorch Hyper 一致性
Linear weight / bias Shard(1) / Replicate() 同左
Embedding weight Shard(0) Shard(0)
Embedding 输出 Partial Partial("sum")
默认 input Shard(-1)(Emb 为 Replicate 同左

SequenceParallel

PyTorch Hyper 一致性
参数 Replicate Replicate
sequence_dim 默认 1 默认 1
use_local_output 默认 False 默认 False

PrepareModuleInput / PrepareModuleOutput / PrepareModuleInputOutput

PyTorch Hyper 一致性
机制 forward pre/post hook 同左
kwarg layouts
redistribute async_op=True 同步 ⚠️
pre-hook 返回 tuple tuple(MS 强制) ⚠️

NoParallel 🔶

PyTorch 官方包 Hyper TorchTitan
导出 hyper_parallel distributed/tensor_parallel.py
用途 MoE Router、replicate Norm 同左

A.3 loss_parallel

接口 / 行为 PyTorch Hyper 一致性 功能描述
上下文 API loss_parallel() 无参 loss_parallel(mesh, strict) ⚠️ Hyper 扩展参数
顶层导出
拦截方式 DTensor.__torch_dispatch__ custom handler __torch_function___OP_DISPATCHER + CE registry ⚠️ 机制不同、语义对齐;详见 #269
is_loss_parallel_active 🔶
label_smoothing
多维 mesh reduction ⚠️ ⚠️
Titan 生产路径 _LossParallelCrossEntropy 不用上下文

A.4 组合层 API(PyTorch 独有 · 本 Issue 关注)

接口 PyTorch Hyper 一致性 功能描述
input_reshard AC backward 输入恢复
_pre_dp_module_transform DDP+TP 参数 localize
DTensorExtensions FSDP flatten/chunk
maybe_enable_async_tp Titan 封装 compile AsyncTP

A.5 Hyper 底层扩展(非 PyTorch TP 包)

接口 与 TP 关系 说明
shard_module 可表达 TP layout 多维 mesh 别名;更低层
DFunction 自定义 TP 算子 _OP_DISPATCHER 注册
parallelize_value_and_grad 独立 API parallelize_module 路径
gather_tensor_parallel_logits 推理 infer/ 模块

A.6 运行时生命周期对照

初始化
  parallelize_module(model, tp_mesh, plan)
  → 递归匹配子模块 → style.apply
  → distribute_module(partition_fn, input_fn, output_fn)
      partition_fn: distribute_tensor(weight, mesh, Shard(...), src_data_rank)
      未切分 param/buffer → Replicate DTensor
  → 根模块注册 forward pre/post hook

Forward
  pre-hook: plain Tensor → DTensor.from_local → redistribute(desired_input)
  → 子模块计算(Linear/Embedding 等经 DTensor dispatch)
  post-hook: redistribute(output_layouts) → optional to_local()

Backward
  梯度经 DTensor autograd + Partial 归约
  Rowwise/Embedding 路径: Partial → all-reduce
  loss_parallel 内: 融合 CE backward,避免额外 all_gather

与 FSDP 组合(推荐顺序)
  parallelize_module(..., mesh["tp"], ...)   # 参数已是 TP 维 DTensor
  → checkpoint_wrapper(可选;非 SP 时需 input_reshard — Hyper 待补)
  → fully_shard(layer, mesh["fsdp"])       # 在 dp_shard 维再切分
要点 PyTorch TP Hyper TP
策略应用 ParallelStyle._apply ParallelStyle.apply
I/O 通信 redistribute(async_op=True) → ACT 同步 redistribute(→ #270
算子分发 DTensor.__torch_dispatch__ __torch_function___OP_DISPATCHER
异步 wait ACT __torch_dispatch__ 待 ACT 接入(#270)
参数来源 distribute_tensor scatter #2 P0
CE loss_parallel() 或 Titan 自定义 Function loss_parallel(mesh, strict)
双后端 Torch only Torch + MindSpore

A.7 与 TorchTitan 的关系

  • TorchTitan 主流模型不直接调用 parallelize_module,而用 ShardingConfig + Module.parallelize() 复刻 Col/Row/SP 语义;实验后端 experiments/transformers_modeling_backend/parallelize.py 仍用 PyTorch parallelize_module
  • Hyper examples/torch/llama3/parallelize.py 与 Titan models/llama3/parallelize.py apply_tp 路径高度对应,可作为迁移模板。
  • Titan 生产 loss:不用 with loss_parallel(),而用 components/loss.py_LossParallelCrossEntropy + ChunkedLossWrapper;Hyper 用户可在 trainer 中包 loss_parallel() 或后续补齐对等 wrapper。
  • Titan enable_async_tensor_parallel 依赖 torch.compile;Hyper 暂无对等开关。
  • Titan ParallelDims.seq_len_divisor = tp * (cp * 2)(SP 开启时);Hyper trainer.ParallelDims 移植同一约束。

A.8 源码索引(便于 Review)

主题 PyTorch Hyper TorchTitan
入口 tensor/parallel/api.py core/tensor_parallel/api.py models/*/parallelize.py
策略 tensor/parallel/style.py core/tensor_parallel/style.py models/common/decoder_sharding.py
CE tensor/parallel/loss.py core/tensor_parallel/loss_parallel.py components/loss.py
AC+TP tensor/parallel/input_reshard.py 声明式路径内
FSDP+TP tensor/parallel/fsdp.py fully_shard + DTensor 参数 distributed/fsdp.py
NoParallel core/tensor_parallel/style.py distributed/tensor_parallel.py
Llama3 示例 models/llama3/parallelize.py examples/torch/llama3/parallelize.py models/llama3/parallelize.py

父文档:HyperParallel / PyTorch / TorchTitan 用户接口对标 · 姊妹专项:DTensor / DeviceMesh #2 · HSDP / FSDP #5 · 子专项:loss_parallel #269 · TP 异步 ACT #270

likedislike
changzheruichangzherui成员
7月1日 修改了issue 的描述
changzherui
changzherui成员
7月1日 评论:

已补充 §3 张量并行 TP 全量接口对标总表(接口索引附录):#9 张量并行 TP 全量接口对标总表

likedislike