news 2026/9/14 19:15:32

MNN PyMNN loss 模块实战:五种损失函数的 Python API、数学实现与训练集成

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
MNN PyMNN loss 模块实战:五种损失函数的 Python API、数学实现与训练集成

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_entropyklmsemaehinge五个常用损失函数。本文以 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 End

def_binary宏(定义于 pymnn/src/util.h 的def_binary处)把 Python 层的函数名批量映射到 Express 层对应的二元算子创建函数:_CrossEntropy_KLDivergence等。因此每个 Python 损失函数本质上都是对底层 Express 计算图的一次“声明式拼接”,而非独立的手写内核。

五个函数的通用接口约定完全一致:

参数类型dtypeshape含义
predictsVarfloat(batch_size, num_classes)输出层的预测值
onehot_targetsVarfloat(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())

几个关键点值得注意:

  1. 标签预处理:数据加载器给出的 label 是整型Var,需先用F.one_hot(F.cast(label, F.int), num_classes, 1, 0)转为(batch, num_classes)的浮点 one-hot,再送入损失函数,这与文档中onehot_targets: Var, dtype=float的签名一致;
  2. 惰性求值:脚本开头F.lazy_eval(True)打开惰性求值,loss.read()才真正触发前向计算并取值,这与 MNN Express 的图化执行模型一致;
  3. 优化器闭环opt = MNN.optim.SGD(model, 0.01, 0.9, 0.0005)创建 SGD 优化器后,opt.step(loss)完成对损失的反传与更新,loss 模块的输出因此天然嵌入了梯度链路;
  4. 布局转换:卷积类算子使用 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 一节列出了同样的函数签名,说明 PyMNNnn.loss.*MNN::Train::_*是同一实现的两层封装;
  • tools/train下的 MNIST、MobileNetV2 训练演示(tools/train/source/demo/MnistUtils.cppMobilenetV2Utils.cpp)也采用相同的 “one-hot 标签 +_CrossEntropy” 组合,与 Python 示例一一对应;
  • 推理端的 ppl_eval 工具(transformers/llm/engine/tools/ppl_eval.cpp)内部还实现了一个带ignore_index的扩展版_CrossEntropy,用_OneHot构造 mask 忽略指定标签,属于 LLM 评测场景的定制化实现,与训练 API 的五函数集合互不干扰。

使用注意事项小结

结合文档示例与源码实现,使用nn.loss时需要注意:

  1. 形状约束predictsonehot_targets必须同为(batch_size, num_classes)二维张量,源码中有明确的dim断言;
  2. 数值约束cross_entropypredicts取对数,应传 softmax 后的概率;klonehot_targets取对数,严格 one-hot 标签会产生inf,软标签场景才合理;
  3. 返回类型:五个函数都返回标量Var,可直接传给MNN.optim.*step接口,也可继续参与图运算(如加权求和组合多任务损失);
  4. 自定义扩展:文档指出“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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/14 19:10:29

从“文献焦虑”到“学术拼图”:书匠策AI文献综述功能拆解

官网:www.shujiangce.com | 微信 公众号 :书匠策AI 写文献综述最诡异的体验是什么? 不是读不懂文献,而是读得越多,脑子越乱。三十篇PDF在文件夹里安静地躺着,每一篇单独看都明白,但当你试图…

作者头像 李华
网站建设 2026/9/14 19:08:05

LangChain Sequential Chain金融场景实战与优化

1. 当LangChain的Sequential Chain在凌晨三点报错时凌晨三点,屏幕的蓝光刺得眼睛生疼。我盯着控制台里那行鲜红的错误提示,第17次尝试修复这个该死的Sequential Chain。咖啡已经喝到第三杯,但大脑依然像被灌了铅——这就是AI工程师的日常&…

作者头像 李华
网站建设 2026/9/14 19:07:52

风管展开下料软件:智能算法提升制造效率与材料利用率

1. 风管展开下料软件的核心价值解析在通风管道制造领域,传统手工放样方式存在三大痛点:一是展开图绘制效率低下,复杂管件需要数小时计算;二是材料利用率普遍低于75%,造成严重浪费;三是人工排料易出错导致返…

作者头像 李华