本文档描述在 HyperParallel 中实现与 PyTorch torch.distributed.tensor.parallel.loss_parallel 语义对齐的 交叉熵(Cross-Entropy)分片 logits 训练路径:在 类别维(词表维)张量并行 下,无需先将 logits 全量聚合即可正确计算 F.cross_entropy / nn.CrossEntropyLoss 及其反向。本文不含具体代码实现,仅供评审与迭代。
torch.distributed.tensor.parallel.loss_parallel
F.cross_entropy
nn.CrossEntropyLoss
列并行张量并行(Tensor Parallelism, TP)下,最后一层线性层通常将 logits 在 类别维 (C)(一般为最后一维,形状 (..., V) 中的 (V))上按 TP 组切分为多个 本地分片。若训练循环直接对本地张量调用 torch.nn.functional.cross_entropy:
(..., V)
torch.nn.functional.cross_entropy
all_gather
PyTorch 在上游提供 loss_parallel() 上下文:在上下文内对 DTensor 路径注册自定义算子实现,使 cross_entropy 在 类别维 Shard 的布局下仍与 单卡全 logits 语义一致,并通过 稳定的分布式 log-softmax 与 带规约的 NLL 控制通信类型与次数。
loss_parallel()
DTensor
cross_entropy
torch/distributed/tensor/parallel/loss.py
ColwiseParallel
Shard(-1)
use_local_output=False
get_train_context(loss_parallel_enabled)
hyper_parallel/core/shard/_op_dispatch.py
OpDispatcher
DTensor.__torch_function__
DistributedOp
preprocess
_dispatch_new
HYPER_PARALLEL_OPS_YAML_DIR
HYPER_PARALLEL_OPS_PYTHON_PATH
parallel_ops_register
register_distributed_op
op_name
dispatch
loss_parallel
DeviceMesh
Shard
Layout
alias_placements
OpDispatcher.dispatch
本期 固定采用方案 A,不再以 YAML 注册 DistributedOp 作为 loss_parallel 的主路径。
_dispatch_layout_infer
contextvars.ContextVar
op_name = platform.get_op_name(op_call)
_dispatch_loss_parallel(op_call, args, kwargs)
与方案 B / C 的关系(非本期主路径):
_dispatch_loss_parallel
ContextVar
all_reduce(MAX)
all_reduce(SUM)
Partial
loss.py
platform.get_op_name
LayoutCacheManager
reduction
disable_loss_parallel
forward_backward_step
validation_context
以下为建议的 接口清单;具体 命名 可与 HyperParallel 命名规范统一(例如前缀 hp_ 或模块 hyper_parallel.distributed.loss_parallel)。
hp_
hyper_parallel.distributed.loss_parallel
可选参数(若需对齐 TorchTitan / 扩展)
mesh
DeviceMesh | None
None
strict
bool
True
False
说明:PyTorch 上游 loss_parallel() 当前多为 无参;HyperParallel 若增加参数,需在文档中标注 与 PyTorch 的差异。
get_train_context
(enable_loss_parallel: bool) -> Callable[[], ContextManager]
enable_loss_parallel
__enter__
建议放在 并行配置结构体(与 TorchTitan ParallelismConfig 对齐)或 HyperParallel 等价配置中:
ParallelismConfig
loss_parallel_strict_mesh
parallelize_module
最后一层 ColwiseParallel(或文档化别名)建议明确:
output_layouts
use_local_output
input_layouts
在 loss_parallel 启用 时,建议文档固定下列支持关系(与 PyTorch loss.py 对齐):
input
target
weight
size_average
ignore_index
reduce
none
mean
sum
label_smoothing
RuntimeError
ValueError
get_op_name
本节列出本期为实现 loss_parallel(方案 A) 需要 新增或冻结契约 的 全部对外接口、配置项、内部扩展点及异常;命名可采用 hyper_parallel.distributed.loss_parallel(或项目统一前缀),下表以 逻辑名 为准。
@contextmanager
TrainContext
with
is_loss_parallel_active
() -> bool
strict=True
fallback_gather
with loss_parallel():
contextlib.nullcontext
__call__() -> AbstractContextManager[None]
with trainer.train_context():
以下字段建议置于 ParallelismConfig(或 HyperParallel 等价 TrainingParallelismConfig)中,类型均为 模块加载时可解析 的静态值。
TrainingParallelismConfig
parallelize
loss_parallel(strict=...)
loss_parallel_fallback_on_layout_error
loss_parallel_enabled
tp_enabled and not disable_loss_parallel
TrainContext | None
Trainer.train_context
下列 非新类型,但为实现 loss_parallel 必须写入文档的取值约定:
Shard(序列维)
Shard(类别维)
Replicate()
use_local_output=True
op_call: Callable
args: tuple
kwargs: dict
Any
_loss_parallel_active
contextvars.ContextVar[bool]
__exit__
LOSS_PARALLEL_OP_NAMES
frozenset[str]
register_loss_parallel_op_names(*names)
loss_parallel_token: int
distributed_log_softmax
dim
mesh_dim
distributed_nll_loss_forward
total_weight
fused_nll_log_softmax_backward
LossParallelLayoutError
LossParallelUnsupportedError
NotImplementedError
CrossEntropyLoss
下列为 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)。
torch.nn.CrossEntropyLoss
本文档描述在 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:all_gather,得到完整词表后再算 softmax,通信与显存开销大;PyTorch 在上游提供
loss_parallel()上下文:在上下文内对DTensor路径注册自定义算子实现,使cross_entropy在 类别维 Shard 的布局下仍与 单卡全 logits 语义一致,并通过 稳定的分布式 log-softmax 与 带规约的 NLL 控制通信类型与次数。1.2 参考实现与生态位置
torch/distributed/tensor/parallel/loss.py(loss_parallel()上下文 + ATen 算子级劫持)。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 目标
F.cross_entropy/nn.CrossEntropyLoss的 标量损失与梯度(对 logits) 与 单进程全 logits 参考实现在约定误差内一致。all_gather作为默认路径(允许可选调试或降级路径显式开启)。DistributedOp分发模型 集成,文档化 约束条件(如一维 TP mesh、shard 维约定)。ColwiseParallel+Shard(-1))及 Trainer / 集成训练入口 的配置项对齐,便于从 TorchTitan 迁移。2.2 非目标(本期可不实现或单列后续)
loss_parallel文档声明 不支持 的能力保持一致或明确报错。3. 术语与符号
DeviceMesh的一维子网格对应。Shardplacement,HyperParallel 中与现有Layout/alias_placements约定一致。4. 方案概述
4.1 总体思路
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之前(或与之等价的单点),增加 条件判断。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,行为与未引入本特性前一致。与方案 B / C 的关系(非本期主路径):
DistributedOp+ YAML):不采用作为 loss_parallel 的默认实现,避免与「仅上下文中启用」强约束下 layout 缓存键、算子表 的耦合复杂化;若未来有算子需 无上下文 也走同一套数学,可再评估是否复用内核代码。loss_parallel()):可作为 Torch 后端优化或对照路径,在单独评估 DTensor 互操作 后作为 可选实现;主路径仍以方案 A 的 Hyper 内聚实现为准。4.1.2 小结
dispatch首段特判 + 独立_dispatch_loss_parallel模块,语义对齐 PyTorch「上下文内注册临时算子行为」的意图,且 实现集中、与 YAML 表正交。ContextVar与训练循环一致性(见第 5 节)。4.2 数学语义(实现契约,非代码)
all_reduce(MAX)与all_reduce(SUM)得到,再构造与各分片一致的 log-softmax 张量。all_reduce(SUM)或使用Partial+ 规则化 reduce)得到与单卡一致的 loss。all_gatherlogits(与 PyTorchloss.py设计动机一致)。5. 组件设计
5.1 上下文管理器
contextvars.ContextVar或等价机制,避免多线程训练场景下状态串扰。5.2 算子分发层
DTensor.__torch_function__→OpDispatcher.dispatch;新增loss_parallel激活时的路由条件。platform.get_op_name与 PyTorchcross_entropy分解路径 的对应表;文档附录维护 「已劫持 / 已自定义算子名列表」。LayoutCacheManager对 CE 路径的缓存键是否包含 上下文标志:若同一 layout 在上下文内外行为不同,缓存键必须区分,避免错误复用。5.3 分布式内核(逻辑模块)
reduction一致的输出;通信:依归约类型而定。5.4 并行策略与模型侧契约
ColwiseParallel(或 HyperParallel 等价样式),输出布局为Shard(-1)(或等价Shard在类别维);use_local_output=False(或等价:保持 DTensor 直至 CE)。disable_loss_parallel(布尔):为 True 时最后一层输出 Replicate 全 logits 或使用本地 Tensor,不依赖loss_parallel上下文(通信与显存代价更高)。5.5 训练框架集成
get_train_context(loss_parallel_enabled)或等价 API;在forward_backward_step(或等价步骤)中对 loss 计算与 backward 使用同一上下文。validation_context的模式对齐)。6. 对外接口与参数
以下为建议的 接口清单;具体 命名 可与 HyperParallel 命名规范统一(例如前缀
hp_或模块hyper_parallel.distributed.loss_parallel)。6.1 上下文管理器
loss_parallel()可选参数(若需对齐 TorchTitan / 扩展)
meshDeviceMesh | NoneNoneNone表示从输入 DTensor 推断(与上游 PyTorch「从一维 mesh 推断」一致)。strictboolTrueFalse时可降级为 警告 + local fallback(若实现)。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_parallelboolFalseTrue:禁用「分片 logits + loss_parallel」路径,最后一层改为 聚合 logits 或等价行为。loss_parallel_strict_meshboolTrue6.4 与
parallelize_module/ 样式的契约参数最后一层
ColwiseParallel(或文档化别名)建议明确:output_layoutsShard(-1)(或等价类别维 Shard)。use_local_outputFalse:输出保持 DTensor,供 CE 消费。input_layouts6.5
F.cross_entropy/nn.CrossEntropyLoss侧支持的参数矩阵在
loss_parallel启用 时,建议文档固定下列支持关系(与 PyTorchloss.py对齐):inputtargetweightsize_averageignore_indexreducereductionnone/mean/sum:需分别说明 分布式下的归约与分母。label_smoothingRuntimeError或明确文档告警。7. 错误与降级策略
loss_parallel上下文中却对 Shard logits 调用 CEValueError,信息中指明期望布局。loss_parallel上下文依赖路径。8. 测试与验收标准
loss_parallel且 复制全 logits 路径与 开启分片路径 在误差范围内一致。9. 风险与依赖
get_op_name与 PyTorch 版本差异10. 文档与交付物
get_op_name→ 语义 对照。11. 新增接口与参数总表(汇总)
本节列出本期为实现 loss_parallel(方案 A) 需要 新增或冻结契约 的 全部对外接口、配置项、内部扩展点及异常;命名可采用
hyper_parallel.distributed.loss_parallel(或项目统一前缀),下表以 逻辑名 为准。11.1 用户可见 API
loss_parallel@contextmanager或返回上下文管理器的工厂get_train_contextTrainContextwith目标,内部在启用时进入loss_parallel。is_loss_parallel_active(可选)() -> boolloss_parallel上下文中(基于ContextVar)。§11.1.1
loss_parallel参数meshDeviceMesh | NoneNoneNone时从参与运算的 DTensor 推断;推断失败且strict=True时抛错。strictboolTrueTrue:布局不满足 一维 mesh + 类别维 Shard 等契约时ValueError/ 专用异常;False:仅告警或走可选降级(若实现fallback_gather)。§11.1.2
get_train_context参数与返回enable_loss_parallelboolTrue:返回的上下文等价于with loss_parallel():;False:空上下文(contextlib.nullcontext等价)。TrainContext__call__() -> AbstractContextManager[None],即with trainer.train_context():形式(与 TorchTitanTrainContext对齐)。11.2 并行 / 训练配置(结构化配置字段)
以下字段建议置于
ParallelismConfig(或 HyperParallel 等价TrainingParallelismConfig)中,类型均为 模块加载时可解析 的静态值。disable_loss_parallelboolFalseTrue:不使用分片 logits +loss_parallel;最后一层策略由parallelize侧改为 Replicate 全 logits 或文档约定的聚合路径;get_train_context中enable_loss_parallel恒为逻辑假(见 §11.3)。loss_parallel_strict_meshboolTrueTrue:强制一维 TP mesh 与类别维 Shard;与loss_parallel(strict=...)可组合,优先级需在实现中固定(建议:配置为根,上下文参数覆盖仅用于测试)。loss_parallel_fallback_on_layout_errorboolFalseTrue时:Shard logits 在未进入上下文或布局非法时,允许all_gather全 logits 再 CE(慢路径,仅用于调试或兼容)。11.3 训练器 / 验证器集成参数(契约)
loss_parallel_enabledbooltp_enabled and not disable_loss_parallel(与 TorchTitan 一致);用于get_train_context(loss_parallel_enabled)。validation_contextTrainContext | NoneTrainer.train_context同一实例 或等价行为,保证验证阶段 CE 与训练一致。11.4
parallelize_module/ 最后一层样式(无新增类时仍须冻结的参数组合)下列 非新类型,但为实现 loss_parallel 必须写入文档的取值约定:
ColwiseParallel(输出层)input_layoutsShard(序列维)等)。output_layoutsShard(-1)(或Shard(类别维)的等价表达)。use_local_outputFalse(输出保持 DTensor)。ParallelismConfig驱动disable_loss_parallelFalse时上述输出布局生效;True时改为Replicate()+use_local_output=True(或项目约定的「全 logits」路径)。11.5 内部扩展点(实现层契约,供模块间对接)
_dispatch_loss_parallelOpDispatcher的方法或模块级函数op_call: Callable,args: tuple,kwargs: dict→Anyop_name命中时由dispatch调用;内部调用 §11.6 内核或委托 PyTorch(若启用方案 C)。_loss_parallel_activecontextvars.ContextVar[bool]Falseloss_parallel__enter__置True,__exit__恢复;支持嵌套时用 计数器 或 token 备份(由 §5.1 固定一种)。LOSS_PARALLEL_OP_NAMESfrozenset[str]或 可注册表platform.get_op_name结果register_loss_parallel_op_names(*names)扩展(测试注入)。LayoutCacheManager集成loss_parallel_token: int(或布尔)11.6 分布式内核子模块(逻辑接口,无具体代码)
distributed_log_softmaxdim、mesh、mesh_dim;出:本地 log-softmax,布局不变。distributed_nll_loss_forwardtotal_weight(若mean)。fused_nll_log_softmax_backward11.7 异常类型(建议新增)
LossParallelLayoutErrorValueErrorLossParallelUnsupportedErrorNotImplementedError或RuntimeErrorlabel_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)。