已开启
【RFC】HP Tensor Parallel ParallelStyle 设计文档 #335
changzherui创建于 8月17日
changzherui
8月17日 评论:
8月17日 评论:
优先测试pytorch后端


8月17日 修改标题为 “【RFC】HP Tensor Parallel ParallelStyle 设计文档”,原标题为“【RFC】HP Tensor Parallel ParallelStyle 转测设计”
8月17日 修改标题为 “【RFC】HP Tensor Parallel ParallelStyle 设计文档”,原标题为“【RFC】HP Tensor Parallel ParallelStyle 转测设计”
HP Tensor Parallel ParallelStyle 设计文档(转测)
0. 基本信息
ParallelStyle(style.py全量)hyper_parallel/core/tensor_parallel/style.pyhyper_parallel/core/tensor_parallel/api.py的parallelize_modulemasterColwise+Rowwise通信调试 STtorch.distributed.tensor.parallel.style/parallelize_modulehyper_parallel/models/qwen3_5/parallelize.pyParallelStyle全量转测本期转测对象是
style.py导出的全部 TP Style,以及把它们应用到模块树上的parallelize_module:1. 背景
Tensor Parallel(TP)把 Linear / Embedding 的权重沿特征维切开,再靠激活上的集合通信把分片结果拼回。HyperParallel 不要求用户手写 all-gather / all-reduce,而是提供一组声明式 Style:
这套 API 对齐 PyTorch DTensor TP:
distribute_tensor;src_data_rank注入由parallelize_module完成。业务侧已经按这个模型拼出完整 Transformer TP plan。以 Qwen3.5 为例:
embed_tokens:RowwiseParallel(Replicate → Shard(1))input_layernorm/post_attention_layernorm/model.norm:SequenceParallel(sequence_dim=1)self_attn/mlp入口:PrepareModuleInput(Shard(1) → Replicate)q/k/v_proj、gate/up_proj:ColwiseParallel()q_norm/k_norm:SequenceParallel(sequence_dim=2, use_local_output=True)o_proj/down_proj:RowwiseParallel(Shard(-1) → Shard(1), use_local_output=False),部分层带reduce_dtype=float32lm_head:ColwiseParallel(Shard(1) → Replicate 或 Shard(-1))因此转测不能只测单个 Linear 切分,还必须覆盖 Style 组合、I/O 布局转换、以及 1-D TP mesh 约束。
相关资料:
torch.distributed.tensor.parallel.style2. 本期目标与非目标
2.1 本期目标
ParallelStyle及上述 7 个具体 Style,均可从hyper_parallel顶层导入。ColwiseParallel/RowwiseParallel对 Linear、Embedding 做列切 / 行切,并挂上输入输出 hook。SequenceParallel复制参数,按序列维切激活,供 LayerNorm / Dropout / RMSNorm 使用。PrepareModuleInput/PrepareModuleOutput/PrepareModuleInputOutput只做边界布局转换,不切参数。NoParallel把参数做成全复制 DTensor,保持 DTensor 语义但不切分,供 Router / 关闭 SP 的 Norm 使用。parallelize_module支持:dict[str, ParallelStyle],key 为 FQN,支持fnmatchglob;DeviceMesh;src_data_rank写入 Style 后再apply。platform.is_linear_module/platform.is_embedding_module。2.2 本期非目标
以下内容不在本次
style.py转测范围内,但测试时需要知道它们会通过同一个parallelize_module/ParallelStyle.apply入口出现:MC2ColwiseParallel/MC2RowwiseParallelmc2_style.pyMC2Linear,并要求输入/输出为序列切分ContextParallel/AsyncContextParallel/ DSA CP /LinearAttentionContextParallelcore/context_parallel/ParallelStyle,走 CP 语义,不是 TP 权重切分TensorParallel(EP 专家权重)core/expert_parallel/expert_parallel.pyw1/w2/w3切分,不是nn.LinearCol/Rowloss_parallel.pylm_head输出是否保持Shard(-1),不是独立 StyleShard/Replicate/Partialparallelize_plan=None只 warning,不做 auto-parallel也不承诺:
parallelize_module(必须先切 1-D 子 mesh,例如mesh["tp"]);tp_size整除时的自动 padding;3. ParallelStyle 语义
3.1 基类契约
约束:
ParallelStyle(),必须实现apply。apply原地改模块(改参数、挂 hook),也可以返回包装后的模块。src_data_rank由parallelize_module(..., src_data_rank=...)写入。0表示从 rank 0 scatter/broadcast 全局参数;None表示各 rank 本地切片、不通信。_src_data_rank_for_tensor()发现tensor.is_meta后强制把src_data_rank置为None。3.2 切分维约定
PyTorch
nn.Linear.weight/ MindSporenn.Dense.weight的逻辑 shape 都按[out_features, in_features]理解:ColwiseParallelShard(0),切 out_featuresShard(0)Replicate()Shard(-1)ColwiseParallelShard(1),切 embedding_dimReplicate()Shard(-1)RowwiseParallelShard(1),切 in_featuresReplicate()Shard(-1)Replicate()RowwiseParallelShard(0),切 vocabReplicate()(apply 时改desired_input_layouts)Replicate()SequenceParallelReplicate()Replicate()Shard(sequence_dim),默认 dim=1use_local_output=FalseNoParallelReplicate()Replicate()Replicate()Replicate()PrepareModule*可整除约束(均匀
Shard,与 DTensor 一致):out_features % tp_size == 0in_features % tp_size == 0embedding_dim % tp_size == 0num_embeddings % tp_size == 0% tp_size == 0不满足时由 DTensor/通信层报错,Style 层不额外 padding。
3.3 激活上的通信
列切 + 行切的标准 MLP:
Qwen3.5 打开 Sequence Parallel 后,行切输出不再还原成
Replicate,而是 reduce-scatter 到Shard(1),让后续 Norm 继续吃序列分片。3.4 真实示例:2 卡 Colwise Linear
(32, 32),对应 out[0, 32)(32,)(32, 32),对应 out[32, 64)(32,)要和单卡参考对比,需要沿最后一维
all_gather再cat。Rowwise Linear(
in=32, out=24):(24, 16),对应 in[0, 16)(24,),Replicate(24, 16),对应 in[16, 32)(24,)反向时,Rowwise 权重梯度在 gather 后需要除以
tp_size再和单卡参考比(现有 ST 已按这个口径断言)。4. 总体设计
4.1 架构与数据流
设计原则:
distribute_tensor和DTensor.redistribute。distribute_module第二次调用同一模块会RuntimeError。partition_fn里切掉的参数,会被distribute_module自动做成Replicate()DTensor。这就是SequenceParallel/NoParallel的参数语义来源。4.2
parallelize_module代码位置:
hyper_parallel/core/tensor_parallel/api.py。行为:
device_mesh=Nonewith mesh:或内部_tensor_parallel_mesh_contextdevice_mesh.ndim > 1ValueError:Tensor Parallel only accepts a 1D DeviceMeshparallelize_plan=NoneParallelStyleplan.src_data_rank = src_data_rank,然后plan.apply(module, mesh)。会原地改传入的 Style 对象dictParallelStyle;key 按.拆成 FQN atom,atom 用fnmatch匹配named_children/ MindSporename_cellsa..bValueErrorTypeError混合并行正确用法:
4.3
distribute_module与参数切分代码位置:
hyper_parallel/core/dtensor/dtensor.py。ColwiseParallel/RowwiseParallel/SequenceParallel/NoParallel都走它:named_modules调用partition_fn(可空)。Replicate()。input_fn/output_fn。_distribute_module_applied=True,禁止二次调用。ColwiseParallel._partition_linear_fn:所有 param(weight、bias)distribute_tensor(..., [Shard(0)], src_data_rank=...)。RowwiseParallel._partition_linear_fn:weight → Shard(1),其它 param(bias)→ Replicate()。Embedding 对应
Shard(1)(列切)或Shard(0)(行切)。SequenceParallel的partition_fn是空操作,随后 replicate 通道把 Norm 权重做成复制 DTensor。NoParallel直接partition_fn=None,全部 replicate。4.4 I/O hook
Colwise / Rowwise / SequenceParallel / NoParallel
通过
distribute_module的input_fn/output_fn注册。MindSpore pre-hook 必须返回 tuple,因此这些 Style 的 input hook 都返回(prepared_first_input, ...)。Colwise / Rowwise 只处理第一个位置参数:
Rowwise Embedding 比较特殊:
nn.Embedding.forward即使 weight 已切,返回的仍是普通 tensor。output hook 会把它标成Partial("sum"),再 redistribute 到output_layouts。非 Embedding 的普通 tensor 输出直接TypeError。Rowwise / PrepareModuleOutput 的
reduce_dtype:当输出仍是 Partial 且目标 layout 不同时,先把 local tensor cast 到reduce_dtype,再 redistribute。Qwen 的o_proj用reduce_dtype=torch.float32做高精度 reduce。PrepareModuleInput
不走
distribute_module,直接platform.register_forward_pre_hook。input_layouts/desired_input_layouts可以是单个Placement或 tuple。None表示该位置参数原样透传。input_kwarg_layouts时走with_kwargs=True的 pre-hook,同时处理 args 和 kwargs。use_local_output=True表示准备完后to_local(),模块 forward 看到的是普通 tensor。这个名字对齐 PyTorch,容易误解:它改的是输入,不是模块输出。构造期校验:
input_layouts就必须给desired_input_layouts;运行期:forward 实参个数必须等于
input_layouts长度,否则ValueError。非 tensor 且带 layout 时AssertionError。PrepareModuleOutput
直接
module.register_forward_hook。单输出返回单个 tensor;多输出返回 tuple。Noneslot 透传。reduce_dtype仅在 Partial 输出且需要 redistribute 时生效。PrepareModuleInputOutput
内部组合上面两个 Style。
use_local_input映射到PrepareModuleInput(..., use_local_output=use_local_input)。4.5 与其它模块的交互
distribute_tensor/redistributeparallelize_modulestyle.apply也可以,但不会自动写src_data_rank(除非调用方自己设)fully_shardmesh["tp"]parallelize_moduleShard(seq),下一层PrepareModuleInput再 gatherlm_head用ColwiseParallel(output_layouts=Shard(-1), use_local_output=False)ParallelStyle,可出现在同一个parallelize_moduleplan 里TensorParallelColwiseParallel混用在同一个nn.Linear上5. 对外接口
全部可从
hyper_parallel导入。5.1 应用入口
5.2 构造参数
ColwiseParallel
input_layoutsReplicate()output_layoutsShard(-1)use_local_outputTrueto_local()desired_input_layouts固定为(Replicate(),),用户不能改。如果输入已经是序列切分,Colwise 会 all-gather 成复制再算。RowwiseParallel
input_layoutsShard(-1)output_layoutsReplicate()Shard(seq_dim)即 reduce-scatterreduce_dtypeNoneuse_local_outputTrueto_local()Linear 的
desired_input_layouts默认(Shard(-1),);Embedding 在apply时改成(Replicate(),)。SequenceParallel
sequence_dim1(B, S, H)的S;Qwenq_norm/k_norm用2use_local_outputFalsePrepareModuleInput
input_layoutsNoneNoneslot 透传desired_input_layoutsNoneinput_kwarg_layoutsNonedesired_input_kwarg_layoutsNoneuse_local_outputFalsePrepareModuleOutput
output_layoutsdesired_output_layoutsreduce_dtypeNoneuse_local_outputTrueto_local()PrepareModuleInputOutput:上述输入侧参数 +
use_local_input(默认False)+ 输出侧参数。NoParallel
input_layoutReplicate()desired_input_layoutReplicate()output_layoutReplicate()use_local_outputTrueto_local()注意:NoParallel 用的是单数
input_layout/output_layout,不是 Col/Row 的复数*_layouts。6. 当前支持矩阵
ParallelStyle抽象契约、src_data_rankColwiseParallelLinear 前向/反向ColwiseParallelEmbedding 前向RowwiseParallelLinear 前向/反向RowwiseParallelEmbeddingCommDebugMode通信计数 STSequenceParallelLayerNorm / DropoutSequenceParallel(sequence_dim=2)q_norm/k_norm使用;无独立 STPrepareModuleInput位置参数 / kwargs /NoneslotPrepareModuleOutput单输出 / 多输出Noneslotmodule.register_forward_hook,需确认 Cell 路径PrepareModuleInputOutput链路NoParallel复制 Linear + SP→NoParallel redistributereduce_dtype(Rowwise / PrepareModuleOutput)src_data_rank=0/Noneparallelize_module有功能 ST;meta tensor 强制Nonefnmatchglob planValueErrorNotImplementedError“接口支持”表示代码走
platform抽象,不区分 PT/MS 分支;是否在真实 MS 多卡上数值正确,需要转测补齐。7. 风险与限制
7.1 只切第一个位置参数
Colwise / Rowwise / SequenceParallel / NoParallel 的 input hook 只处理
inputs[0]。多输入模块(例如带mask的 Attention 根模块)应使用PrepareModuleInput,不要假设 Col/Row 会转换全部参数。7.2
use_local_output语义不统一TrueFalseFalseQwen 里大量
use_local_output=False,就是为了让激活以 DTensor 形式在 Style 之间传递。测试时必须区分「返回 local」和「返回 DTensor」。7.3 Rowwise Embedding 的 Partial 包装
Embedding 前向不会自动产出 DTensor 输出。Rowwise 依赖 output hook 把普通 tensor 标成
Partial("sum")。如果模块类型检测失败(platform.is_embedding_module为 false),会变成TypeError而不是静默错误结果。7.4
distribute_module只能调用一次同一模块不能先
ColwiseParallel.apply再PrepareModuleOutput.apply走distribute_module。PrepareModule*不走distribute_module,所以可以挂在已经被 Col/Row 切过的模块外面;但两个都会调用distribute_module的 Style 不能叠在同一个模块上。7.5 MindSpore hook 差异
PrepareModuleInput使用platform.register_forward_pre_hook,有with_kwargs。PrepareModuleOutput直接调用module.register_forward_hook。转测 MS 时要把 kwargs pre-hook、多输出 forward hook 作为必测项,不能只测 PT。
7.6 梯度比较口径
tp_size再比单卡(现有 NPU ST 口径)。7.7 混合并行顺序
推荐:先
parallelize_module(..., mesh["tp"]),再fully_shard(..., mesh["dp"])。现有 4 卡 STtest_tp_fsdp_mlp_fwd_bwd_precision_npu覆盖 MLP 这一路径。PP / CP 与 TP 的组合不在本次 Style 单测范围,但 plan 里可以同时出现 CP Style。8. 验证设计与当前结果
8.1 已有覆盖(开发侧)
UT(
tests/ut/core/tensor_parallel/,CPU mock,不启分布式):test_style.pysrc_data_rank、PrepareModule* 构造校验与 identity 前向test_colwise_parallel.pytest_rowwise_parallel.pytest_sequence_parallel.pydistribute_module回调、输入类型校验test_no_parallel.pytest_api.pyparallelize_moduleglob、2-D mesh 拒绝、空 plan warning、单 Style 根应用PyTorch ST(
tests/torch/tensor_parallel/):test_tp_styles_distributed.pytest_tp_sequence_parallel_distributed.pytest_prepare_module_io_distributed.pyNoneslot、与 Col/Row 组合、4 卡 MLP blocktest_parallelize_module_distributed.pysrc_data_rank、单 Style 根test_tp_hybrid_distributed.py精度口径:NPU float32 vs CPU 单卡参考,
rtol=1.5e-4,atol=1e-5。MindSpore ST:
tests/mindspore/st/dtensor/_test_comm_debug_mode_mlp.py:2 卡Colwise+RowwiseMLP,只断言 collective 次数,不比数值。8.2 转测建议用例
下列用例按测试可直接拆 Feature / Description / Expectation。未标注“已有”的视为转测补齐重点。
A. 接口与 fail-closed
ParallelStyle()TypeError,信息含abstractparallelize_moduleValueError,提示使用device_mesh["tp"]TypeError""或"a..b"ValueErrorLayerNorm/ReLU/ 自定义 CellNotImplementedError,信息含Linear and Embeddinginput_layouts与desired_*长度不同AssertionErrorValueError,含same lengthValueError,含tensor or DTensorTypeErrordistribute_moduleapplyRuntimeErrorparallelize_plan=NoneA6、A2、A11 必须在 PT 和 MS 上都测。fail-closed 的验收表现就是立刻抛上述异常,不能静默按未切分模块继续算。
B. 参数切分几何
weight.placements == (Shard(0),),local shape[out/tp, in];bias 同样Shard(0)Shard(1),local[out, in/tp];biasReplicate(),local 完整outShard(1),local[vocab, embed/tp]Shard(0),local[vocab/tp, embed]Replicate()DTensor,local 完整 hiddenReplicate()src_data_rank=0src_data_rank=Noneis_meta权重切分不发起源 rank 通信,不 hangB4、B7、B8、B9 是当前缺口。
C. 数值精度(多卡 vs 单卡参考)
参考实现:同一组 float32 权重在 CPU 单进程上跑
F.linear/F.embedding/LayerNorm。/ tp_sizelinear1 Colwise + relu + linear2 RowwiseShard(0)→Replicate再列切RowwiseParallel(reduce_dtype=fp32)SequenceParallel(sequence_dim=2)(B, H, S)或 Qwen RMSNorm 形状C1–C4、C6–C12、C15 在 PT 上已有;转测重点是 MS 复现 C1–C6、C8、C13、C14,以及 PT 补 C5/C13/C14。
建议 dtype:主路径 float32;C13 覆盖 fp16/bf16 +
reduce_dtype=float32。D. 组合与业务 plan
"layers.*.mlp.gate_proj": ColwiseParallel()命中所有层parallelize_module(linear, mesh, ColwiseParallel())PrepareModuleInput(Shard(1)→Replicate)+ Colwise q/k/v + Rowwise o_proj(output=Shard(1),use_local_output=False) + SequenceParallel normlm_head+ Loss Parallel 布局ColwiseParallel(input_layouts=Shard(1), output_layouts=Shard(-1), use_local_output=False),输出保持 DTensor 且最后一维切分NoParallel()替代SequenceParallel(),输入 replicateD3 是最接近
parallelize_qwen3_5_tp的最小可测单元,建议作为 MS/PT 共同验收。E. 平台矩阵
8.3 如何验证(对应测试常问问题)
module.weight,确认是DTensor,看placements和to_local().shape。allclose。Replicate输出必须有 all-reduce / reduce-scatter;可用CommDebugMode或 profiler 看 collective 次数(MS 已有 MLP 通信 ST)。9. 验收标准
9.1 功能验收
hyper_parallel导入,且都是ParallelStyle子类。NotImplementedError。sequence_dim上切激活,参数保持复制;LayerNorm 前向/反向与单卡切片一致。Noneslot 不改对应输入/输出;kwargs 路径可用。parallelize_module支持单 Style、dict、fnmatch,拒绝 N-D mesh。src_data_rank=0/None都能得到正确 local shard;meta 参数不通信。9.2 兼容性验收
nn.Linear/nn.Dense行为不变。platform(模块类型、hook 注册)。9.3 明确报错(必须 fail-closed)
ParallelStyleTypeErrorNotImplementedErrorparallelize_module收到 N-D meshValueErrordict[str, Style]TypeErrorValueErrorAssertionErrorValueErrorValueErrorTypeErrordistribute_moduleRuntimeErrorAssertionError9.4 精度口径
rtol=1.5e-4,atol=1e-5(与现有 ST 一致)reduce_dtype=fp32rtol=1e-3,atol=1e-3,但必须优于不升精度 reduce 的误差tp_size9.5 转测完成定义
reduce_dtype、sequence_dim=2。TensorParallel算作本次style.py转测通过条件。