MNN 训练优化器使用指南:SGD、ADAM、损失函数与学习率调度
2026/9/14 4:51:31 网站建设 项目流程

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)内部的典型流程:

  1. OpGrad::grad(loss, trainable(), ...)基于 loss 对全部可训练参数求梯度;
  2. 对每个参数执行regularizeParameters(param, grad)叠加权重衰减(正则化)项;
  3. 调用onComputeUpdateValue(param, grad)计算带动量/二阶矩的更新量;
  4. 执行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)在计算更新量之前,先向原始梯度叠加正则项:

  • L1grad + weightDecay * sign(param)
  • L2grad + weightDecay * param
  • L1L2grad + 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外,其余损失均要求predictsoneHotTargets二维张量(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布局的输入会自动_ConvertNCHW再计算(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 使用注意事项

  1. ADAM 的setMomentum2默认值为0.999,而文档示例取0.99,实际训练时应按任务调参;setMomentum的默认值是0,使用 ADAM 时必须显式设置 β₁。
  2. 权重衰减与正则化方法绑定mWeightDecay的实际作用方式由RegularizationMethod决定,L1 作用于参数符号、L2 作用于参数本身(见 SGD.cpp)。
  3. 损失必须为标量:所有内置损失末尾都有_ReduceMean(..., {})将结果归约为标量,自定义损失也需保持该约束。
  4. 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询