📎 姊妹专项: #2 DTensor / DeviceMesh 对标 PyTorch · #5 HSDP / FSDP 对标 PyTorch 📎 子专项(上游 mindspore/hyper-parallel): #269 loss_parallel / 分布式 CE · #270 TP redistribute 异步(ACT)
📎 姊妹专项: #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 生态迁移成本。
torch.distributed.tensor.parallel
上级文档: 本 Issue 是 HyperParallel / PyTorch / TorchTitan 用户接口对标 在 TP 模块的子专项,聚焦「声明式 TP 策略 + 边界 layout 重分布 + loss parallel + 与 FSDP/AC 组合」层的差距分析与补齐计划。
分析依据(源码):
torch/distributed/tensor/parallel/
tensor/parallel/loss.py
input_reshard.py
ddp.py
fsdp.py
hyper_parallel/core/tensor_parallel/
platform/*/loss_parallel_ops.py
core/shard/
torchtitan/distributed/tensor_parallel.py
models/common/decoder_sharding.py
models/*/parallelize.py
components/loss.py
一致性图例(附录中使用): ✅ 一致 · ⚠️ 部分一致 · ❌ 仅一侧有 · 🔶 Hyper 扩展
本 Issue 覆盖以下子模块的对标与补齐:
parallel/api.py
core/tensor_parallel/api.py
parallel/style.py
core/tensor_parallel/style.py
parallel/loss.py
core/tensor_parallel/loss_parallel.py
_LossParallelCrossEntropy
parallel/input_reshard.py
parallel/ddp.py
parallel/fsdp.py
fully_shard
apply_fsdp_to_decoder
core/shard/api.py
shard_module
DFunction
以下设计保持不变,补齐工作在此基础上演进:
parallelize_module()
ParallelStyle.apply()
_apply()
ShardingConfig
Module.parallelize()
ParallelStyle
get_platform()
nn.Module
Cell
distribute_module
__torch_dispatch__
理解以下差异,是阅读后文差距与路线图的前提:
redistribute
async_op=True
AsyncCollectiveTensor
wait_tensor()
__torch_function__
_OP_DISPATCHER
use_local_output=True
to_local()
nn.Linear
distribute_tensor(..., src_data_rank=...)
distribute_tensor
parallelize_module
loss_parallel()
examples/torch/llama3/parallelize.py
parallelize_qwen3
Shard(0)
Shard(1)
Replicate
Shard(-1)
Partial
Shard(seq_dim)
典型 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]
DTensorBase.__torch_function__
DTensor.__torch_dispatch__
redistribute(async_op=True)
Linear
mm
_ce_op_registry
__ms_dispatch__
推荐异步路径(#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 架构冲突)
parallelize_module(module, mesh, parallelize_plan)
dict
src_data_rank
None
ColwiseParallel
RowwiseParallel
SequenceParallel
use_local_output
PrepareModuleInput
PrepareModuleOutput
PrepareModuleInputOutput
partition_fn
input_fn
output_fn
_distribute_module_applied
loss_parallel
cross_entropy
ParallelDims.tp
get_mesh("tp")
trainer.parallel_dims
apply_tp
NoParallel
tensor.parallel
loss_parallel(mesh=..., strict=...)
is_loss_parallel_active()
Dense
Embedding
Module
ShardingPlan
"dp"
"tp"
gather_tensor_parallel_logits
hyper_parallel/infer/
tp
parallelize_plan
layers.*
use_local_output=False
DTensorExtensions
input_reshard
maybe_enable_async_tp
_pre_dp_module_transform
SpmdLayout
parallelize_qwen3_5
NotImplementedError
from hyper_parallel import parallelize_module, ColwiseParallel, ...
from hyper_parallel.core.tensor_parallel import loss_parallel
__init__
parallelize_module(model, mesh["tp"], plan)
ndim > 1
CP(可选)→ TP → AC → compile(可选)→ FSDP
在架构约束下,以下差距需通过功能补齐弥合;❌ 不支持 · ⚠️ 有替代但不等价。
tensor.parallel.loss_parallel
__all__
tensor/parallel/fsdp.py
local_map
decoder_sharding
label_smoothing
parallelize_<model>
parallelize_plan=None
_apply
torch_xla
原则: 不改 ParallelStyle + distribute_module + 双栈架构,在现有 parallelize_module 上增量补齐。
async_op
saved_tensors_hooks
PrepareAttentionTP
parallelize_*
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
详细技术参考;日常决策优先看第三~五章。
module
device_mesh
ndim>1
ParallelStyle | dict
layers.*.mlp
apply
style.src_data_rank
Replicate()
Partial("sum")
sequence_dim
hyper_parallel
distributed/tensor_parallel.py
loss_parallel(mesh, strict)
is_loss_parallel_active
parallelize_value_and_grad
infer/
初始化 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 维再切分
ParallelStyle._apply
ParallelStyle.apply
experiments/transformers_modeling_backend/parallelize.py
models/llama3/parallelize.py
with loss_parallel()
ChunkedLossWrapper
enable_async_tensor_parallel
torch.compile
ParallelDims.seq_len_divisor = tp * (cp * 2)
trainer.ParallelDims
tensor/parallel/api.py
tensor/parallel/style.py
tensor/parallel/input_reshard.py
distributed/fsdp.py
父文档:HyperParallel / PyTorch / TorchTitan 用户接口对标 · 姊妹专项:DTensor / DeviceMesh #2 · HSDP / FSDP #5 · 子专项:loss_parallel #269 · TP 异步 ACT #270
已补充 §3 张量并行 TP 全量接口对标总表(接口索引附录):#9 张量并行 TP 全量接口对标总表
一、目标与背景
本 Issue 目标: 在保持 Hyper 现有架构不变的前提下,完善 张量并行(TP) 能力,使其在语义与 API 上尽可能对齐 PyTorch
torch.distributed.tensor.parallel及 TorchTitan 训练编排习惯,降低 PyTorch 生态迁移成本。上级文档: 本 Issue 是 HyperParallel / PyTorch / TorchTitan 用户接口对标 在 TP 模块的子专项,聚焦「声明式 TP 策略 + 边界 layout 重分布 + loss parallel + 与 FSDP/AC 组合」层的差距分析与补齐计划。
分析依据(源码):
torch/distributed/tensor/parallel/、tensor/parallel/loss.py、input_reshard.py、ddp.py、fsdp.pyhyper_parallel/core/tensor_parallel/、platform/*/loss_parallel_ops.py、core/shard/torchtitan/distributed/tensor_parallel.py、models/common/decoder_sharding.py、models/*/parallelize.py、components/loss.py一致性图例(附录中使用): ✅ 一致 · ⚠️ 部分一致 · ❌ 仅一侧有 · 🔶 Hyper 扩展
二、范围与架构约束
2.1 模块范围
本 Issue 覆盖以下子模块的对标与补齐:
parallel/api.pycore/tensor_parallel/api.pymodels/*/parallelize.pyparallel/style.pycore/tensor_parallel/style.pymodels/common/decoder_sharding.pyparallel/loss.pycore/tensor_parallel/loss_parallel.py+platform/*/loss_parallel_ops.pycomponents/loss.py(_LossParallelCrossEntropy)parallel/input_reshard.pyparallel/ddp.py、parallel/fsdp.pyfully_shard与 DTensor 参数共存apply_fsdp_to_decoder在 TP 之后core/shard/api.py(shard_module)、DFunction2.2 架构约束(补齐前提)
以下设计保持不变,补齐工作在此基础上演进:
parallelize_module()+ParallelStyle.apply(),不改为 PyTorch 私有_apply()命名,也不引入 TorchTitan 新一代ShardingConfig+Module.parallelize()作为唯一路径ParallelStyle经get_platform()适配 PyTorchnn.Module与 MindSporeCelldistribute_module+ forward hook 完成,不改为全量__torch_dispatch__shard_module/DFunction作为更低层扩展保留,与 declarative TP 并存2.3 与 PyTorch 的三条结构性差异
理解以下差异,是阅读后文差距与路线图的前提:
redistribute时默认async_op=True,collective 提交后由 PyTorchAsyncCollectiveTensor(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=True时to_local()透传 ACT 给nn.Linear即可重叠。distribute_tensor(..., src_data_rank=...)从逻辑全局张量 scatter;Hyperdistribute_tensor当前 无 scatter/broadcast 通信(#2 P0),导致parallelize_module接口虽已对齐,但「rank0 单份权重一次分片」路径仍不等价。ShardingConfig+Module.parallelize(),生产 loss 用_LossParallelCrossEntropy而非loss_parallel()上下文;Hyper 提供examples/torch/llama3/parallelize.py等 直接调用parallelize_module的参考实现,尚无 Titan 级parallelize_qwen3全模型 TP 编排(如 Qwen3.5 线性注意力层)。2.4 并行模式速览
Shard(0);EmbeddingShard(1)Replicate→ 出Shard(-1)Shard(1);EmbeddingShard(0)Shard(-1)→ 出Replicate/PartialReplicateShard(seq_dim)Replicate2.5 算子分发与异步 wait 架构
DTensorBase.__torch_function__→_OP_DISPATCHER(YAML 注册表)DTensor.__torch_dispatch__(全量 ATen)redistribute(async_op=True)→ ACT(#270 待补)AsyncCollectiveTensor__torch_dispatch__(下游Linear/mm等)_OP_DISPATCHER+_ce_op_registry(#269)__ms_dispatch__(MSAsyncCollectiveTensor)三、对标结论(速览)
3.1 已对齐的能力(可直接复用思路)
parallelize_module(module, mesh, parallelize_plan)dict+ fnmatch glob、原地修改src_data_rankdistribute_tensor;None表示仅用本地数据ColwiseParallel/RowwiseParallel/SequenceParalleluse_local_output语义与 PyTorch 一致PrepareModuleInput/PrepareModuleOutput/PrepareModuleInputOutputdistribute_module三阶段partition_fn+input_fn+output_fn;防重入_distribute_module_appliedloss_parallel上下文cross_entropy等 CE 入口,分布式 vocab-shard softmaxParallelDims.tp+get_mesh("tp")trainer.parallel_dims移植自 Titan,mesh 轴名一致examples/torch/llama3/parallelize.py对标 Titanapply_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、非严格模式Dense/Embedding/Cell与 PyTorchLinear/Embedding/Module同一套 styleshard_module+ShardingPlan:多维 mesh 别名("dp","tp")的低层分片,与 declarative TP 互补DFunction:用户自定义分布式 autograd,经_OP_DISPATCHER与 TP 算子生态衔接gather_tensor_parallel_logits:hyper_parallel/infer/用于生成阶段 vocab gather3.3 按场景一致性总表
tpmeshparallelize_plandict +layers.*globuse_local_output=False保持 DTensor 传递loss_parallel()+ vocab-shard logitsparallelize_moduledistribute_tensor通信fully_shard(先 TP 后 FSDP)DTensorExtensions级 state_dict/flatten 辅助input_reshard;Hyper 无maybe_enable_async_tp;Hyper 无_pre_dp_module_transform;Hyper 训练主路径走 FSDPModule.parallelize()ShardingConfig/SpmdLayout体系parallelize_qwen3_5显式NotImplementedErrorNoParallel;需与 EP 编排联调3.4 迁移三原则
from hyper_parallel import parallelize_module, ColwiseParallel, ...;loss_parallel需from hyper_parallel.core.tensor_parallel import loss_parallel(顶层__init__暂未导出)。parallelize_module(model, mesh["tp"], plan);Hyper 对ndim > 1显式抛错(与 PyTorch 文档约束一致,但 PyTorch 运行时未必强制)。CP(可选)→ TP → AC → compile(可选)→ FSDP;勿在 TP 之前fully_shard。四、能力差距(PyTorch 可达 · Hyper 未达)
4.1 阻塞级差距(对应 P0)
distribute_tensorscatter 语义parallelize_module权重初始化redistribute(async_op=True)loss_parallel顶层导出tensor.parallel.loss_parallel__all__loss_parallel4.2 高级并行差距(对应 P1)
input_reshardDTensorExtensions(FSDP+TP)tensor/parallel/fsdp.py_pre_dp_module_transformmaybe_enable_async_tplocal_map/ head 维 SPdecoder_shardingPrepareModuleInputlabel_smoothing、多维 reduction_LossParallelCrossEntropyloss_parallel替代4.3 完备性与生态差距(对应 P2)
parallelize_<model>ShardingConfigModule.parallelize()gather_tensor_parallel_logits与训练 CE 统一decoder_shardingparallelize_plan=Noneauto-plan4.4 明确不在本 Issue 补齐范围
_apply+ MRO 风格ShardingConfig/SpmdLayouttorch_xlaTP五、补齐路线图(P0 → P1 → P2)
5.1 优先级总览
5.2 P0 — 必须补
distribute_tensor通信partition_fn验证loss_parallel顶层导出redistributeasync_op)5.3 P1 — 应补
input_reshardsaved_tensors_hookspack/unpack 逻辑DTensorExtensions对等fully_shardload 桥接maybe_enable_async_tp对等PrepareAttentionTP或文档化 Titan 映射NoParallel+ EP 联调5.4 P2 — 建议补
parallelize_qwen3_5TPparallelize_*loss_parallel集成gather_tensor_parallel_logits与 CE 共享语义5.5 里程碑
5.6 执行优先级结论
distribute_tensor通信(#2 P0)+ #269 P0 — 不改架构下解锁「加载权重 → parallelize → 训练」主路径。input_reshard— 解锁通信重叠与 AC+TP;复用 PyTorch ACT,不扩展 Hyper DTensor 为__torch_dispatch__。附录 A:逐类逐接口对照
A.1 入口 API:
parallelize_modulemodulenn.ModuleModule(Torch/MS)device_meshNone→ 当前 meshndim>1抛错parallelize_planParallelStyle | dictlayers.*.mlpsrc_data_rankparallelize_plan=None_applyapplystyle.src_data_rankA.2 并行策略类
ColwiseParallelShard(0)Shard(0)Shard(1)Shard(1)Replicate()Replicate()Shard(-1)Shard(-1)redistributeasync_op=TrueLinear,EmbeddingDense/EmbeddingRowwiseParallelShard(1)/Replicate()Shard(0)Shard(0)Partial("sum")Shard(-1)(Emb 为Replicate)SequenceParallelReplicateReplicatesequence_dimuse_local_outputPrepareModuleInput/PrepareModuleOutput/PrepareModuleInputOutputredistributeasync_op=TrueNoParallel🔶hyper_paralleldistributed/tensor_parallel.pyA.3
loss_parallelloss_parallel()无参loss_parallel(mesh, strict)DTensor.__torch_dispatch__custom handler__torch_function__→_OP_DISPATCHER+ CE registryis_loss_parallel_activelabel_smoothing_LossParallelCrossEntropyA.4 组合层 API(PyTorch 独有 · 本 Issue 关注)
input_reshard_pre_dp_module_transformDTensorExtensionsmaybe_enable_async_tpA.5 Hyper 底层扩展(非 PyTorch TP 包)
shard_moduleDFunction_OP_DISPATCHER注册parallelize_value_and_gradparallelize_module路径gather_tensor_parallel_logitsinfer/模块A.6 运行时生命周期对照
ParallelStyle._applyParallelStyle.applyredistribute(async_op=True)→ ACTredistribute(→ #270)DTensor.__torch_dispatch____torch_function__→_OP_DISPATCHER__torch_dispatch__distribute_tensorscatterloss_parallel()或 Titan 自定义 Functionloss_parallel(mesh, strict)A.7 与 TorchTitan 的关系
parallelize_module,而用ShardingConfig+Module.parallelize()复刻 Col/Row/SP 语义;实验后端experiments/transformers_modeling_backend/parallelize.py仍用 PyTorchparallelize_module。examples/torch/llama3/parallelize.py与 Titanmodels/llama3/parallelize.pyapply_tp 路径高度对应,可作为迁移模板。with loss_parallel(),而用components/loss.py的_LossParallelCrossEntropy+ChunkedLossWrapper;Hyper 用户可在 trainer 中包loss_parallel()或后续补齐对等 wrapper。enable_async_tensor_parallel依赖torch.compile;Hyper 暂无对等开关。ParallelDims.seq_len_divisor = tp * (cp * 2)(SP 开启时);Hypertrainer.ParallelDims移植同一约束。A.8 源码索引(便于 Review)
tensor/parallel/api.pycore/tensor_parallel/api.pymodels/*/parallelize.pytensor/parallel/style.pycore/tensor_parallel/style.pymodels/common/decoder_sharding.pytensor/parallel/loss.pycore/tensor_parallel/loss_parallel.pycomponents/loss.pytensor/parallel/input_reshard.pytensor/parallel/fsdp.pyfully_shard+ DTensor 参数distributed/fsdp.pycore/tensor_parallel/style.pydistributed/tensor_parallel.pymodels/llama3/parallelize.pyexamples/torch/llama3/parallelize.pymodels/llama3/parallelize.py父文档:HyperParallel / PyTorch / TorchTitan 用户接口对标 · 姊妹专项:DTensor / DeviceMesh #2 · HSDP / FSDP #5 · 子专项:loss_parallel #269 · TP 异步 ACT #270