本文档描述在 HyperParallel 中实现与 PyTorch torch.distributed.tensor.parallel.loss_parallel 语义对齐的 交叉熵(Cross-Entropy)分片 logits 训练路径:在 类别维(词表维)张量并行 下,无需先将 logits 全量聚合即可正确计算 F.cross_entropy / nn.CrossEntropyLoss 及其反向。本文不含具体代码实现,仅供评审与迭代。
1. 背景
1.1 问题陈述
列并行张量并行(Tensor Parallelism, TP)下,最后一层线性层通常将 logits 在 类别维 (C)(一般为最后一维,形状 (..., V) 中的 (V))上按 TP 组切分为多个 本地分片。若训练循环直接对本地张量调用 torch.nn.functional.cross_entropy:
- 要么需要在 CE 前对 logits 沿 (V) 维
all_gather,得到完整词表后再算 softmax,通信与显存开销大;
- 要么数值错误(每个 rank 只对本地列做 softmax)。
PyTorch 在上游提供 loss_parallel() 上下文:在上下文内对 DTensor 路径注册自定义算子实现,使 cross_entropy 在 类别维 Shard 的布局下仍与 单卡全 logits 语义一致,并通过 稳定的分布式 log-softmax 与 带规约的 NLL 控制通信类型与次数。
1.2 参考实现与生态位置
- PyTorch:
torch/distributed/tensor/parallel/loss.py(loss_parallel() 上下文 + ATen 算子级劫持)。
- TorchTitan:最后一层
ColwiseParallel 配置 Shard(-1) 输出与 use_local_output=False;训练器通过 get_train_context(loss_parallel_enabled) 在 forward / backward 外包裹 loss_parallel()(或等价物)。
1.3 HyperParallel 现状
hyper_parallel/core/shard/_op_dispatch.py 中 OpDispatcher 经 DTensor.__torch_function__ 分发算子;扩展方式包括 YAML + DistributedOp 子类、preprocess → _dispatch_new、以及环境变量 HYPER_PARALLEL_OPS_YAML_DIR / HYPER_PARALLEL_OPS_PYTHON_PATH 注入外部算子定义。
parallel_ops_register 支持 register_distributed_op;若某 op_name 已在注册表中且 YAML 未声明,dispatch 会将该算子 自动纳入 layout 推断路径。
- 当前 未 提供与
loss_parallel 等价的 专用 CE 分布式内核 与 训练框架级默认封装。
2. 目的与范围
2.1 目标
| 编号 |
目标 |
| G1 |
在 logits 沿类别维为 Shard 的 DTensor 前提下,F.cross_entropy / nn.CrossEntropyLoss 的 标量损失与梯度(对 logits) 与 单进程全 logits 参考实现在约定误差内一致。 |
| G2 |
避免在 CE 前对 logits 做 全词表 all_gather 作为默认路径(允许可选调试或降级路径显式开启)。 |
| G3 |
对外提供 显式上下文(或等价机制),语义对齐 PyTorch:进入上下文则启用分布式 CE 规则;退出则恢复,且 forward 与 backward 须在启用规则的前提下成对成立。 |
| G4 |
与 HyperParallel DeviceMesh / Layout / DistributedOp 分发模型 集成,文档化 约束条件(如一维 TP mesh、shard 维约定)。 |
| G5 |
与 并行化模块(最后一层 ColwiseParallel + Shard(-1))及 Trainer / 集成训练入口 的配置项对齐,便于从 TorchTitan 迁移。 |
2.2 非目标(本期可不实现或单列后续)
- Label smoothing、target 为概率分布(非 index)等与 PyTorch
loss_parallel 文档声明 不支持 的能力保持一致或明确报错。
- 多维 DeviceMesh 上类别维 Shard 的一般情形(若上游仅保证 一维 mesh,本期可 显式限制 并文档化)。
- MindSpore 后端是否与本特性同期交付:建议在文档中单列为 平台矩阵,分期实现。
- MoE、Pipeline Parallel、Context Parallel 与 loss_parallel 组合的全矩阵验证:本期给出 支持矩阵 与 已知限制 即可。
3. 术语与符号
| 术语 |
含义 |
| 类别维 / class 维 |
logits 上 softmax 作用的维度,通常为 最后一维,长度 (V)(词表大小)。 |
| TP 组 / TP mesh |
持有 logits 不同列分片的设备集合;实现上与 DeviceMesh 的一维子网格对应。 |
| Shard(C) |
在类别维上的 Shard placement,HyperParallel 中与现有 Layout / alias_placements 约定一致。 |
| 稳定 softmax |
先 全局 max 再 全局 sum exp 的 log-softmax 分解,避免溢出;分布式实现通过 集合通信 拼出全局量。 |
| MaskPartial 语义 |
PyTorch DTensor 中的命名;指 仅部分 rank 对全局标签索引有贡献 时的 gather + 跨 rank 规约 组合。HyperParallel 可实现 等价数值行为,类名可不照搬。 |
4. 方案概述
4.1 总体思路
- 布局前提:模型最后一层输出为 DTensor,在 类别维 Shard;target 为 整型类别索引,在 mesh 上视为 复制(Replicate)(或等价:各 rank 持有相同 labels 张量)。
- 启用方式:训练侧使用
loss_parallel() 上下文(或 HyperParallel 命名,见第 6 节)包裹 forward + loss + backward 中与 CE 相关的区间;上下文与 OpDispatcher.dispatch 中的路由联动(见下文 方案 A)。
4.1.1 选定实现:方案 A(dispatch 上下文分支 + 专用内核)
本期 固定采用方案 A,不再以 YAML 注册 DistributedOp 作为 loss_parallel 的主路径。
| 要点 |
说明 |
| 截获点 |
在 OpDispatcher.dispatch 内、白名单 / 随机算子等既有分支之后、进入 _dispatch_layout_infer 之前(或与之等价的单点),增加 条件判断。 |
| 条件 |
(1)loss_parallel 上下文已激活(如 contextvars.ContextVar);(2)当前 op_name = platform.get_op_name(op_call) 属于 预置的 CE 相关算子集合(例如 cross_entropy 以及分解路径上的 log-softmax / NLL 相关 ATen 名,以实装时固化表为准)。 |
| 行为 |
命中时 不 走默认的 DistributedOp + layout 缓存 路径,转而调用 _dispatch_loss_parallel(op_call, args, kwargs)(名称可调整)同一套分布式实现:内部完成 稳定分布式 log-softmax、index NLL、融合反向(与第 4.2 节数学契约一致)。 |
| 未命中 |
仍走现有 _dispatch_layout_infer,行为与未引入本特性前一致。 |
| 退出上下文 |
无 CE 特判,避免影响其它算子。 |
与方案 B / C 的关系(非本期主路径):
- 方案 B(为相关算子注册
DistributedOp + YAML):不采用作为 loss_parallel 的默认实现,避免与「仅上下文中启用」强约束下 layout 缓存键、算子表 的耦合复杂化;若未来有算子需 无上下文 也走同一套数学,可再评估是否复用内核代码。
- 方案 C(委托 PyTorch
loss_parallel()):可作为 Torch 后端优化或对照路径,在单独评估 DTensor 互操作 后作为 可选实现;主路径仍以方案 A 的 Hyper 内聚实现为准。
4.1.2 小结
- 方案 A = 显式上下文 +
dispatch 首段特判 + 独立 _dispatch_loss_parallel 模块,语义对齐 PyTorch「上下文内注册临时算子行为」的意图,且 实现集中、与 YAML 表正交。
- 实施时需同步:(1)CE 算子名表;(2)layout 缓存是否绕过或带上下文位;(3)多线程下
ContextVar 与训练循环一致性(见第 5 节)。
4.2 数学语义(实现契约,非代码)
- Log-softmax:各 rank 仅持有 (z) 的部分分量时,全局 (\max) 与 (\sum \exp(z-\max)) 分别通过
all_reduce(MAX) 与 all_reduce(SUM) 得到,再构造与各分片一致的 log-softmax 张量。
- NLL:对每个样本的全局类别 (y),仅在持有对应列的分片上 非零贡献;通过 本地 gather 与 跨 rank 规约(如 对标量 loss 项
all_reduce(SUM) 或使用 Partial + 规则化 reduce)得到与单卡一致的 loss。
- 反向:建议采用 NLL 与 log-softmax 反向融合 的策略,避免朴素 autograd 分解引入 额外
all_gather logits(与 PyTorch loss.py 设计动机一致)。
5. 组件设计
5.1 上下文管理器
| 项目 |
说明 |
| 职责 |
进入时 启用分布式 CE 规则(注册分发钩子或切换调度分支);退出时 撤销,保证不影响其它模块与非 CE 算子路径。 |
| 线程安全 |
建议使用 contextvars.ContextVar 或等价机制,避免多线程训练场景下状态串扰。 |
| 嵌套语义 |
建议定义:可重入计数或 禁止嵌套二选一,并在文档中固定一种行为。 |
5.2 算子分发层
| 项目 |
说明 |
| 入口 |
保持现有 DTensor.__torch_function__ → OpDispatcher.dispatch;新增 loss_parallel 激活时的路由条件。 |
| 算子集合 |
需在实现阶段固化 platform.get_op_name 与 PyTorch cross_entropy 分解路径 的对应表;文档附录维护 「已劫持 / 已自定义算子名列表」。 |
| 缓存 |
LayoutCacheManager 对 CE 路径的缓存键是否包含 上下文标志:若同一 layout 在上下文内外行为不同,缓存键必须区分,避免错误复用。 |
5.3 分布式内核(逻辑模块)
| 模块 |
职责 |
| 分布式 log-softmax |
输入:类别维 Shard 的本地 logits;输出:同 Shard 的 log-softmax;通信:MAX + SUM all_reduce。 |
| 分布式 NLL(index target) |
输入:分片 logits 或中间量、全局 target、ignore_index、可选 weight;输出:标量或与 reduction 一致的输出;通信:依归约类型而定。 |
| 融合反向 |
输入:上游梯度、forward 保存的中间态;输出:对 logits 分片的梯度;通信:与实现对齐的最小集合。 |
5.4 并行策略与模型侧契约
| 项目 |
说明 |
| 最后一层 |
与 TorchTitan 对齐:ColwiseParallel(或 HyperParallel 等价样式),输出布局为 Shard(-1)(或等价 Shard 在类别维);use_local_output=False(或等价:保持 DTensor 直至 CE)。 |
| 关闭路径 |
配置项 disable_loss_parallel(布尔):为 True 时最后一层输出 Replicate 全 logits 或使用本地 Tensor,不依赖 loss_parallel 上下文(通信与显存代价更高)。 |
5.5 训练框架集成
| 项目 |
说明 |
| Trainer |
提供 get_train_context(loss_parallel_enabled) 或等价 API;在 forward_backward_step(或等价步骤)中对 loss 计算与 backward 使用同一上下文。 |
| 验证 / inference |
若验证阶段也计算 CE,需同样包裹上下文(与 TorchTitan validator 传入 validation_context 的模式对齐)。 |
6. 对外接口与参数
以下为建议的 接口清单;具体 命名 可与 HyperParallel 命名规范统一(例如前缀 hp_ 或模块 hyper_parallel.distributed.loss_parallel)。
6.1 上下文管理器
| 接口 |
类型 |
说明 |
loss_parallel() |
上下文管理器(无参或可扩展可选参数,见下表) |
启用分布式 CE 语义;必须与分片 logits + index target 的前提配合使用。 |
可选参数(若需对齐 TorchTitan / 扩展)
| 参数名 |
类型 |
默认值 |
含义 |
mesh |
DeviceMesh | None |
None |
显式指定 TP mesh;None 表示从输入 DTensor 推断(与上游 PyTorch「从一维 mesh 推断」一致)。 |
strict |
bool |
True |
布局不满足约定时是否 抛错;False 时可降级为 警告 + local fallback(若实现)。 |
说明:PyTorch 上游 loss_parallel() 当前多为 无参;HyperParallel 若增加参数,需在文档中标注 与 PyTorch 的差异。
6.2 训练上下文工厂
| 接口 |
签名(逻辑) |
说明 |
get_train_context |
(enable_loss_parallel: bool) -> Callable[[], ContextManager] |
当 enable_loss_parallel 为 True 时,返回的上下文在 __enter__ 中调用 loss_parallel();为 False 时返回 空上下文。 |
6.3 并行 / 训练配置项
建议放在 并行配置结构体(与 TorchTitan ParallelismConfig 对齐)或 HyperParallel 等价配置中:
| 配置项 |
类型 |
默认值 |
含义 |
disable_loss_parallel |
bool |
False |
True:禁用「分片 logits + loss_parallel」路径,最后一层改为 聚合 logits 或等价行为。 |
loss_parallel_strict_mesh |
bool |
True |
是否严格要求 一维 TP mesh 与 Shard 在类别维。 |
6.4 与 parallelize_module / 样式的契约参数
最后一层 ColwiseParallel(或文档化别名)建议明确:
| 参数 |
含义 |
output_layouts |
启用 loss_parallel 时为 Shard(-1)(或等价类别维 Shard)。 |
use_local_output |
False:输出保持 DTensor,供 CE 消费。 |
input_layouts |
与上游 序列并行 / TP 衔接,与现有 TorchTitan 文档一致。 |
6.5 F.cross_entropy / nn.CrossEntropyLoss 侧支持的参数矩阵
在 loss_parallel 启用 时,建议文档固定下列支持关系(与 PyTorch loss.py 对齐):
| PyTorch 参数 |
支持策略 |
input |
必须为 类别维 Shard 的 DTensor(在约定 mesh 上)。 |
target |
类别索引;Replicate 或与文档一致的布局。 |
weight |
可选;若为 DTensor,应为 Replicate;实现需说明 本地 weight 切片方式。 |
size_average |
废弃,遵循 PyTorch 行为。 |
ignore_index |
应支持。 |
reduce |
废弃,遵循 PyTorch 行为。 |
reduction |
none / mean / sum:需分别说明 分布式下的归约与分母。 |
label_smoothing |
不支持:建议 RuntimeError 或明确文档告警。 |
7. 错误与降级策略
| 场景 |
建议行为 |
未在 loss_parallel 上下文中却对 Shard logits 调用 CE |
报错或 回退全 gather(若配置允许),避免静默数值错误。 |
| 多维 mesh 或 Shard 不在类别维 |
ValueError,信息中指明期望布局。 |
| target 为 float(概率) |
不支持,与 PyTorch 对齐。 |
| 全局 disable_loss_parallel |
跳过 loss_parallel 上下文依赖路径。 |
8. 测试与验收标准
| 类型 |
内容 |
| 单元测试 |
分布式 log-softmax / NLL 子模块与 单卡参考对比(给定相同全局 logits 与 target)。 |
| 集成测试 |
TP 度为 2/4,词表维 shard,对比 loss 值与 logits.grad(或等价梯度检查)。 |
| 回归 |
关闭 loss_parallel 且 复制全 logits 路径与 开启分片路径 在误差范围内一致。 |
| 缓存 |
切换上下文前后 同一算子缓存未错误命中。 |
9. 风险与依赖
| 风险 |
缓解 |
get_op_name 与 PyTorch 版本差异 |
版本矩阵测试;文档列出支持的 torch 版本。 |
| Layout 缓存与上下文耦合 |
缓存键包含 上下文标志 或 禁用相关缓存。 |
| 与 FSDP / EP 组合 |
在 支持矩阵 中逐项验证;未验证组合 报错或文档免责。 |
10. 文档与交付物
| 交付物 |
说明 |
| 本文档 |
设计与接口契约。 |
| 用户文档 |
简述 何时启用、如何配置最后一层、Trainer 用法、与 TorchTitan 对照表。 |
| 附录:算子名映射表 |
实现完成后补充 get_op_name → 语义 对照。 |
11. 新增接口与参数总表(汇总)
本节列出本期为实现 loss_parallel(方案 A) 需要 新增或冻结契约 的 全部对外接口、配置项、内部扩展点及异常;命名可采用 hyper_parallel.distributed.loss_parallel(或项目统一前缀),下表以 逻辑名 为准。
11.1 用户可见 API
| 序号 |
接口名(建议) |
形态 |
参数 / 返回值 |
语义 |
| L1 |
loss_parallel |
@contextmanager 或返回上下文管理器的工厂 |
见 §11.1.1 |
进入:激活 CE 分布式调度;退出:恢复。嵌套语义见 §5.1。 |
| L2 |
get_train_context |
可调用对象工厂:TrainContext |
见 §11.1.2 |
根据配置生成 训练 / 验证 共用的 with 目标,内部在启用时进入 loss_parallel。 |
| L3 |
is_loss_parallel_active(可选) |
() -> bool |
无参数 |
调试或断言:当前线程是否在 loss_parallel 上下文中(基于 ContextVar)。 |
§11.1.1 loss_parallel 参数
| 参数名 |
类型 |
默认值 |
必填 |
说明 |
mesh |
DeviceMesh | None |
None |
否 |
显式指定 TP 子 mesh;None 时从参与运算的 DTensor 推断;推断失败且 strict=True 时抛错。 |
strict |
bool |
True |
否 |
True:布局不满足 一维 mesh + 类别维 Shard 等契约时 ValueError / 专用异常;False:仅告警或走可选降级(若实现 fallback_gather)。 |
§11.1.2 get_train_context 参数与返回
| 参数名 |
类型 |
默认值 |
说明 |
enable_loss_parallel |
bool |
— |
True:返回的上下文等价于 with loss_parallel():;False:空上下文(contextlib.nullcontext 等价)。 |
| 返回值(逻辑类型) |
说明 |
TrainContext |
协议:__call__() -> AbstractContextManager[None],即 with trainer.train_context(): 形式(与 TorchTitan TrainContext 对齐)。 |
11.2 并行 / 训练配置(结构化配置字段)
以下字段建议置于 ParallelismConfig(或 HyperParallel 等价 TrainingParallelismConfig)中,类型均为 模块加载时可解析 的静态值。
| 字段名 |
类型 |
默认值 |
说明 |
disable_loss_parallel |
bool |
False |
True:不使用分片 logits + loss_parallel;最后一层策略由 parallelize 侧改为 Replicate 全 logits 或文档约定的聚合路径;get_train_context 中 enable_loss_parallel 恒为逻辑假(见 §11.3)。 |
loss_parallel_strict_mesh |
bool |
True |
True:强制一维 TP mesh 与类别维 Shard;与 loss_parallel(strict=...) 可组合,优先级需在实现中固定(建议:配置为根,上下文参数覆盖仅用于测试)。 |
loss_parallel_fallback_on_layout_error |
bool |
False |
可选;True 时:Shard logits 在未进入上下文或布局非法时,允许 all_gather 全 logits 再 CE(慢路径,仅用于调试或兼容)。 |
11.3 训练器 / 验证器集成参数(契约)
| 位置 |
参数名 |
类型 |
说明 |
| Trainer 构造 |
由配置派生的 loss_parallel_enabled |
bool |
逻辑式:tp_enabled and not disable_loss_parallel(与 TorchTitan 一致);用于 get_train_context(loss_parallel_enabled)。 |
| Validator 构造 |
validation_context |
TrainContext | None |
与 Trainer.train_context 同一实例 或等价行为,保证验证阶段 CE 与训练一致。 |
11.4 parallelize_module / 最后一层样式(无新增类时仍须冻结的参数组合)
下列 非新类型,但为实现 loss_parallel 必须写入文档的取值约定:
| 样式 / 模块 |
参数名 |
启用 loss_parallel 时的取值 |
ColwiseParallel(输出层) |
input_layouts |
与上游 SP/TP 一致(常为 Shard(序列维) 等)。 |
| 同上 |
output_layouts |
Shard(-1)(或 Shard(类别维) 的等价表达)。 |
| 同上 |
use_local_output |
False(输出保持 DTensor)。 |
ParallelismConfig 驱动 |
disable_loss_parallel |
False 时上述输出布局生效;True 时改为 Replicate() + use_local_output=True(或项目约定的「全 logits」路径)。 |
11.5 内部扩展点(实现层契约,供模块间对接)
| 序号 |
名称(建议) |
形态 |
参数 / 成员 |
说明 |
| I1 |
_dispatch_loss_parallel |
OpDispatcher 的方法或模块级函数 |
**op_call: Callable, args: tuple, kwargs: dict → Any |
方案 A 核心:仅在上下文激活 + op_name 命中时由 dispatch 调用;内部调用 §11.6 内核或委托 PyTorch(若启用方案 C)。 |
| I2 |
_loss_parallel_active |
contextvars.ContextVar[bool] |
初始 False |
loss_parallel __enter__ 置 True,__exit__ 恢复;支持嵌套时用 计数器 或 token 备份(由 §5.1 固定一种)。 |
| I3 |
LOSS_PARALLEL_OP_NAMES |
frozenset[str] 或 可注册表 |
元素为 platform.get_op_name 结果 |
CE 分解路径上的算子全集(至少覆盖 forward/backward 所需条目);允许 register_loss_parallel_op_names(*names) 扩展(测试注入)。 |
| I4 |
LayoutCacheManager 集成 |
缓存键扩展 |
可选字段:loss_parallel_token: int(或布尔) |
上下文内外 布局推断结果不同时,禁止缓存键冲突(见 §9)。 |
11.6 分布式内核子模块(逻辑接口,无具体代码)
| 序号 |
逻辑模块名(建议) |
职责 |
输入 / 输出(抽象) |
| K1 |
distributed_log_softmax |
类别维 Shard 上的稳定 log-softmax |
入:本地 logits、dim、mesh、mesh_dim;出:本地 log-softmax,布局不变。 |
| K2 |
distributed_nll_loss_forward |
index target + 可选 weight + reduction |
入:分片 logits 或中间张量、target、ignore_index、weight、reduction;出:loss 张量及 total_weight(若 mean)。 |
| K3 |
fused_nll_log_softmax_backward |
融合反向 |
入:grad_output、forward 保存态;出:对 logits 分片的梯度。 |
11.7 异常类型(建议新增)
| 异常名(建议) |
基类 |
触发条件 |
LossParallelLayoutError |
ValueError |
一维 mesh、Shard 维、DTensor 类型 等与契约不符。 |
LossParallelUnsupportedError |
NotImplementedError 或 RuntimeError |
label_smoothing、概率 target 等明确不支持的功能。 |
11.8 F.cross_entropy / CrossEntropyLoss 参数(沿用 PyTorch,无新关键字)
下列为 PyTorch 已有参数;本特性 不新增关键字,仅约束 支持矩阵(见 §6.5):
torch.nn.functional.cross_entropy:input, target, weight, size_average, ignore_index, reduce, reduction, label_smoothing。
torch.nn.CrossEntropyLoss 构造参数:与模块初始化一致(weight, size_average, ignore_index, reduce, reduction, label_smoothing)。
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 114
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/114