MNN PyMNN loss 模块实战:五种损失函数的 Python API、数学实现与训练集成
【免费下载链接】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 的 PyMNN 接口中,nn.loss是模型训练链路的核心组件之一,提供cross_entropy、kl、mse、mae、hinge五个常用损失函数。本文以 docs/pymnn/loss.md 的 API 说明为主体,结合 Loss.cpp 中的底层实现,讲清每个损失的数学定义、输入约束、返回值语义,以及如何在 PyMNN 训练循环中与数据加载器、优化器配合使用,最终形成一套可直接复制到自有训练脚本中的损失计算方案。
loss 模块概览
loss 模块是 MNN 模型训练使用的模块,提供了多个损失函数。从 Python 侧看,它属于MNN.nn命名空间下的loss子模块(module loss),所有函数均为二元操作:接收预测值与 one-hot 标签两个Var,返回一个标量Var形式的损失值。
模块的 Python 绑定非常简洁,见 pymnn/src/loss.h:
// loss Module Start def_binary(Loss, cross_entropy, _CrossEntropy, kl, _KLDivergence, mse, _MSE, mae, _MAE, hinge, _Hinge ) static PyMethodDef PyMNNLoss_methods[] = { register_methods(Loss, cross_entropy, "cross_entropy loss", kl, "kl loss", mse, "mse loss", mae, "mae loss", hinge, "hinge loss" ) }; // loss Module Enddef_binary宏(定义于 pymnn/src/util.h 的def_binary处)把 Python 层的函数名批量映射到 Express 层对应的二元算子创建函数:_CrossEntropy、_KLDivergence等。因此每个 Python 损失函数本质上都是对底层 Express 计算图的一次“声明式拼接”,而非独立的手写内核。
五个函数的通用接口约定完全一致:
| 参数 | 类型 | dtype | shape | 含义 |
|---|---|---|---|---|
predicts | Var | float | (batch_size, num_classes) | 输出层的预测值 |
onehot_targets | Var | float | (batch_size, num_classes) | one-hot 编码的标签 |
返回值统一为Var,即标量损失,可直接传给优化器的step接口或参与更复杂的图运算。
五种损失函数的定义与实现
下面逐个对照官方文档的示例与 Loss.cpp 中的 Express 表达式实现。所有实现都位于MNN::Train命名空间,声明见 Loss.hpp。
cross_entropy:交叉熵损失
计算交叉熵损失,是分类任务最常用的损失。文档示例:
>>> predict = np.random.random([2,3]) >>> onehot = np.array([[1., 0., 0.], [0., 1., 0.]]) >>> nn.loss.cross_entropy(predict, onehot) array(4.9752955, dtype=float32)实现如下(Loss.cpp 中_CrossEntropy):
Express::VARP _CrossEntropy(Express::VARP predicts, Express::VARP oneHotTargets) { MNN_ASSERT(predicts->getInfo()->dim.size() == 2); MNN_ASSERT(predicts->getInfo()->dim == oneHotTargets->getInfo()->dim); auto loss = _Negative(_ReduceMean(_ReduceSum(_Log(predicts) * oneHotTargets, {1}), {})); return loss; }即数学形式 $\mathcal{L} = -\frac{1}{B}\sum_{b}\sum_{c} y_{bc}\log \hat{p}_{bc}$:对预测值取对数、与 one-hot 标签逐元素相乘(只保留真实类别项),沿类别维求和、再对 batch 维求均值取负。由此可以得出两条使用约束:
predicts必须是概率分布(每个元素大于 0 且行内和为 1),通常在模型末尾接softmax,示例脚本中网络forward最后一步即为x = F.softmax(x, 1);- 两个输入必须都是二维张量且形状相同,断言失败会在 debug 构建下触发。
kl:KL 散度损失
计算 KL 损失(相对标签分布的 KL 散度)。文档示例:
>>> predict = np.random.random([2,3]) >>> onehot = np.array([[1., 0., 0.], [0., 1., 0.]]) >>> nn.loss.kl(predict, onehot) array(inf, dtype=float32)实现为:
Express::VARP _KLDivergence(Express::VARP predicts, Express::VARP oneHotTargets) { MNN_ASSERT(predicts->getInfo()->dim.size() == 2); MNN_ASSERT(predicts->getInfo()->dim == oneHotTargets->getInfo()->dim); auto loss = _ReduceMean(_ReduceSum(_Multiply(predicts, _Log(predicts) - _Log(oneHotTargets)), {1}), {}); return loss; }数学形式 $\mathcal{L} = \frac{1}{B}\sum_{b}\sum_{c} \hat{p}{bc}\left(\log \hat{p}{bc} - \log y_{bc}\right)$。注意示例返回inf并非异常:one-hot 标签中含 0,而 $\log 0 = -\infty$,当预测值在错误类别上非零时对应项发散。因此从源码结构看,该函数更适合作为软标签(soft labels)或知识蒸馏场景使用——例如 MNN 训练工具中的蒸馏损失_DistillLoss就是把教师 logits 过 softmax 得到软目标后计算 KL 散度(见 docs/train/distl.md)。直接用严格 one-hot 标签调用会得到文档示例中的inf,这是预期行为而非 bug。
mse:均方误差
计算 MSE 损失。文档示例:
>>> predict = np.random.random([2,3]) >>> onehot = np.array([[1., 0., 0.], [0., 1., 0.]]) >>> nn.loss.mse(predict, onehot) array(1.8694793, dtype=float32)实现为:
Express::VARP _MSE(Express::VARP predicts, Express::VARP oneHotTargets) { MNN_ASSERT(predicts->getInfo()->dim.size() == 2); MNN_ASSERT(predicts->getInfo()->dim == oneHotTargets->getInfo()->dim); auto loss = _ReduceMean(_ReduceSum(_Square(predicts - oneHotTargets), {1}), {}); return loss; }即 $\mathcal{L} = \frac{1}{B}\sum_{b}\sum_{c}(\hat{p}{bc} - y{bc})^2$。对 one-hot 回归化训练(如多标签分类的 sigmoid 输出)或回归任务适用。
mae:平均绝对误差
计算 MAE 损失。文档示例:
>>> predict = np.random.random([2,3]) >>> onehot = np.array([[1., 0., 0.], [0., 1., 0.]]) >>> nn.loss.mae(predict, onehot) array(2.1805272, dtype=float32)实现为:
Express::VARP _MAE(Express::VARP predicts, Express::VARP oneHotTargets) { MNN_ASSERT(predicts->getInfo()->dim.size() == 2); MNN_ASSERT(predicts->getInfo()->dim == oneHotTargets->getInfo()->dim); auto loss = _ReduceMean(_ReduceSum(_Abs(predicts - oneHotTargets), {1}), {}); return loss; }即 $\mathcal{L} = \frac{1}{B}\sum_{b}\sum_{c}|\hat{p}{bc} - y{bc}|$。相比 MSE 对离群点更鲁棒,梯度幅值恒定,适合标签噪声较多的场景。
hinge:铰链损失
计算 Hinge 损失。文档示例:
>>> predict = np.random.random([2,3]) >>> onehot = np.array([[1., 0., 0.], [0., 1., 0.]]) >>> nn.loss.hinge(predict, onehot) array(2.791432, dtype=float32)实现为:
Express::VARP _Hinge(Express::VARP predicts, Express::VARP oneHotTargets) { MNN_ASSERT(predicts->getInfo()->dim.size() == 2); MNN_ASSERT(predicts->getInfo()->dim == oneHotTargets->getInfo()->dim); auto loss = _ReduceMean(_ReduceSum(_Maximum(_Const(0.), _Const(1.) - predicts * oneHotTargets), {1}), {}); return loss; }即 $\mathcal{L} = \frac{1}{B}\sum_{b}\sum_{c}\max(0, 1 - \hat{p}{bc},y{bc})$,是 SVM 的合页损失在多类 one-hot 上的逐类推广:只要预测值与标签乘积小于 1 就产生惩罚,鼓励预测尽量推过 1 的间隔。
在训练循环中使用 loss:官方示例解析
loss 函数返回的标量Var是接入优化器的入口。pymnn/examples/MNNTrain/mnist/train_mnist.py 展示了最完整的集成方式(LeNet-5 训练 MNIST):
nn = MNN.nn F = MNN.expr # open lazy evaluation for train F.lazy_eval(True) class Net(nn.Module): """construct a lenet 5 model""" def __init__(self): super(Net, self).__init__() self.conv1 = nn.conv(1, 20, [5, 5]) self.conv2 = nn.conv(20, 50, [5, 5]) self.fc1 = nn.linear(800, 500) self.fc2 = nn.linear(500, 10) def forward(self, x): x = F.relu(self.conv1(x)) x = F.max_pool(x, [2, 2], [2, 2]) x = F.relu(self.conv2(x)) x = F.max_pool(x, [2, 2], [2, 2]) # MNN use NC4HW4 format for convs, so we need to convert it to NCHW before entering other ops x = F.convert(x, F.NCHW) x = F.reshape(x, [0, -1]) x = F.relu(self.fc1(x)) x = self.fc2(x) x = F.softmax(x, 1) # loss 要求 predicts 为概率分布 return x训练函数中的损失计算与优化:
def train_func(net, train_dataloader, opt): net.train(True) train_dataloader.reset() for i in range(train_dataloader.iter_number): example = train_dataloader.next() data = example[0][0] # 输入 Var label = example[1][0] # 标签 Var(int) predict = net.forward(data) target = F.one_hot(F.cast(label, F.int), 10, 1, 0) # one-hot 编码 loss = nn.loss.cross_entropy(predict, target) # 标量 Var opt.step(loss) # 反传 + 参数更新 if i % 100 == 0: print("train loss: ", loss.read())几个关键点值得注意:
- 标签预处理:数据加载器给出的 label 是整型
Var,需先用F.one_hot(F.cast(label, F.int), num_classes, 1, 0)转为(batch, num_classes)的浮点 one-hot,再送入损失函数,这与文档中onehot_targets: Var, dtype=float的签名一致; - 惰性求值:脚本开头
F.lazy_eval(True)打开惰性求值,loss.read()才真正触发前向计算并取值,这与 MNN Express 的图化执行模型一致; - 优化器闭环:
opt = MNN.optim.SGD(model, 0.01, 0.9, 0.0005)创建 SGD 优化器后,opt.step(loss)完成对损失的反传与更新,loss 模块的输出因此天然嵌入了梯度链路; - 布局转换:卷积类算子使用 NC4HW4 布局,进入全连接层前要
F.convert(x, F.NCHW),否则形状对不上,损失函数中dim断言也会失败。
迁移学习的场景可参考 pymnn/examples/MNNTrain/mobilenet_finetune/mobilenet_transfer.py,其训练循环与 MNIST 完全同构:
predict = net.forward(F.convert(data, F.NC4HW4)) # 特征提取器输入为 NC4HW4 target = F.one_hot(F.cast(label, F.int), num_classes, 1, 0) loss = nn.loss.cross_entropy(predict, target) opt.step(loss)区别只是分类头类别数num_classes由自定义数据集决定,以及输入数据在进入 MobileNetV2 特征提取器前要先F.convert(data, F.NC4HW4)。
与 MNN 训练工具链的对应关系
PyMNN 的 loss 模块并不是孤立的 API。从源码结构看,它与 C++ 训练工具共享同一套损失实现:
- C++ 侧声明集中在 tools/train/source/optimizer/Loss.hpp,除上述五个函数外还包含蒸馏损失
_DistillLoss(温度缩放下的 KL + 交叉熵加权组合,见 docs/train/distl.md); - docs/train/optim.md 的 Loss 一节列出了同样的函数签名,说明 PyMNN
nn.loss.*与MNN::Train::_*是同一实现的两层封装; tools/train下的 MNIST、MobileNetV2 训练演示(tools/train/source/demo/MnistUtils.cpp、MobilenetV2Utils.cpp)也采用相同的 “one-hot 标签 +_CrossEntropy” 组合,与 Python 示例一一对应;- 推理端的 ppl_eval 工具(transformers/llm/engine/tools/ppl_eval.cpp)内部还实现了一个带
ignore_index的扩展版_CrossEntropy,用_OneHot构造 mask 忽略指定标签,属于 LLM 评测场景的定制化实现,与训练 API 的五函数集合互不干扰。
使用注意事项小结
结合文档示例与源码实现,使用nn.loss时需要注意:
- 形状约束:
predicts与onehot_targets必须同为(batch_size, num_classes)二维张量,源码中有明确的dim断言; - 数值约束:
cross_entropy对predicts取对数,应传 softmax 后的概率;kl对onehot_targets取对数,严格 one-hot 标签会产生inf,软标签场景才合理; - 返回类型:五个函数都返回标量
Var,可直接传给MNN.optim.*的step接口,也可继续参与图运算(如加权求和组合多任务损失); - 自定义扩展:文档指出“loss 模块是模型训练使用的模块,提供了多个损失函数”,若五个内置损失不满足需求,可以直接用
MNN.expr的算子(_Log、_Square、_Abs、_ReduceSum等)按 Loss.cpp 同样的模式拼装出自己的损失表达式,训练工具链也鼓励“自行设计”。
综上,PyMNN 的 loss 模块以五个语义清晰的二元损失函数覆盖了分类、回归与蒸馏等主流训练需求,其实现均为 Express 算子的组合,行为可预测、易于扩展;配合官方 MNIST 与 MobileNet 微调示例中的 “forward → one_hot → loss → opt.step” 范式,即可在 MNN 体系内完整地搭建端到端训练与微调流程。
【免费下载链接】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),仅供参考