mindspore-ai / mindspore-ai/hyper-parallel

HyperParallel:`loss_parallel` 类功能设计文档

Open
#741 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

本文档描述在 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)。


schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 114
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/114

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Read hyper_parallel/core/shard/_op_dispatch.py and PyTorch's torch/distributed/tensor/parallel/loss.py first. Trace OpDispatcher.dispatch and the existing DTensor/DistributedOp routing, then turn the stated Scheme A contracts into a reviewed design document covering supported layouts, context behavior, integration points, and tests. Done means the design and its acceptance criteria are agreed.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems, documentation
Issue type
Documentation
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.