已开启
HyperParallel:`loss_parallel` 类功能设计文档 #114
changzherui创建于  4月29日
changzherui
changzherui成员
4月29日 创建

本文档描述在 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 参考实现与生态位置

  • PyTorchtorch/distributed/tensor/parallel/loss.pyloss_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.pyOpDispatcherDTensor.__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 smoothingtarget 为概率分布(非 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 总体思路

  1. 布局前提:模型最后一层输出为 DTensor,在 类别维 Shardtarget整型类别索引,在 mesh 上视为 复制(Replicate)(或等价:各 rank 持有相同 labels 张量)。
  2. 启用方式:训练侧使用 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-softmaxindex 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_parallelTrue 时,返回的上下文在 __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 meshShard 在类别维

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 必须为 类别维 ShardDTensor(在约定 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(若配置允许),避免静默数值错误。
多维 meshShard 不在类别维 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 子 meshNone 时从参与运算的 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_contextenable_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: dictAny 方案 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 入:本地 logitsdimmeshmesh_dim;出:本地 log-softmax,布局不变。
K2 distributed_nll_loss_forward index target + 可选 weight + reduction 入:分片 logits 或中间张量、targetignore_indexweightreduction;出:loss 张量total_weight(若 mean)。
K3 fused_nll_log_softmax_backward 融合反向 入:grad_output、forward 保存态;出:对 logits 分片的梯度

11.7 异常类型(建议新增)

异常名(建议) 基类 触发条件
LossParallelLayoutError ValueError 一维 meshShard 维DTensor 类型 等与契约不符。
LossParallelUnsupportedError NotImplementedErrorRuntimeError label_smoothing概率 target 等明确不支持的功能。

11.8 F.cross_entropy / CrossEntropyLoss 参数(沿用 PyTorch,无新关键字)

下列为 PyTorch 已有参数;本特性 不新增关键字,仅约束 支持矩阵(见 §6.5):

torch.nn.functional.cross_entropyinput, target, weight, size_average, ignore_index, reduce, reduction, label_smoothing

torch.nn.CrossEntropyLoss 构造参数:与模块初始化一致(weight, size_average, ignore_index, reduce, reduction, label_smoothing)。


likedislike
changzheruichangzherui成员
4月29日 将 changzherui1 设为负责人
Zzhangyuguo
5月27日 关联了pull request:fix: add loss parallel support for tensor parallel training