ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ONNX Runtime ORTModule ModuleWithLoss 包装器:在 ORT 内计算损失并启用标签稀疏优化

2026/9/13 16:39:05 拓冰建站 浏览量
ONNX Runtime ORTModule ModuleWithLoss 包装器:在 ORT 内计算损失并启用标签稀疏优化 ONNX Runtime ORTModule ModuleWithLoss 包装器在 ORT 内计算损失并启用标签稀疏优化【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime导读本文围绕 ONNX Runtime 训练模块ORTModule的ModuleWithLoss包装器展开说明如何参照 Optimum 中的 ModuleWithLoss 思路将损失loss计算并入 ONNX RuntimeORT的执行图内从而让 ORT 的图优化能力尤其是标签稀疏优化 Label Sparsity Optimization在训练中生效。读完本文后你将掌握什么场景下必须使用该包装器、如何用几行nn.Module代码实现它、如何把它接入 ORTModule 的标准训练循环以及标签稀疏优化在 ORT 底层是如何被检测与执行的。为什么需要 ModuleWithLoss损失计算在 ONNX 图内的可见性ORTModule 训练的核心机制是把 PyTorch 模型的forward通过 torch.onnx.export 导出为 ONNX 图再交给 ONNX Runtime 的图优化器与执行引擎完成前向、反向与权重更新。因此只有出现在模型 forward 中的计算才会进入 ONNX 图ORT 才能针对它做优化。文档明确给出了两个必须做包装器适配的典型场景损失不在模型的 forward 路径中计算。如果训练脚本在模型输出之后、用独立的 Python 代码计算 loss例如loss criterion(logits, labels)写在模型调用之外ONNX Runtime 无法在 ONNX 图中看到这段损失计算自然也无法对其施加任何图优化。forward 虽然计算了损失但同时也返回了后续计算不需要的额外输出。这种情况下直接使用原始模型包装 ORTModule反向传播backward阶段会在 CUDA 设备上为这些无用张量保留不必要的显存造成浪费。把包装器的forward只返回标量 loss可以让导出图只保留必需的输出。一句话概括包装器的目的是把模型前向 损失计算整体收进一个nn.Module的forward里使损失成为 ONNX 图的一部分既减少显存占用又为标签稀疏优化等 compute optimizer 提供入口。Step 1实现ModuleWithLoss包装类实现方式非常直接继承torch.nn.Module在__init__中持有原始模型在forward中先调用模型得到 logits再用nn.CrossEntropyLoss计算损失最后只返回 lossclass ModuleWithLoss(nn.Module): def __init__(self, model): super().__init__() self.model model def forward(self, inputs, labels): # Perform the forward pass of the model lm_logits self.model(inputs) # Compute the cross-entropy loss loss nn.CrossEntropyLoss()(lm_logits, labels) return loss要点说明包装器只返回 loss 一个输出这正是文档第二个场景的解法即便原模型 forward 内部算出了 logits 等中间量只要它们不作为包装器输出就不会在反向阶段造成额外的 CUDA 显存开销。CrossEntropyLoss需要出现在包装器的forward里而不是训练脚本里这样它才能被导出进 ONNX 图成为SoftmaxCrossEntropyLossInternal/SoftmaxCrossEntropyLoss节点供后续的 InsertGatherBeforeSceLoss 图变换识别。实际项目中inputs可能是input_ids、attention_mask等多个张量可按需扩展forward签名文档给出的代码按(inputs, labels)二元组组织是最小可运行形态。Step 2定义一个可运行的示例训练脚本文档同时给出了一套配套的训练脚本骨架包含模型定义与训练循环两部分。定义PretrainedModel模型类同样是普通nn.Moduleforward 中依次经过 transformer 层得到 hidden states再通过 lm_head语言模型头输出 logits# Define the model architecture class PretrainedModel(nn.Module): ... def forward(self, input_ids, attention_mask): ... transformer_outputs self.transformer( input_ids, attention_maskattention_mask, ... ) hidden_states transformer_outputs[0] lm_logits self.lm_head(hidden_states) return lm_logits注意此处forward只返回 logits语言建模场景中最终计算损失所需的标签labels由训练循环另行传入包装器。这种模型只出 logits、包装器负责 loss的分层正是避免额外输出污染 ONNX 图的关键。训练循环包装器 ORTModule训练循环的关键点在于包装顺序先把PretrainedModel包成ModuleWithLoss再由ORTModule包裹最外层# Training loop model PretrainedModel(...) model ModuleWithLoss(model) optimizer torch.optim.Adam(model.parameters(), lr0.001) model ORTModule(model) for inputs, labels in dataloader: optimizer.zero_grad() # Compute the forward pass and cross-entropy loss loss model(inputs, labels) # Backward pass and optimization step loss.backward() optimizer.step()需要理解的两个顺序细节ModuleWithLoss必须在ORTModule之前包装ORTModule(model)包裹的是一个输入为 (inputs, labels)、输出为 loss的完整模块这样 ORTModule 导出的 ONNX 图才会同时包含模型前向与损失计算。optimizer必须在包ORTModule之前创建ModuleWithLoss内部的model.parameters()与包装后模型参数一致Adam(model.parameters(), ...)拿到的是原始可学习参数ORTModule包装不改变参数对象因此梯度更新链路不受影响。标签稀疏优化的底层原理从 Hook 检测到图变换使用ModuleWithLoss的收益之一是激活 ORT 的标签稀疏优化。该优化在仓库中有完整的实现链路从上到下可以分为三层。1. 运行时检测导出阶段注册的 label sparsity hook在 ORTModule 导出模型前GraphTransitionManager._add_check_label_sparsity_hook 会遍历模块树为每个torch.nn.modules.loss.CrossEntropyLoss子模块注册 forward pre-hook。hook 在导出运行时会统计标签密度valid_token torch.count_nonzero(label_input - module.ignore_index) total_token label_input.numel() label_density float(valid_token) / float(total_token) * 100当label_density 90即稀疏度超过 10%时优化被判定为值得开启结果写入runtime_inspector._sceloss_module_to_ignore_density_map同时用FlagAndPrintDensity.apply(...)包裹标签张量为后续图变换打标记见 _graph_transition_manager.py。这些 hook 在导出结束后会被移除第 781-786 行不影响正常训练运行。2. 选项传递密度信息进入 runtime options检测到的各 SCE 模块密度会以模块名:密度%的逗号分隔字符串写入_runtime_options.label_sparsity_ratio见 _graph_execution_manager.py。对应的配置项定义在 _options.pyenable_compute_optimizercompute optimizer 总开关默认Trueenable_label_sparse_optimizer标签稀疏优化开关默认Truelabel_sparsity_ratio记录实际检测到的标签密度字符串类型默认空串。3. 图变换InsertGatherBeforeSceLoss真正修改 ONNX 图的是名为 InsertGatherBeforeSceLoss 的GraphTransformer。它会把形如下面的 SCE 子图logits [token_count, classes] labels [token_count] \ / SCE Node(ignore_index-100) / \ loss (scalar 或 [token_count]) log_prob [token_count, classes]变换为先求有效 token索引、再做 ShrunkenGather 裁剪的稀疏计算图labels [token_count] -- Sub(ignore_index) -- NonZero -- Squeeze | indices of valid token [valid_token_count] | logits [token_count, classes] -- ShrunkenGather ShrunkenGather [valid_token_count, classes] [valid_token_count] \ / SCE Node (ignore_index-100)由于valid_token_count token_count恒成立Transformer 场景下token_count通常等于 batch size × sequence lengthclasses通常等于词表大小插入 ShrunkenGather 后损失计算的 FLOP 直接下降上游的 Gather 类图优化还会进一步削减其他算子的计算量这就是标签稀疏优化的收益来源。该变换只在同时满足以下条件时才会应用见 sceloss_compute_optimization.hSCE 节点的 reduction 属性为sum或mean保证 loss 是标量第二个输出log_prob不是图输出、也不被其他节点消费存在ignore_index且是常量标量如常见的-100标签输入不是由ShrunkenGather产生防止变换重复套用其后跟有FlagAndPrintDensity标记即运行时 hook 检测到了稀疏标签。使用前提与注意事项并非所有场景都需要包装器只有当损失不在 forward 内计算、或 forward 输出了后续不需要的张量时才需要ModuleWithLoss适配。常规模型 forward 内已算好 loss 且只返回 loss的用法直接ORTModule(model)即可。保持 forward 输出精简包装器 forward 只应返回 loss任何多余输出都会在反向时占用 CUDA 显存与包装的初衷相悖。优化是条件触发的标签稀疏优化由运行时的标签密度阈值 90%自动判定且要求 SCE 满足 reduction、ignore_index、log_prob 输出等图结构条件不满足条件时 ORT 会照常执行原始 SCE 子图训练结果不受影响。文档代码按最小示例组织文中的PretrainedModel的 transformer 层与 lm_head 部分以...占位实际使用时需按你的模型结构补全forward的参数如input_ids、attention_mask也需要与你的数据加载器输出对齐。延伸阅读ModuleWithLoss 包装器文档本文的原始依据ORTModule 导出与 hook 检测实现orttraining/orttraining/python/training/ortmodule/_graph_transition_manager.pycompute optimizer 运行时选项orttraining/orttraining/python/training/ortmodule/options.py标签稀疏图变换源码sceloss_compute_optimization.h 与 sceloss_compute_optimization.ccORTModule 训练整体指南docs/ORTModule_Training_Guidelines.md【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考