十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

ONNX Runtime ORTTraining:用 onnxblock 与 generate_artifacts 离线生成模型训练图

ONNX Runtime ORTTraining:用 onnxblock 与 generate_artifacts 离线生成模型训练图 ONNX Runtime ORTTraining用 onnxblock 与 generate_artifacts 离线生成模型训练图【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime本文围绕 ONNX Runtime 仓库中的 onnxblock 离线工具 README 展开讲解如何从一个前向推理的 ONNX 基础模型出发离线生成 ORT Training API 所需的训练模型、评估模型、优化器模型与 checkpoint 四类产物。读完本文你可以掌握artifacts.generate_artifacts的完整参数用法、如何用onnxblock编写自定义损失例如多输出加权 MSE 平均以及生成产物如何衔接在线训练循环Python API 与 C Trainer。一、训练需要准备哪些文件ORT Training API 执行训练时需要以下文件训练 ONNX 模型training model包含基础模型图、损失子图和梯度反向图评估 ONNX 模型eval model可选仅包含基础模型图与损失子图用于训练过程中的评估优化器 ONNX 模型optimizer model包含优化器更新图checkpoint 文件目录存放模型参数。这些文件不需要手工构造——仓库中的离线工具onnxruntime.training.artifacts模块 onnxblock模块负责一键生成。核心实现在 artifacts.py 与 onnxblock 包包 docstring 即“Offline tooling for generating files needed by ort training apis”。二、前置条件准备一个前向 ONNX 基础模型离线工具的起点是一个只做前向推理的 ONNX 模型它将被用作构造训练模型的基底。README 明确指出如果从 PyTorch 导出模型建议使用以下参数export_params: Truetraining: torch.onnx.TrainingMode.TRAININGdo_constant_folding: False示例导出命令torch.onnx.export(model, sample_inputs, base_model.onnx, export_paramsTrue, trainingtorch.onnx.TrainingMode.TRAINING, do_constant_foldingFalse)关键在于参数必须内嵌在导出的模型中即export_paramsTrue因为generate_artifacts需要把模型参数从 initializer 中抽取出来、转换为训练模型的输入进而计算它们的梯度。关闭常量折叠则保证权重不会被常量传播合并掉从而保持参数节点可被自动求导处理。三、快速上手generate_artifacts 生成四类训练产物假设前向模型已经生成README 给出的最小完整用法如下from onnxruntime.training import artifacts # 加载 onnx 模型 model_path model.onnx base_model onnx.load(model_path) # 定义需要计算梯度的参数 requires_grad [weight1, bias1, weight2, bias2] # 定义冻结不训练的参数 frozen_params [weight3, bias3] # 生成训练产物 artifacts.generate_artifacts(base_model, requires_gradrequires_grad, frozen_paramsfrozen_params, lossartifacts.LossType.CrossEntropyLoss, optimizerartifacts.OptimType.AdamW) # 成功后会在当前工作目录生成 4 个文件 # training_model.onnx, eval_model.onnx, checkpoint/, optimizer_model.onnxgenerate_artifacts的完整签名摘自 artifacts.py如下各参数含义比 README 描述得更细参数类型/默认值说明modelonnx.ModelProto或str基础模型或模型路径。若模型大于 2GB 必须传路径源码中阈值USE_PATH_THRESHOLD 2147483648字节artifacts.py#L16requires_gradlist[str] \| None需要计算梯度的参数名列表frozen_paramslist[str] \| None需要冻结的参数名列表与requires_grad有交集会直接抛RuntimeErrorartifacts.py#L164-L168lossLossType/onnxblock.Block/None损失函数枚举或自定义 onnxblockNone时不添加损失节点内部用PassThrough代替optimizerOptimType/onnxblock.Block/None优化器枚举或自定义 onnxblockNone时跳过优化器模型生成artifact_directorystr \| None产物保存目录None时写当前工作目录prefixstr 产物文件名前缀ort_formatbool False是否同时用convert_onnx_models_to_ort导出 ORT 格式Fixed 优化风格custom_op_librarystr \| None自定义算子库路径供梯度图构建与优化阶段使用additional_output_nameslist[str] \| None在损失输出之外追加的训练/评估模型输出名nominal_checkpointbool False额外生成名义 checkpoint用于端侧降低训练模型构建开销、减小随应用打包的 checkpoint 体积loss_input_nameslist[str] \| None指定只把哪些图输出送入损失函数None时所有图输出都参与损失计算内置的损失与优化器枚举分别是LossTypeMSELoss、CrossEntropyLoss、BCEWithLogitsLoss、L1Lossartifacts.py#L19-L28每个枚举在内部都映射到onnxblock.loss下对应的 Block 实现loss.pyOptimTypeAdamW、SGDartifacts.py#L31-L38。所有生成的 ModelProto 会沿用基础模型定义的 opset。若目标路径已存在产物会被覆盖源码中有明确的 overwrite 日志这点在脚本化重复生成时值得注意。四、onnxblock 机制Block 是如何“堆叠”出训练图的generate_artifacts并非简单的图拼接它内部的调用链见 artifacts.py#L187-L197是用onnxblock.base(model, model_path)上下文管理器把基础模型注册为全局可操作模型实现见 model_accessor.py它会先deepcopy模型避免污染原始 ModelProto把基础模型所有图输出喂给一个内部的_TrainingBlock该 Block 内部再调用你指定的损失 BlockTrainingBlock.__call__完成构建、shape 推断、输出注册后调用_training_graph_utils.build_gradient_graph反向构图见 onnxblock.py#L180-L211再追加梯度累加节点最终产出 training model 与 eval model 两个 ModelProto。从源码注释看onnxblock.py#L197-L202梯度图构建包含三步把模型参数移为模型输入、对模型执行 orttraining 图变换、把梯度图追加到优化后的模型上。构建完成后输入顺序为“用户输入 参数输入”输出顺序为“用户输出 参数梯度”——这正是在线训练时按名喂入数据、取回梯度的契约。每个 Block 的基类是 blocks.py 中的Block子类实现build(*args, **kwargs)方法向全局图追加节点__call__会在build之后自动执行onnx.checker.check_model校验保证任何时刻被操作的模型都是合法 ONNX 模型。常用的积木块包括二元运算Add/Sub/Mul/Div、一元运算Sigmoid/Log/Abs/Neg、ReduceMean/ReduceSum、Pow、Constant创建 float initializer、Clip、Cast、Linear随机权重 Gemm 层、InputLike按已有输入/输出克隆一个新图输入以及PassThrough等。损失函数 Block 本身就是用这些积木组合的。以MSELoss为例loss.py#L12-L49它把SubPow(2.0) 归约ReduceMean/ReduceSum/PassThrough对应reductionmean/sum/none组合成一个子图CrossEntropyLoss则直接添加SoftmaxCrossEntropyLoss节点并自动为标签输入推导 INT64 类型、去掉最后一个维度即predictions: (N, C) - labels: (N,)BCEWithLogitsLoss由Sigmoid Log 加减乘子图手工组合而成。五、进阶场景自定义损失函数加权多输出损失当损失不是内置枚举能表达的比如模型有两个输出、需要分别计算 MSE 后加权平均loss 0.4 * mse_loss1(output1, target1) 0.6 * mse_loss2(output2, target2)README 给出的方案就是继承onnxblock.Block组合积木块import onnxruntime.training.onnxblock as onnxblock from onnxruntime.training import artifacts # 自定义损失块接收两个输入对两者的损失做加权平均 class WeightedAverageLoss(onnxblock.Block): def __init__(self): self._loss1 onnxblock.loss.MSELoss() self._loss2 onnxblock.loss.MSELoss() self._w1 onnxblock.blocks.Constant(0.4) self._w2 onnxblock.blocks.Constant(0.6) self._add onnxblock.blocks.Add() self._mul onnxblock.blocks.Mul() def build(self, loss_input_name1, loss_input_name2): # build 方法定义块如何堆叠在 loss_input_name1 / loss_input_name2 之上 return self._add( self._mul(self._w1(), self._loss1(loss_input_name1, target_nametarget1)), self._mul(self._w2(), self._loss2(loss_input_name2, target_nametarget2)) ) my_custom_loss WeightedAverageLoss() base_model onnx.load(model.onnx) requires_grad [weight1, bias1, weight2, bias2] frozen_params [weight3, bias3] artifacts.generate_artifacts(base_model, requires_gradrequires_grad, frozen_paramsfrozen_params, lossmy_custom_loss, optimizerartifacts.OptimType.AdamW)要点说明build方法的返回值是从该块输出继续构图或开始反向传播的张量名因此WeightedAverageLoss.build返回Add节点的输出名Constant(0.4)这类无参块调用self._w1()即创建一个 shape 为[1]的 float initializer实现见 blocks.py#L274-L290MSELoss(...)支持target_name关键字参数指定目标输入名若图中尚无该输入Block 会自动通过InputLike克隆一个新图输入loss.py#L46-L48所以示例中的target1/target2无需预先在基础模型中声明generate_artifacts对自定义损失的判定逻辑是loss要么是LossType枚举、要么必须是onnxblock.Block实例否则抛RuntimeErrorartifacts.py#L119-L133——也就是说自定义损失必须由 Block 自身控制损失节点的创建。同理loss与optimizer两个参数都接受onnxblock.Block优化器侧也可以完全自定义例如接入onnxblock.optim.AdamW其节点参数bias_correction、betas(0.9, 0.999)、eps1e-6、weight_decay0.0、可选clip_grad见 optim.py#L249-L294。内置的OptimType.AdamW/OptimType.SGD会生成独立的optimizer_model.onnx其中分别放置com.microsoft域的AdamWOptimizer与SGDOptimizerV2节点optim.py#L102-L197模型输入固定包含learning_ratefloat[1]、stepint64[1]以及params/gradients/first_order_moments/second_order_moments等序列输入。checkpoint 的保存由 checkpoint_utils.py 完成底层调用 pybind 的save_checkpoint将“可训练参数 冻结参数”分别序列化写入 checkpoint 目录若nominal_checkpointTrue会额外生成一份只含名义信息的 checkpoint适合端侧场景artifacts.py#L234-L237。六、产物如何接入在线训练循环四类产物生成后即可被在线训练 API 加载执行。仓库提供了两层参考实现1. Python 训练 API 循环见 api/README.mdfrom onnxruntime.training.api import Module, Optimizer, CheckpointState # 加载 checkpoint 状态 state CheckpointState.load_checkpoint(checkpoint.ckpt) # 创建 Module 与 Optimizer model Module(training_model.onnx, state, eval_model.onnx) optimizer Optimizer(optimizer.onnx, model) # 训练模式执行一个训练步 model.train() training_model_outputs model(训练模型输入) # 优化器步 optimizer.step() # 评估模式执行评估步 model.eval() eval_model_outputs model(评估模型输入) # 假设训练模型输出第一个元素是 loss print(Loss : , training_model_outputs[0]) # 保存 checkpoint CheckpointState.save_checkpoint(state, checkpoint_export.ckpt)2. C 参考 Trainertrainer.cc它展示了生产级训练循环的完整形态——用Ort::TrainingSession一次性传入训练图、评估图可选与优化器图CheckpointState::LoadCheckpoint加载初始 checkpointtrainer.cc#L221-L240每个 batch 调用session.TrainStep(inputs)取得 loss并按gradient_accumulation_steps决定是否执行OptimizerStep()、SchedulerStep()线性学习率调度由RegisterLinearLRScheduler注册与LazyResetGrad()trainer.cc#L277-L314按eval_interval周期性执行EvalStep按checkpoint_interval周期性SaveCheckpoint并支持在 checkpoint 中附带epoch、loss、framework等属性。七、小结与适用前提本工具链的定位是离线生成 在线执行generate_artifacts负责把“基础模型 损失 Block 参数梯度策略”编译成四个训练产物在线侧只负责按步执行传入模型若超过 2GB必须传文件路径而非ModelProto底层依赖 ONNX 的infer_shapes/check_model路径版本 APIrequires_grad与frozen_params不可有交集且默认所有参数都不参与求梯度必须显式声明见TrainingBlock.requires_gradonnxblock.py#L121-L139自定义损失/优化器统一走onnxblock.Block扩展点积木式组合保证了每步操作后的模型可被 ONNX checker 校验更多复杂场景大模型、多机、ZeRO 等可参考仓库内orttraining/orttraining/python/training/下的 ORTModule、amp、optim 等模块本文所述的 onnxblock 离线工具是其中最轻量、面向单进程场景的入口。【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表