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_ptr<Express::Module>,通过module()->parameters()拿到模型全部参数,再筛选出"可训练"参数集合mTrainable(详见 ParameterOptimizer.cpp)。文档示例中的solver->append(model->parameters())即完成"设置模型中需要优化的参数"这一步。 step(loss)是统一的更新入口:一次调用完成"前向求 loss → 反向求梯度 → 正则化 → 动量/二阶矩修正 → 更新参数"的全链路。- 正则化方式枚举:
enum RegularizationMethod { L1, L2, L1L2 },即支持 L1、L2、L1+L2 三种,默认 L2。 - 训练步数管理:
currentStep()/setCurrentStep(int)维护全局优化步数mStep,ADAM 的偏差校正(bias correction)依赖该步数。
从SGD::onGetNextParameter的实现(SGD.cpp)可以还原出step(loss)内部的典型流程:
OpGrad::grad(loss, trainable(), ...)基于 loss 对全部可训练参数求梯度;- 对每个参数执行
regularizeParameters(param, grad)叠加权重衰减(正则化)项; - 调用
onComputeUpdateValue(param, grad)计算带动量/二阶矩的更新量; - 执行
param - updateValue得到新参数值,返回给上层写回。
此外,框架还提供了工厂方法用于一行创建优化器(同样定义在 ParameterOptimizer.hpp):
static ParameterOptimizer* createSGD(std::shared_ptr<Express::Module> module, float lr, float momentum, float weightDecay, RegularizationMethod method); static ParameterOptimizer* createADAM(std::shared_ptr<Express::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_ptr<SGD> 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_ptr<SGD> solver(new SGD(model));,将模型直接传入构造函数(参见 SGD.hpp 中SGD(std::shared_ptr<Express::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/L1L2 |
2.3 更新公式的源码级解读
SGD 的核心更新逻辑在SGD::onComputeUpdateValue(SGD.cpp):
mHistory[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。
正则化(权重衰减)如何生效:regularizeParameters(SGD.cpp)在计算更新量之前,先向原始梯度叠加正则项:
L1:grad + weightDecay * sign(param)L2:grad + weightDecay * paramL1L2:grad + weightDecay * sign(param) + weightDecay * param
注意 L1L2 模式下两个分量共用同一个mWeightDecay系数。需要区分 L1/L2 强度时,可从源码结构推断应在调用侧自行扩展或分别维护系数。
梯度阻断(进阶):SGD还提供setGradBlockName(std::vector<std::string> block)(SGD.hpp),用于指定不需要参与梯度计算(即反向传播时被阻断)的算子名称,配合OpGrad::grad使用,适合冻结部分子网络的需求。
三、ADAM 优化器使用详解
3.1 完整使用示例
以下代码同样来自 docs/train/optim.md:
// 新建ADAM优化器 std::shared_ptr<SGD> 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_ptr<SGD> solver(new ADAM)利用了ADAM 继承自 SGD这一设计(见 ADAM.hpp),ADAM复用了 SGD 的学习率、一阶动量、权重衰减与正则化配置,仅重写更新值计算逻辑;你也可以显式写为std::shared_ptr<ADAM> solver(new ADAM(model));。
3.2 ADAM 特有参数与默认值
对照 ADAM.hpp 的成员定义:
| 配置项 | Setter 接口 | 默认值 | 说明 |
|---|---|---|---|
| 一阶动量 β₁ | setMomentum(float)(继承自 SGD) | 0 | 一阶矩衰减系数,文档示例取0.9 |
| 二阶动量 β₂ | setMomentum2(float momentum2) | 0.999 | 二阶矩衰减系数,文档示例取0.99,源码默认0.999 |
| 数值稳定项 ε | setEps(float eps) | 1e-8 | 防止除零;getEps()可查询 |
| 学习率 | setLearningRate(float)(继承自 SGD) | 0.001f | 同 SGD |
3.3 ADAM 更新公式的源码级解读
ADAM::onComputeUpdateValue(ADAM.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 correction):
correction使用当前优化步数step(来自currentStep(),即基类维护的mStep)对一、二阶矩初始阶段的偏差进行补偿,这是标准 Adam 论文的修正项。因此 ADAM 的训练步数管理是必须的,框架在 ParameterOptimizer.hpp 提供currentStep()/setCurrentStep()。 - 两套历史缓存:ADAM 在 SGD 的
mHistory(一阶矩)之外,额外维护mHistory2(二阶矩),两者都在构造时为每个可训练参数初始化为全零张量(ADAM.cpp)。 - 权重衰减在 ADAM 中的处理:在
onMakeParameterUpdateGraphByGrad路径中(ADAM.cpp),先执行gradWithDecay = grad + weightDecay * param(L2 形式),再将该梯度同时送入一阶矩与二阶矩的更新,属于 "L2 正则化耦合进梯度" 的实现方式。
四、学习率调度器(LrScheduler)
虽然优化器本身只需设置一个基础学习率,但 MNN 训练框架同时提供静态学习率调度工具类LrScheduler(定义见 LearningRateScheduler.hpp),便于在训练过程中按迭代步数调整学习率:
// 多段衰减:在指定步数处将学习率乘以对应倍数 static float multiStep(const float baseLr, const int step, std::vector<int> stepIterations, std::vector<float> 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适合平滑衰减策略。
五、内置损失函数(Loss)
5.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 各损失函数的数学形式与实现要点
- 交叉熵
_CrossEntropy(Loss.cpp):-mean(sum(log(predicts) * oneHotTargets, dim=1))。适用于分类任务,要求predicts已通过 Softmax 归一化。 - KL 散度
_KLDivergence(Loss.cpp):mean(sum(predicts * (log(predicts) - log(oneHotTargets)), dim=1))。适用于分布匹配,是蒸馏损失的核心组件。 - 均方误差
_MSE(Loss.cpp):mean(sum((predicts - oneHotTargets)^2, dim=1))。适用于回归任务。 - 平均绝对误差
_MAE(Loss.cpp):mean(sum(|predicts - oneHotTargets|, dim=1))。对离群点比 MSE 更鲁棒。 - Hinge
_Hinge(Loss.cpp):mean(sum(max(0, 1 - predicts * oneHotTargets), dim=1))。适用于最大间隔类任务(如 SVM 式目标)。 - 蒸馏损失
_DistillLoss(Loss.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.cpp:MNIST 训练/测试工具,展示了数据加载、模型构建、loss 计算与优化器配合的完整范式;
- MobilenetV2Utils.cpp:MobileNetV2 训练工具,演示了
_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),仅供参考