已开启
[基础能力4] loss_parallel / 分布式 CE:当前差异梳理与补齐计划 #269
changzherui创建于 7月1日
7月1日 修改标题为 “[基础能力4] loss_parallel / 分布式 CE:当前差异梳理与补齐计划”,原标题为“[基础能力4] loss_parallel / 分布式 CE 对标 PyTorch & TorchTitan:差距梳理与补齐路线图”
7月1日 修改标题为 “[基础能力4] loss_parallel / 分布式 CE:当前差异梳理与补齐计划”,原标题为“[基础能力4] loss_parallel / 分布式 CE 对标 PyTorch & TorchTitan:差距梳理与补齐路线图”
7月1日 修改了issue 的描述
7月1日 将 changzherui1 设为负责人
7月1日 修改了issue 的描述
一、背景与目标
本 Issue 目标: 在保持 Hyper 现有
loss_parallel架构(上下文管理器 +_OP_DISPATCHER路由 + 双栈 kernel)不变的前提下,根据下文当前差异逐项补齐loss_parallel/ 分布式 Cross-Entropy 能力。背景:
lm_head常用 ColwiseParallel,logits 在 vocab 维为Shard(-1),每个 rank 仅持有V/tp切片。直接F.cross_entropy语义错误,必须在 TP 维做分布式 softmax(all-reduce max / sumexp / gather)。with loss_parallel(): F.cross_entropy(...),经 DTensor custom handler 拦截_log_softmax/nll_loss。components/loss.py的_LossParallelCrossEntropy(plain tensor +tp_group)+ 外层ChunkedLossWrapper(seq 分块降显存)。loss_parallel上下文 +DistributedCrossEntropyFunction(Torch/MindSpore 双栈),核心 CE 算法与 Titan_LossParallelCrossEntropy同类,但在 mesh/labels/reduction 语义、框架集成、ChunkedLoss 上与 PyTorch / Titan 存在下述差异。分析依据(源码):
torch/distributed/tensor/parallel/loss.pyhyper_parallel/core/tensor_parallel/loss_parallel.py、platform/*/loss_parallel_ops.py、_ce_op_registry.pytorchtitan/components/loss.py(CrossEntropyLoss、_LossParallelCrossEntropy、ChunkedLossWrapper)图例: ✅ Hyper 已有 · ⚠️ 部分具备 · ❌ Hyper 缺失 · 🔶 Hyper 扩展
二、范围与约束
2.1 模块范围
loss_parallel()loss_parallel(mesh, strict)DistributedCrossEntropyFunction_LossParallelCrossEntropy_LossParallelCrossEntropy.apply(...)CrossEntropyLoss+ChunkedLossWrapper_OP_DISPATCHER+_ce_op_registry2.2 约束
with loss_parallel()上下文 +_OP_DISPATCHER拦截。platform/*/loss_parallel_ops.py。loss_parallel.py内核。parallelize_module/ lm_head 编排见 #6;TP 边界通信重叠见 #270(与本 Issue 正交,不影响 CE 正确性)。2.3 三端现状对比
2.4 算子拦截机制(与 PyTorch 差异)
DTensorBase.__torch_function__→_OP_DISPATCHERDTensor.__torch_dispatch___ce_op_registry+ 上下文loss_parallel()_log_softmax/nll_lossDistributedCrossEntropyFunction(双栈)说明: Hyper 不在 DTensor 上实现全套
__torch_dispatch__;CE 语义由_OP_DISPATCHER在上下文激活时路由到DistributedCrossEntropyFunction。与 #270 的 ACT 异步 wait 无关——CE 路径走融合 kernel,不依赖 redistribute 重叠。三、当前差异、影响与补齐方向
3.1 Hyper 已具备(本 Issue 不改动)
ceil(V/tp)chunk 语义ignore_index=-100log_softmax/nll_lossloss_parallel(mesh, strict)3.2 差异总表
tensor.parallel公开components/losshyper_parallel.__all__F.cross_entropy__all__+ 文档local_map支持多维mesh_dim=0写死_find_tp_mesh_dim;DTensor labels 校验;非 TP 维Partial输出_LossParallelCrossEntropy.applyloss_parallel_cross_entropy(...)reduction语义none/sum多维 layout 完整;mean仅 1Dsum+/global_valid_tokensmean/sum/none有;loss 非 DTensormean可能不正确;与 Titan token 归一化路径不同CrossEntropyLoss组件CrossEntropyLossChunkedLossWrappernn.CrossEntropyLoss模块cross_entropy_loss_ce_op_registrygather_tensor_parallel_logits3.3 按场景影响
tp_mesh+with loss_parallel()CrossEntropyLoss+ChunkedLossWrapperwith loss_parallel()+ 2D meshreduction="mean"+ 2D mesh四、补齐计划(P0 → P1 → P2)
4.1 优先级总览
4.2 P0
loss_parallelhyper_parallel.__all__mesh_dim=0写死_find_tp_mesh_dim(placements, class_dim)reduction="sum"非 TP 维Partial4.3 P1
loss_parallel_cross_entropy()DistributedCrossEntropyFunctionCrossEntropyLoss组件Shard(-1);sum+global_valid_tokensnn.CrossEntropyLoss、MindSpore 变体reduction="mean"多维约束NotImplementedError(与 PyTorch 相同限制)4.4 P2
ChunkedLossWrapperLLMTrainer可选 ChunkedLosstp_cp_example、三端差异说明4.5 里程碑
4.6 不在本 Issue 范围
label_smoothingparallelize_modulegather_tensor_parallel_logits重写五、测试计划
loss_parallel(1D/2D);与 Titan_LossParallelCrossEntropytests/torch/loss_parallel、tests/mindspore/st/loss_parallel附录:三端 API 对照
loss_parallel()loss_parallel(mesh, strict)distributed_cross_entropy_LossParallelCrossEntropy.applyCrossEntropyLossChunkedLossWrapperis_loss_parallel_active()关联:changzherui1/hyper-parallel#1 · changzherui1/hyper-parallel#6 · mindspore/hyper-parallel#270