已开启
[接口梳理3] 张量并行 TP 全量接口对标总表(Hyper · PyTorch · TorchTitan) #9
changzherui创建于  7月1日
changzherui
changzherui成员
7月1日 创建

一、文档说明

📎 上级文档: HyperParallel / PyTorch / TorchTitan 用户接口对标 #1

📎 姊妹专项(差距与路线图): 张量并行 TP 对标 PyTorch #6

本 Issue 定位: 在 #1 模块划分下,汇总 §3 张量并行(TP) 子模块的全量对外接口对标总表(表头:接口名 · 参数 · 功能 · hyper · pytorch · titan),作为 #6 差距梳理的接口索引附录。

分析依据(源码,2026-07):

侧 路径
HyperParallel hyper_parallel/__init__.py、hyper_parallel/core/tensor_parallel/、platform/*/loss_parallel_ops.py
PyTorch torch/distributed/tensor/parallel/(api.py、style.py、loss.py、input_reshard.py、fsdp.py、ddp.py)
TorchTitan torchtitan/distributed/tensor_parallel.py、models/common/decoder_sharding.py、models/*/parallelize.py、components/loss.py

图例: ✅ 有对标 · ⚠️ 部分对标 · ❌ 无 · 🔶 Hyper 扩展

列说明:

  • hyper:HyperParallel 导出路径或实现位置
  • pytorch:PyTorch torch.distributed.tensor.parallel 对标 API
  • titan:TorchTitan 实际调用的接口(含编排层封装,非 Titan 自有原语)

架构备注: Hyper 策略类使用公开 apply()(PyTorch 为 _apply());I/O 边界 redistribute 当前全同步(PyTorch 默认 async_op=True → ACT),见 #270。


二、全量接口对标总表

接口名 参数 功能 hyper pytorch titan
入口 API
parallelize_module module, device_mesh, parallelize_plan, *, src_data_rank=0 对根模块或子模块应用声明式 TP 策略 hyper_parallel.parallelize_module torch.distributed.tensor.parallel.parallelize_module 实验路径仍直接调用;主流用 Module.parallelize()
parallelize_module · device_mesh DeviceMesh | None 1-D TP mesh;None 取当前 mesh 上下文 ✅;ndim>1 显式抛错 ✅;文档要求 1-D parallel_dims.get_mesh("tp")
parallelize_module · parallelize_plan ParallelStyle | dict[str, ParallelStyle] | None 单策略或路径→策略字典 ✅ ✅ ShardingConfig 映射为等价 layout
parallelize_module · fnmatch glob layers.* 等 子模块路径 glob 匹配 ✅ ✅ 声明式 sharding 内部
parallelize_module · src_data_rank int | None,默认 0 写入 style 后传入 distribute_tensor ✅ ✅ 同左
parallelize_module · plan=None — warn + no-op ✅ ✅ —
parallelize_module · plan="" 空路径 — 对当前模块应用策略 ❌ 空路径抛 ValueError ✅ shortcut —
_tensor_parallel_mesh_context device_mesh 测试/库内部 mesh 上下文 api._tensor_parallel_mesh_context(未导出) — —
并行策略基类
ParallelStyle ABC TP 策略抽象基类 hyper_parallel.ParallelStyle tensor.parallel.ParallelStyle 继承 PyTorch 基类(Titan NoParallel)
ParallelStyle.apply module, device_mesh 应用策略(in-place) ✅ 公开 apply _apply(私有) _apply
ParallelStyle.src_data_rank 类属性 parallelize_module 写入 ✅ ✅ ✅
ColwiseParallel
ColwiseParallel *, input_layouts, output_layouts, use_local_output 列切分 Linear/Embedding hyper_parallel.ColwiseParallel ✅ colwise_config() / 等价
Linear weight — Shard(0) ✅ ✅ ✅
Embedding weight — Shard(1) ✅ ✅ ✅
默认 input / output — Replicate() → Shard(-1) ✅ ✅ ✅
use_local_output 默认 True post-hook to_local() ✅ ✅ ✅
支持模块 — Linear、Embedding ✅ + MS Dense/Embedding Linear、Embedding 同 PyTorch
I/O redistribute — 边界 layout 转换 ⚠️ 同步 async_op=True 经 PyTorch style
RowwiseParallel
RowwiseParallel *, input_layouts, output_layouts, use_local_output=True 行切分 Linear/Embedding ✅ ✅ rowwise_config()
Linear weight / bias — Shard(1) / Replicate() ✅ ✅ ✅
Embedding weight — Shard(0) ✅ ✅ ✅
Embedding 输出 Partial — Partial("sum") 再 redistribute ✅ ✅ ✅
Embedding 默认 input — Replicate()(非 Shard(-1)) ✅ ✅ ✅
SequenceParallel
SequenceParallel *, sequence_dim=1, use_local_output=False 参数 replicate、序列维 activation shard ✅ ✅ dense_sequence_parallel_placement()
参数 partition_fn — 无切分(replicate) ✅ ✅ ✅
PrepareModuleInput
PrepareModuleInput input_layouts, desired_input_layouts, input_kwarg_layouts, desired_input_kwarg_layouts, use_local_output=False 仅 hook:输入 annotate + redistribute ✅ ✅ Attention/FFN 入口
kwarg layouts — with_kwargs=True pre-hook ✅ ✅ 声明式路径
None layout 槽 — 跳过该输入 ✅ ✅ ✅
PrepareModuleOutput
PrepareModuleOutput output_layouts, desired_output_layouts, use_local_output=True 仅 hook:输出 redistribute ✅ ✅ ✅
PrepareModuleInputOutput
PrepareModuleInputOutput 组合上述 + use_local_input 输入输出边界一体 ✅ ✅ ✅
NoParallel 🔶
NoParallel *, input_layout, output_layout, desired_input_layout, use_local_output=True 参数 replicate + DTensor 边界 ✅ hyper_parallel.NoParallel ❌ 官方包未导出 ✅ torchtitan/distributed/tensor_parallel.py
典型用途 — MoE Router、未开 SP 的 Norm ✅ — ✅
Loss Parallel
loss_parallel 无参(PT)/ *, mesh=None, strict=True(Hyper) 上下文内启用分布式 CE core.tensor_parallel.loss_parallel(未进顶层 __all__) tensor.parallel.loss_parallel 生产多用 _LossParallelCrossEntropy
is_loss_parallel_active — 是否在 loss_parallel 上下文 ✅ 🔶 ❌ —
get_loss_parallel_count — 嵌套深度(内部/调试) loss_parallel.py(未导出) ❌ —
拦截 cross_entropy / CrossEntropyLoss — vocab-shard 融合 softmax ⚠️ _OP_DISPATCHER + CE registry DTensor.__torch_dispatch__ handler _LossParallelCrossEntropy
多维 mesh CE — batch/CP 非 TP 维 reduction ⚠️ ✅ ✅
label_smoothing — CE 参数 ❌ ❌(文档标注不支持) ⚠️
PyTorch 组合层 API(TP 模块相关)
input_reshard module, tp_device_mesh, input_reshard_dim=None AC 下 backward 恢复 replicate 输入 ❌ tensor.parallel.input_reshard 声明式路径内处理
DTensorExtensions FSDP 扩展类 TP+FSDP 参数 flatten/chunk、DCP 元数据 ❌ tensor.parallel.fsdp.DTensorExtensions 经 PyTorch FSDP
_pre_dp_module_transform module DDP+TP:DTensor 参数 localize ❌ tensor.parallel.ddp(内部) Hyper 主路径 FSDP
_data_parallel_utils sync_grad_hook 等 DP 与 DTensor 梯度同步辅助 ❌ 内部 —
TorchTitan 编排(仅 Titan)
maybe_enable_async_tp parallelism, compile_config, tp_mesh 开启 Inductor micro-pipeline TP ❌ — ✅ distributed/tensor_parallel.py
parallelize_llama / model.parallelize parallel_dims, … CP→TP→AC→compile→FSDP 流水线 ⚠️ examples/torch/llama3/parallelize.py — ✅ models/llama3/parallelize.py
parallelize_qwen3_5 model, mesh, cfg 全模型 TP+AC+FSDP ❌ TP NotImplementedError — Titan 有 qwen 路径
decoder_sharding colwise_config, rowwise_config, set_gqa_attention_sharding, … 声明式 ShardingConfig / SpmdLayout ❌ — ✅ models/common/decoder_sharding.py
Module.parallelize(parallel_dims) — 新一代 SPMD 声明式 TP ❌ — ✅ spmd_backend=full_dtensor
ShardingConfig / SpmdLayout — 模块级声明式分片协议 ❌ — ✅ Titan 主路径
_LossParallelCrossEntropy — 无上下文分布式 CE Function ⚠️ 可用 loss_parallel() — ✅ components/loss.py
ChunkedLossWrapper — 分块 loss + 梯度累积 ❌ — ✅
CrossEntropyLoss(Titan 组件) — Trainer 级 loss 封装 ❌ — ✅
Titan NoParallel input_layout, output_layout, use_local_output Titan 版 replicate 策略 Hyper 有更完整 desired_input_layout ❌ ✅
Hyper 底层扩展(与 TP 强相关)
shard_module module, device_mesh, sharding_plan 多维 mesh 别名低层分片 hyper_parallel.shard_module 🔶 ❌ TP 包无 —
ShardingPlan — 参数/算子分片计划 core/shard/sharding_plan.py — Titan 另有 protocols/sharding.py
DFunction — 用户自定义分布式 autograd hyper_parallel.DFunction 🔶 — —
parallelize_value_and_grad — 值+梯度并行(非 declarative TP) hyper_parallel.parallelize_value_and_grad 🔶 ❌ —
gather_tensor_parallel_logits logits, config 推理阶段 vocab gather hyper_parallel.infer 🔶 — loss 组件内处理
Trainer / Mesh 编排(Hyper)
ParallelDims dp_replicate, dp_shard, cp, tp, pp, ep, … 并行度校验 + mesh 构建 hyper_parallel.trainer.parallel_dims 🔶 — ✅ torchtitan.distributed.ParallelDims
ParallelDims.get_mesh "tp" 等 取命名子 mesh ✅ — ✅
ParallelDims.seq_len_divisor — SP 开启时 tp * (cp*2) ✅ 移植自 Titan — ✅

三、差距速查(与 #6 联动)

本节从 总表 中筛出 hyper 列为 ❌ 或 ⚠️ 的项,按来源拆为四块。详细 P0/P1 补齐路线见 #6;loss_parallel 子专项见 #269;TP 异步 ACT 见 #270。

3.1 仅 PyTorch 有 · Hyper 无

接口 / 参数 PyTorch 路径 功能 影响
input_reshard tensor.parallel.input_reshard AC + 非 SP 模块 backward 输入恢复 replicate AC+TP 显存与正确性
DTensorExtensions tensor.parallel.fsdp FSDP+TP 参数 flatten/chunk、checkpoint 元数据 TP 后 fully_shard 互操作
_pre_dp_module_transform tensor.parallel.ddp DDP 自动 localize DTensor 参数 DDP+TP(Hyper 主路径为 FSDP)
parallelize_plan="" 空路径 shortcut parallel.api 对当前模块直接应用 style 编排写法差异
PyTorch 官方 loss_parallel 顶层导出 tensor.parallel 与 TP 并列导出 Hyper 仅在子包
ParallelStyle._apply 命名 私有 _apply — 仅命名;Hyper 用 apply

3.2 仅 TorchTitan 有 · Hyper 无

Titan 底层 TP 原语多数直接调用 PyTorch parallelize_module / ParallelStyle。下表为 Titan 训练编排层 特有、Hyper 无对等公开 API 的项。

接口 / 能力 Titan 路径 功能 影响
maybe_enable_async_tp distributed/tensor_parallel.py torch.compile + micro-pipeline TP compile 场景 TP 性能
Module.parallelize(parallel_dims) 各 model 声明式 ShardingConfig 一站式 TP 主流模型迁移路径
decoder_sharding 全套 models/common/decoder_sharding.py Col/Row/SP/Attention/FFN 声明式 layout 复杂 attention 变体
ShardingConfig / SpmdLayout protocols/ 非 parallelize_module 的 SPMD 协议 Titan full_dtensor 后端
_LossParallelCrossEntropy + ChunkedLossWrapper components/loss.py 无上下文 CE + 分块 loss Trainer 集成
parallelize_llama 全流水线 models/llama3/parallelize.py CP→TP→AC→compile→FSDP Hyper 有 examples 级 parallelize_llama3
parallelize_qwen3_5 TP Titan qwen 路径 含线性注意力层 TP Hyper NotImplementedError
Titan NoParallel(distributed/tensor_parallel.py) 独立实现 MoE Router replicate Hyper 有对等且更完整

3.3 PyTorch 有 · Hyper 有但语义偏弱(⚠️)

接口 / 参数 PyTorch Hyper 现状 差距说明
Col/Row/Seq/Prepare I/O redistribute async_op=True → ACT 全同步 通信-计算无重叠;→ #270
src_data_rank + distribute_tensor rank0 scatter/broadcast ⚠️ 仅本地 slice(#2 P0) 全局权重一次分片路径不等价
loss_parallel() 无参;多维 mesh 支持 mesh/strict 扩展;未顶层导出 → #269
多维 mesh CE reduction 完整 偏 1D TP mesh batch/CP 维 reduction → #269
label_smoothing 文档标注不支持 ❌ 未实现 训练配置迁移
TP + fully_shard DTensorExtensions 桥接 ⚠️ 有 e2e 测试、无 Extensions checkpoint/load 互操作
parallelize_module dict 空 token PyTorch warn + skip Hyper 部分路径 ValueError 空路径 "" 行为不一致
CE 拦截机制 __torch_dispatch__ custom handler __torch_function__ + _OP_DISPATCHER 机制不同;语义对齐待验

3.4 Hyper 有 · PyTorch 用户 API 无直接对标

接口 说明
NoParallel(官方 TP 包) 参数 replicate + 可配 I/O layout;PyTorch 官方未导出;Titan 有独立实现
loss_parallel(mesh=..., strict=...) 显式 mesh、非严格模式
is_loss_parallel_active() 调试/断言 API
ParallelStyle.apply(公开) PyTorch 为 _apply
shard_module + ShardingPlan 多维 mesh 别名低层分片
DFunction 自定义分布式 autograd,衔接 _OP_DISPATCHER
parallelize_value_and_grad 值+梯度并行独立 API
gather_tensor_parallel_logits 推理 vocab gather(infer/)
ParallelDims(trainer) 移植自 Titan 的并行度校验与 mesh 构建
MindSpore 双栈 同一套 ParallelStyle 适配 Cell/Dense;pre-hook 须返回 tuple
examples/torch/llama3/parallelize.py 对标 Titan apply_tp 的参考编排

四、顶层导出清单(hyper_parallel.__all__,本模块相关)

"parallelize_module",
"ColwiseParallel", "RowwiseParallel", "SequenceParallel",
"PrepareModuleInput", "PrepareModuleInputOutput", "PrepareModuleOutput",
"NoParallel", "ParallelStyle",

子包 hyper_parallel.core.tensor_parallel 额外导出:

"loss_parallel", "is_loss_parallel_active"

未导出但常用:

  • loss_parallel / is_loss_parallel_active(需 from hyper_parallel.core.tensor_parallel import ...)
  • _tensor_parallel_mesh_context(测试内部)
  • shard_module、DFunction、parallelize_value_and_grad(低层扩展,在顶层 __all__)
  • gather_tensor_parallel_logits(hyper_parallel.infer)
  • ParallelDims(hyper_parallel.trainer.parallel_dims)

PyTorch 对标导出(torch.distributed.tensor.parallel):

"parallelize_module", "loss_parallel",
"ColwiseParallel", "RowwiseParallel", "SequenceParallel",
"PrepareModuleInput", "PrepareModuleInputOutput", "PrepareModuleOutput",
"ParallelStyle"
# 注意:无 NoParallel;input_reshard 在子模块单独导入

五、维护说明

  • 本表为 §3 张量并行(TP) 的接口索引;差距分析与 P0/P1 补齐路线见 #6。
  • loss_parallel 细节与 CE 算子 parity 见 #269;TP 边界异步 redistribute 见 #270。
  • DTensor distribute_tensor 通信(影响 TP 参数初始化)归属 #2 / #7,本表仅标注依赖关系。
  • FSDP+TP 组合 与 fully_shard 衔接见 #8。
  • 源码变更导致接口漂移时,请同步更新本 Issue 总表。

父文档:#1 · 姊妹专项:#6 TP 差距与路线图 · 关联:#7 DTensor 总表 · #8 FSDP 总表 · #269 loss_parallel · #270 TP 异步 ACT

likedislike