
ONNX Runtime 训练 API 实战基于 ONNX 模型构建可运行的训练循环【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime本文基于 ONNX Runtime 仓库中orttraining/orttraining/python/training/api/README.md的入门指南展开讲解如何用 ONNX Runtime Training APIModule、Optimizer、CheckpointState围绕四件套训练工件训练模型、评估模型、优化器模型、checkpoint 文件搭建完整的训练循环并结合 module.py、checkpoint_state.py、optimizer.py 等源码深入解析每个 API 的参数含义、底层行为与进阶能力梯度重置、参数缓冲、推理模型导出、学习率调度帮助读者把这份简明指南落地为可复现、可运行的训练代码。一、训练前需要准备哪些工件原始文档 README.md 开宗明义ORT 训练 API 执行训练需要以下文件训练 ONNX 模型training model必需评估 ONNX 模型eval model可选优化器 ONNX 模型optimizer model可选checkpoint 文件必需。这套用 ONNX 图表达模型与优化器、用 checkpoint 承载参数状态的设计意味着你不需要在 Python 中手写张量运算训练与优化都发生在 ONNX 计算图里。上述工件无需手工构造仓库提供了统一的生成入口 generate_artifacts它从基础 ONNX 模型出发产出四类工件from onnxruntime.training import generate_artifacts, LossType, OptimType generate_artifacts( model, # onnx.ModelProto 或模型文件路径2GB 时用路径 requires_grad[...], # 需要计算梯度的参数名列表 frozen_params[...], # 需要冻结的参数名列表 lossLossType.MSELoss, # 损失函数 optimizerOptimType.AdamW, # 优化器类型 artifact_directoryart, # 工件输出目录默认为当前工作目录 prefix, # 工件文件名前缀 ort_formatFalse, # 是否以 ORT 格式保存工件 nominal_checkpointFalse, # 是否额外生成 nominal checkpoint )从 artifacts.py 的文档字符串可以确认该函数生成的四类工件与 README 的清单一一对应训练模型包含基础模型图、loss 子图和梯度图评估模型包含基础模型图和 loss 子图无梯度checkpoint存放模型参数优化器模型包含优化器计算图。其中LossType提供MSELoss、CrossEntropyLoss、BCEWithLogitsLoss、L1Loss四种枚举OptimType提供AdamW与SGD两种枚举也可以直接传入 onnxblock 自定义块。更细粒化的建模方式可参考文档指向的 onnxblock README。二、训练循环完整的可运行代码工件生成之后README 给出的核心训练循环如下已按源码核对补全语义from onnxruntime.training.api import Module, Optimizer, CheckpointState # 1. 加载 checkpoint 状态 state CheckpointState.load_checkpoint(checkpoint.ckpt) # 2. 创建 Module 和 Optimizer model Module(training_model.onnx, state, eval_model.onnx) optimizer Optimizer(optimizer.onnx, model) # 3. 训练模式 前向/反向一步 model.train() training_model_outputs model(inputs to your training model) # 4. 优化器更新参数 optimizer.step() # 5. 评估模式 评估一步 model.eval() eval_model_outputs model(inputs to your eval model) # 假设 loss 是训练模型输出的第一个元素 print(Loss : , training_model_outputs[0]) # 6. 保存 checkpoint CheckpointState.save_checkpoint(state, checkpoint_export.ckpt)Module、Optimizer、CheckpointState三个类以及下文提到的LinearLRScheduler都在init.py 中通过__all__统一导出因此from onnxruntime.training.api import ...即可使用。整个循环的执行顺序值得注意先model.train()再调用模型完成一次前向反向然后才optimizer.step()。反向得到的梯度存放在 checkpoint state 中step()读取梯度更新参数若顺序颠倒先step()后训练优化器读到的是上一轮的陈旧梯度。三、Module训练/评估双模会话Module 是对训练/评估 ONNX 模型的封装其构造函数签名为Module( train_model_uri: os.PathLike, # 训练模型路径必需 state: CheckpointState, # checkpoint 状态对象必需 eval_model_uri: os.PathLike | None, # 评估模型路径可选 device: str cpu, # 设备默认 cpu session_options: SessionOptions | None None, # 会话选项可选 )设备字符串的解析规则从源码可以看到device支持type:id格式代码对device执行split(:)第一段是设备类型如cpu、cuda第二段是设备编号缺省为0。因此cuda:1会解析为cuda类型的 1 号设备cpu则为默认 CPU。__call__的双路径行为Module实现了__call__根据输入类型自动走两条执行路径见 module.py 中_take_generic_step与_take_step_with_ortvalues输入含numpy.ndarray走_take_generic_step内部调用train_step/eval_step输出统一转换为 numpy 数组单输出直接返回数组多输出返回元组输入全部为OrtValue走_take_step_with_ortvalues调用train_step_with_ort_values/eval_step_with_ort_values输出保持为OrtValue元组单输出则直接返回OrtValue。这条路径可以保持张量驻留设备显存避免反复的 CPU-GPU 拷贝适合高性能推理式数据流。模式切换与梯度管理model.train(modeTrue)/model.eval()切换内部training标志决定__call__走训练步还是评估步且input_names()/output_names()会随模式返回对应模型的输入输出名lazy_reset_grad()延迟重置梯度——设置内部状态使梯度在下次train()调用、新梯度计算之前才被重置对需要累积梯度或部分重置的场景有用。参数缓冲与推理导出Module还提供三组深度能力源码中均有明确文档get_parameters_size(trainable_onlyTrue)返回参数的元素个数float32 计get_contiguous_parameters(trainable_onlyFalse)将参数拷贝到一个连续缓冲区OrtValue中返回便于一次性读写全部参数copy_buffer_to_parameters(buffer, trainable_onlyTrue)把缓冲区内容回拷进训练会话参数。源码注释特别指出若模块加载自nominal checkpoint仅含参数名义信息、不含完整数据需要调用此函数把更新后的参数加载回 checkpoint 以补全状态——这正好对应generate_artifacts的nominal_checkpoint选项export_model_for_inferencing(inference_model_uri, graph_output_names)训练完成后导出纯推理模型。它遍历训练图、定位产生指定输出名的节点并裁剪掉其后与推理无关的节点梯度、优化器相关子图得到干净的推理 ONNX。四、CheckpointState训练会话状态的容器CheckpointState 承载训练会话的全部状态模型参数、梯度、优化器状态和用户自定义属性。加载与保存# 加载支持完整 checkpoint 或 nominal checkpoint state CheckpointState.load_checkpoint(checkpoint.ckpt) # 保存include_optimizer_stateTrue 时连优化器状态一起落盘 CheckpointState.save_checkpoint(state, checkpoint_export.ckpt)注意save_checkpoint的include_optimizer_state参数默认是False默认只保存参数等状态不保存优化器动量等状态如果需要在断点处完整恢复训练含 AdamW 的一阶/二阶矩应显式传True。参数的字典式访问state.parameters返回一个类字典对象支持按名访问与写回for name, param in state.parameters: print(name, param.requires_grad) # 是否参与梯度 data param.data # numpy 数组 grad param.grad # numpy 数组无梯度时为 None if weight in state.parameters: state.parameters[weight] new_value # numpy 数组直接写回其中Parameter对象暴露name、data可读可写、grad、requires_grad四个属性Parameters支持__getitem__/__setitem__/__contains__/ 迭代与len()写回时若参数名不存在会抛KeyError。自定义属性state.properties是同样的字典式接口仅支持int、float、str三种值类型可用于把 epoch 数、随机种子等标量信息持久化进 checkpoint随save_checkpoint一起落盘。五、Optimizer 与学习率调度OptimizerOptimizer 绑定一个优化器 ONNX 模型和待训练的Moduleoptimizer Optimizer(optimizer.onnx, model) optimizer.step() # 按本轮梯度更新参数 optimizer.set_learning_rate(1e-4) lr optimizer.get_learning_rate()step()的语义是沿已计算梯度方向走一步具体更新规则由构造时传入的优化器模型决定即generate_artifacts选择的AdamW或SGD。构造时它还会复用 Module 的设备与SessionOptions保证优化器与模型在同一设备执行。LinearLRSchedulerinit.py 还导出了 LinearLRSchedulerfrom onnxruntime.training.api import LinearLRScheduler scheduler LinearLRScheduler(optimizer, warmup_step_count, total_step_count, initial_lr)其行为定义在源码文档中warmup 阶段学习率从 0 线性升到initial_lr随后按线性衰减的乘数从initial_lr一路降到 0全程以 warmup 前的初始学习率为基准。它要求训练循环中每一步都调用scheduler.step()与optimizer.step()配合完成一次完整的参数更新学习率调整。六、验证与延伸测试代码中的真实调用如果想看上述 API 在 C 侧对应的调用链与更完整的用法仓库提供了 trainer 测试orttraining/orttraining/test/training_api/trainer/它对 Python 层Module/Optimizer背后的 C 训练会话进行了直接测试可作为理解底层行为C.Module、C.Optimizer、C.CheckpointState、C.LinearLRScheduler均定义于 onnxruntime_pybind_module 注册的 pybind 绑定的参照。从 Python 源码结构看整条链路是generate_artifactsartifacts.py产出工件 →CheckpointState.load_checkpoint载入状态 →Module按当前模式执行train_step/eval_step→Optimizer.step消费梯度 →CheckpointState.save_checkpoint持久化 →export_model_for_inferencing产出纯推理模型。七、使用要点小结工件清单严格遵循 README训练模型 checkpoint 为必需评估模型与优化器模型可选但缺省训练循环如二节示例两者都用到模式切换必须放在对应__call__之前train()后调用是训练步eval()后调用是评估步输入是 numpy 时输出也是 numpy追求性能可全程使用OrtValue输入输出保持设备侧张量使用 nominal checkpoint 的场景记得用get_contiguous_parameters/copy_buffer_to_parameters完成参数读写闭环断点续训请显式开启include_optimizer_stateTrue保存优化器状态训练结束后用export_model_for_inferencing剥离训练专用节点得到可部署的推理 ONNX 模型。【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考