已开启
HyperParallel:`ExpertTensorParallel`(EP⊗TP)设计文档 #104
changzherui创建于 4月27日
4月27日 将 changzherui1 设为负责人
5月29日 关联了pull request:feat(examples): add MoE Expert Parallel correctness demo
6月2日 关联了pull request:feat(moe): add PP+EP 1F1B correctness verification example
本文档说明 Expert Tensor Parallel(ETP) 在 HyperParallel 中的设计定位,并与 TorchTitan
distributed/expert_parallel.py中的ExpertTensorParallel对齐说明。不要求重复粘贴实现代码,侧重背景、职责边界、接口与 mesh 契约、测试与验收。1. 背景
1.1 问题域:MoE 的两类并行
w1/w3列切、w2行切),通信多为 TP 组内 all-gather / reduce-scatter。仅 EP 时,单专家算子可能过大;仅 TP 时,专家总数受单卡显存限制。ETP 同时在 专家维 与 专家内部维 切分,需要 二维 DeviceMesh,且必须澄清:token 路由通信只随 EP 维度走,不参与 TP 维度的无关 collective。
1.2 TorchTitan 参考语义
TorchTitan 中结构如下(逻辑分层):
ExpertParallel:GroupedExperts权重Shard(0)(专家维);_token_dispatch/_token_combine在传入的 1D EP mesh 上做计数交换、可微 all-to-all、permute/unpermute。ExpertTensorParallel(ExpertParallel):_partition_fn:对w1/w2/w3使用[Shard(0), Shard(1)]或[Shard(0), Shard(2)](与 SwiGLU 形状约定一致:EP 切专家维,TP 切 hidden / 行维)。_token_dispatch/_token_combine:传入device_mesh["ep"],即 仅在 EP 子 mesh 上调用父类逻辑,避免用 整幅[ep, tp]mesh 做 token collective。TorchTitan 注释要点:「for dp2ep with TP」——在 DP 转 EP 且仍保留 TP 的场景下,专家权重二维切分、token 仍按 EP 组交换。
1.3 HyperParallel 中的位置
ExpertParallel:与 TorchTitan 同类 all-to-all EP 语义(实现细节可依赖平台 collective)。TensorParallel(MoE 专用):仅 TP、Shard(1)/Shard(2),无 dispatch(EP 度为 1 时切专家算子)。ExpertTensorParallel:继承ExpertParallel,覆写_partition_fn为二维切分,覆写_token_dispatch/_token_combine将device_mesh["ep"]传给父类。2. 解决的问题与不当用法
ExpertParallel+ 1Depmesh,勿用 ETPTensorParallel+ 1Dtpmesh3. 目标与非目标
3.1 目标
[num_experts, …]形状约定下w1/w3:[Shard(0), Shard(1)],w2:[Shard(0), Shard(2)](维度索引与 PyTorch / HyperParallelGroupedExperts定义一致)。ExpertParallel完全同款算法,仅 collective group =mesh["ep"].get_group()(或平台等价 API)。ParallelStyle.apply(module, device_mesh),与ColwiseParallel等一致;device_mesh为 2D,且含命名维"ep"与"tp"(或与实现约定的子 mesh 名字一致)。3.2 非目标
NoParallel、TP router):不在本文档范围;由上层 MoE parallelize 组合。DeepEPExpertParallel;HyperParallel 若引入,应 独立样式,不混在本 ETP 基线语义中。4. 接口与 Mesh 契约
4.1 类与继承
4.2
apply(module, device_mesh)moduleGroupedExperts(或与其实现 相同参数名与形状约定 的 MoE 专家子模块)。device_meshdevice_mesh["ep"]、device_mesh["tp"]切片为 1D 子 mesh;ndim == 2,维名("ep", "tp")(见init_device_mesh约定)。禁止:对 1D mesh 调用 ETP(应在
parallelize_module或apply内 显式校验 并给出可读错误)。4.3 权重分片语义(GroupedExperts)
与
docs/expert_parallel.md表格一致:w1[E, H, D][Shard(0), Shard(1)]w3[E, H, D][Shard(0), Shard(1)]w2[E, D, H][Shard(0), Shard(2)]其中维
0为 专家维(EP),1/2为 TP 切分维(列 / 行与 Megatron SwiGLU 一致)。4.4 Token 路径语义
ExpertParallel相同步骤(计数 all_to_all、可微 token all_to_all、permute / unpermute);唯一区别是device_mesh实参 为device_mesh["ep"]。5. 与 TorchTitan 的差异与对齐点
distribute_module+_applyapply→distribute_module(与现有 EP 一致)device_mesh["ep"]w1/w2/w3二维 Shard[Shard(0), Shard(1)]/[Shard(0), Shard(2)]TensorParallel类TensorParallel(MoE)6. 实现要点检查清单(落地代码时)
parallelize_module/ TP 校验:现有api.py可能 限制仅 1D mesh;对 ETP 需在 MoE 专用路径或apply入口 使用 2D mesh,避免被 TP 1D 校验误伤(若尚未放开,属 已知集成项)。GroupedExperts.forward:权重为 DTensor 时to_local()后再走 grouped matmul(已有平台约定)。distribute_module:避免对同一 module 重复包装(沿用全局_distribute_module_applied约定)。all_to_all_single_autograd的可微行为需与 ExpertParallel 单测策略一致。7. 测试设计
7.1 已有单测映射(维护建议)
仓库中
tests/ut/core/expert_parallel/test_expert_parallel.py已包含 C4:ExpertTensorParallel:_partition_fn:二维 Shard placement 断言;_token_dispatch/_token_combine:对device_mesh["ep"]委托的 spy/mock。设计文档要求:分区与委托行为变更时必须同步更新上述用例。
7.2 建议增补(若尚未覆盖)
ep_tpmesh 上,apply后w1.placements为(Shard(0), Shard(1))(或平台等价表示)。device_mesh["ep"],断言 dispatch/combine 未对tp维调用 collective。7.3 文档
docs/expert_parallel.md:已含 ETP 小节;本设计文档批准后,可补充 「与 TorchTitan 对齐」 一句及 mesh 命名硬性要求。auto_parallel/fast-tuner/docs/torchtitan.md如有 MoE 策略说明,可加交叉引用。8. 风险与开放问题
parallelize_module仅接受 1D TP meshExpertTensorParallel().apply(..., mesh_2d)直连,或扩展 API 显式区分 TP 与 EP⊗TP。"ep"/"tp"ep_mesh_dim/tp_mesh_dim字符串参数(RFC)。ExpertParallel相同:num_experts % ep_degree == 0,且隐藏维可被 tp_degree 整除(具体由GroupedExperts与 TP shard 维推导)。9. 验收标准
10. 参考文献
torchtitan/distributed/expert_parallel.py—ExpertTensorParallel、ExpertParallel、TensorParallel(GroupedExperts)。hyper_parallel/core/expert_parallel/expert_parallel.py、docs/expert_parallel.md。DeviceMesh子 mesh 索引、distribute_tensor多维Shard。