
MNN 训练优化器使用指南SGD、ADAM、损失函数与学习率调度【免费下载链接】MNNMNN: A blazing-fast, lightweight inference engine battle-tested by Alibaba, powering high-performance on-device LLMs and Edge AI.项目地址: https://gitcode.com/GitHub_Trending/mn/MNN导读本指南围绕 MNN 训练框架MNN::Train中优化器Optimizer与损失函数Loss的使用展开涵盖 SGD with Momentum、ADAM 的完整配置流程以及框架内置的多种损失函数交叉熵、KL 散度、MSE、MAE、Hinge、蒸馏损失的接口与底层实现。读者读完本文后可以在自己的 MNN 训练或微调任务中正确创建优化器、绑定模型可训练参数、配置动量/权重衰减/正则化方式与学习率并通过solver-step(loss)完成一次完整的梯度计算与参数更新同时掌握内置损失函数的数学形式与调用前提。文中所有接口均与当前仓库 tools/train/source/optimizer 下的真实源码一一对应。一、优化器的核心抽象与整体流程在 MNN 的训练体系中优化器统一继承自抽象基类ParameterOptimizer定义见 ParameterOptimizer.hpp其核心职责是根据 loss 计算梯度并更新参数。ParameterOptimizer的关键设计点如下构造时绑定 Module优化器构造时接收一个std::shared_ptrExpress::Module通过module()-parameters()拿到模型全部参数再筛选出可训练参数集合mTrainable详见 ParameterOptimizer.cpp。文档示例中的solver-append(model-parameters())即完成设置模型中需要优化的参数这一步。step(loss)是统一的更新入口一次调用完成前向求 loss → 反向求梯度 → 正则化 → 动量/二阶矩修正 → 更新参数的全链路。正则化方式枚举enum RegularizationMethod { L1, L2, L1L2 }即支持 L1、L2、L1L2 三种默认 L2。训练步数管理currentStep()/setCurrentStep(int)维护全局优化步数mStepADAM 的偏差校正bias correction依赖该步数。从SGD::onGetNextParameter的实现SGD.cpp可以还原出step(loss)内部的典型流程OpGrad::grad(loss, trainable(), ...)基于 loss 对全部可训练参数求梯度对每个参数执行regularizeParameters(param, grad)叠加权重衰减正则化项调用onComputeUpdateValue(param, grad)计算带动量/二阶矩的更新量执行param - updateValue得到新参数值返回给上层写回。此外框架还提供了工厂方法用于一行创建优化器同样定义在 ParameterOptimizer.hppstatic ParameterOptimizer* createSGD(std::shared_ptrExpress::Module module, float lr, float momentum, float weightDecay, RegularizationMethod method); static ParameterOptimizer* createADAM(std::shared_ptrExpress::Module module, float lr, float momentum, float momentum2, float weightDecay, float eps, RegularizationMethod method);如果你不想逐个set配置项可以直接用这两个工厂函数一次性传入学习率、动量、权重衰减与正则化方法。二、SGD with Momentum 使用详解2.1 完整使用示例以下代码来自 docs/train/optim.md展示了 SGD 优化器从创建到更新的完整流程// 新建SGD优化器 std::shared_ptrSGD solver(new SGD); // 设置模型中需要优化的参数 solver-append(model-parameters()); // 设置momentum和weight decay solver-setMomentum(0.9f); solver-setWeightDecay(0.0005f); // 设置正则化方法默认L2 solver-setRegularizationMethod(RegularizationMethod::L2); // 设置学习率 solver-setLearningRate(0.001); // 根据loss计算梯度并更新参数 solver-step(loss);其中solver-append(model-parameters())对应ParameterOptimizer构造阶段对 Module 参数的收集逻辑实际工程中更常见的写法是std::shared_ptrSGD solver(new SGD(model));将模型直接传入构造函数参见 SGD.hpp 中SGD(std::shared_ptrExpress::Module module)的声明以及 SGD.cpp 中构造时自动为每个可训练参数初始化全零历史动量缓存mHistory[p] _Const(0.0f, ...)的实现。2.2 可配置参数与默认值对照 SGD.hpp 的成员定义SGD 优化器的全部配置项与默认值如下配置项Setter 接口默认值说明学习率setLearningRate(float rate)0.001f控制每次更新的步长currentLearningRate()可查询当前值动量setMomentum(float momentum)0经典 momentum一般取0.9附近getMomentum()可查询权重衰减setWeightDecay(float decay)0正则化强度配合RegularizationMethod使用正则化方法setRegularizationMethod(RegularizationMethod)L2可选L1/L2/L1L22.3 更新公式的源码级解读SGD 的核心更新逻辑在SGD::onComputeUpdateValueSGD.cppmHistory[param] lr * grad mMomentum * mHistory[param]; return mHistory[param]; // 上层执行 param param - updateValue对应经典动量 SGD 公式v_t lr * g_t momentum * v_{t-1}θ_t θ_{t-1} - v_t。正则化权重衰减如何生效regularizeParametersSGD.cpp在计算更新量之前先向原始梯度叠加正则项L1grad weightDecay * sign(param)L2grad weightDecay * paramL1L2grad weightDecay * sign(param) weightDecay * param注意 L1L2 模式下两个分量共用同一个mWeightDecay系数。需要区分 L1/L2 强度时可从源码结构推断应在调用侧自行扩展或分别维护系数。梯度阻断进阶SGD还提供setGradBlockName(std::vectorstd::string block)SGD.hpp用于指定不需要参与梯度计算即反向传播时被阻断的算子名称配合OpGrad::grad使用适合冻结部分子网络的需求。三、ADAM 优化器使用详解3.1 完整使用示例以下代码同样来自 docs/train/optim.md// 新建ADAM优化器 std::shared_ptrSGD solver(new ADAM); // 设置模型中需要优化的参数 solver-append(model-parameters()); // 设置ADAM的两个momentum设置weight decay solver-setMomentum(0.9f); solver-setMomentum2(0.99f); solver-setWeightDecay(0.0005f); // 设置正则化方法默认L2 solver-setRegularizationMethod(RegularizationMethod::L2); // 设置学习率 solver-setLearningRate(0.001); // 根据loss计算梯度并更新参数 solver-step(loss);文档中的std::shared_ptrSGD solver(new ADAM)利用了ADAM 继承自 SGD这一设计见 ADAM.hppADAM复用了 SGD 的学习率、一阶动量、权重衰减与正则化配置仅重写更新值计算逻辑你也可以显式写为std::shared_ptrADAM solver(new ADAM(model));。3.2 ADAM 特有参数与默认值对照 ADAM.hpp 的成员定义配置项Setter 接口默认值说明一阶动量 β₁setMomentum(float)继承自 SGD0一阶矩衰减系数文档示例取0.9二阶动量 β₂setMomentum2(float momentum2)0.999二阶矩衰减系数文档示例取0.99源码默认0.999数值稳定项 εsetEps(float eps)1e-8防止除零getEps()可查询学习率setLearningRate(float)继承自 SGD0.001f同 SGD3.3 ADAM 更新公式的源码级解读ADAM::onComputeUpdateValueADAM.cpp实现如下mHistory[param] beta1 * mHistory[param] (1 - beta1) * grad; // 一阶矩 m mHistory2[param] beta2 * mHistory2[param] (1 - beta2) * grad^2; // 二阶矩 v correction sqrt(1 - beta2^step) / (1 - beta1^step); // 偏差校正 updateValue lr * correction * m / (sqrt(v) eps);几点值得注意的工程细节偏差校正bias correctioncorrection使用当前优化步数step来自currentStep()即基类维护的mStep对一、二阶矩初始阶段的偏差进行补偿这是标准 Adam 论文的修正项。因此 ADAM 的训练步数管理是必须的框架在 ParameterOptimizer.hpp 提供currentStep()/setCurrentStep()。两套历史缓存ADAM 在 SGD 的mHistory一阶矩之外额外维护mHistory2二阶矩两者都在构造时为每个可训练参数初始化为全零张量ADAM.cpp。权重衰减在 ADAM 中的处理在onMakeParameterUpdateGraphByGrad路径中ADAM.cpp先执行gradWithDecay grad weightDecay * paramL2 形式再将该梯度同时送入一阶矩与二阶矩的更新属于 L2 正则化耦合进梯度 的实现方式。四、学习率调度器LrScheduler虽然优化器本身只需设置一个基础学习率但 MNN 训练框架同时提供静态学习率调度工具类LrScheduler定义见 LearningRateScheduler.hpp便于在训练过程中按迭代步数调整学习率// 多段衰减在指定步数处将学习率乘以对应倍数 static float multiStep(const float baseLr, const int step, std::vectorint stepIterations, std::vectorfloat lrMulti); // 逆时间衰减baseLr * pow((1 gamma * step), -power) static float inv(const float baseLr, const int step, const float gamma, const float power); // 指数衰减baseLr * pow(gamma, step) static float exp(const float baseLr, const int step, const float gamma);典型用法是在每个训练迭代中用当前step调用调度函数计算新的学习率再通过solver-setLearningRate(lr)写回优化器。其中multiStep适合分段常数衰减策略exp/inv适合平滑衰减策略。五、内置损失函数Loss5.1 接口一览文档列出的全部损失函数接口定义见 Loss.hpp实现见 Loss.cpp如下VARP _CrossEntropy(Express::VARP predicts, Express::VARP oneHotTargets); VARP _KLDivergence(Express::VARP predicts, Express::VARP oneHotTargets); VARP _MSE(Express::VARP predicts, Express::VARP oneHotTargets); VARP _MAE(Express::VARP predicts, Express::VARP oneHotTargets); VARP _Hinge(Express::VARP predicts, Express::VARP oneHotTargets); VARP _DistillLoss(Express::VARP studentLogits, Express::VARP teacherLogits, Express::VARP oneHotTargets, const float temperature, const float alpha);共同约束除_DistillLoss外其余损失均要求predicts与oneHotTargets为二维张量dim.size() 2且形状一致源码中通过MNN_ASSERT强制校验如 Loss.cpp。5.2 各损失函数的数学形式与实现要点交叉熵_CrossEntropyLoss.cpp-mean(sum(log(predicts) * oneHotTargets, dim1))。适用于分类任务要求predicts已通过 Softmax 归一化。KL 散度_KLDivergenceLoss.cppmean(sum(predicts * (log(predicts) - log(oneHotTargets)), dim1))。适用于分布匹配是蒸馏损失的核心组件。均方误差_MSELoss.cppmean(sum((predicts - oneHotTargets)^2, dim1))。适用于回归任务。平均绝对误差_MAELoss.cppmean(sum(|predicts - oneHotTargets|, dim1))。对离群点比 MSE 更鲁棒。Hinge_HingeLoss.cppmean(sum(max(0, 1 - predicts * oneHotTargets), dim1))。适用于最大间隔类任务如 SVM 式目标。蒸馏损失_DistillLossLoss.cpp组合教师网络与学生网络的软目标 KL 散度与真实标签交叉熵softTargets softmax(teacherLogits / temperature); studentPredict softmax(studentLogits / temperature); loss1 temperature^2 * KLDivergence(studentPredict, softTargets); // 蒸馏项 loss2 CrossEntropy(softmax(studentLogits), oneHotTargets); // 监督项 loss alpha * loss1 (1 - alpha) * loss2;参数语义temperature温度控制软标签的平滑程度alpha在蒸馏项与监督项之间取权重源码通过MNN_ASSERT(alpha 0 alpha 1)约束取值范围。此外该函数对NC4HW4布局的输入会自动_Convert到NCHW再计算Loss.cpp体现了与 MNN 内部张量布局体系的兼容性。5.3 自行设计 Loss文档指出目前支持的 Loss也可自行设计。由于所有损失函数本质上都是基于 MNN 表达式系统Express的算子组合_ReduceMean、_Log、_Square、_Softmax、_Scalar等你完全可以复用 Loss.cpp 的模式用表达式算子拼装自定义损失例如构造VARP myLoss _ReduceMean(...)得到标量 loss 后直接作为solver-step(myLoss)的入参。损失函数最终输出的必须是标量dim.size() 0这是优化器反向求导的输入前提。六、端到端使用要点与验证6.1 完整调用链回顾一个典型的 MNN 训练迭代可以归纳为// 1. 前向得到模型输出 auto predicts model-forward(inputs); // 2. 计算损失标量 auto loss _CrossEntropy(predicts, oneHotTargets); // 3. 优化器一步更新内部完成 反向梯度 → 正则化 → 动量/二阶矩 → 参数写回 solver-step(loss);step的具体实现为ParameterOptimizer::step(Express::VARP loss)声明于 ParameterOptimizer.hpp实现于 ParameterOptimizer.cpp是文档示例中solver-step(loss)这一行的落点。6.2 可运行的工程参照仓库内提供了完整可编译的训练示例作为参照MnistUtils.cppMNIST 训练/测试工具展示了数据加载、模型构建、loss 计算与优化器配合的完整范式MobilenetV2Utils.cppMobileNetV2 训练工具演示了_CrossEntropy等损失与 SGD/ADAM 的搭配quanByMSE.cpp基于 MSE 损失的量化校准示例体现了以 MSE 作为优化目标的实际应用。如需验证优化器行为可关注仓库 test/grad 与 test/expr 目录下的测试用例它们对梯度与表达式算子即优化器底层依赖的正确性进行了覆盖。6.3 使用注意事项ADAM 的setMomentum2默认值为0.999而文档示例取0.99实际训练时应按任务调参setMomentum的默认值是0使用 ADAM 时必须显式设置 β₁。权重衰减与正则化方法绑定mWeightDecay的实际作用方式由RegularizationMethod决定L1 作用于参数符号、L2 作用于参数本身见 SGD.cpp。损失必须为标量所有内置损失末尾都有_ReduceMean(..., {})将结果归约为标量自定义损失也需保持该约束。step(loss)会推进内部步数该步数同时驱动 ADAM 的偏差校正与学习率调度因此不要绕过step手动混用更新逻辑。结语本文以 docs/train/optim.md 为主线结合 tools/train/source/optimizer 下的源码完整覆盖了 MNN 训练框架中 SGD with Momentum 与 ADAM 优化器的配置方法、更新公式、默认参数与正则化机制以及六种内置损失函数的数学形式与实现细节。无论你是要在端侧设备上微调分类模型还是借助蒸馏损失压缩教师网络均可参照文中代码直接落地并通过solver-step(loss)一键完成参数更新。【免费下载链接】MNNMNN: A blazing-fast, lightweight inference engine battle-tested by Alibaba, powering high-performance on-device LLMs and Edge AI.项目地址: https://gitcode.com/GitHub_Trending/mn/MNN创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考