当前位置: 首页 > news >正文

别再死记硬背NLL公式了!用PyTorch手把手带你复现一个分类任务(附完整代码)

从零实现NLL损失函数:PyTorch实战图像分类任务

刚接触机器学习的同学一定对"负对数似然损失"这个术语不陌生,但真正理解它如何在实际代码中发挥作用的人却不多。今天我们不谈复杂的数学推导,而是直接动手用PyTorch实现一个完整的分类任务,让你亲眼看到NLLLoss是如何工作的。

很多教程一上来就抛出NLL的数学公式,让人望而生畏。其实,理解一个概念最好的方式就是亲手实现它。我们将从数据加载开始,一步步构建模型、定义损失函数,直到完成训练循环。在这个过程中,你会遇到几个常见的"坑",比如忘记添加LogSoftmax层,或者混淆了NLLLoss和CrossEntropyLoss的区别——别担心,我都会带你一一解决。

1. 环境准备与数据加载

首先确保你已经安装了最新版的PyTorch。如果你使用conda环境,可以通过以下命令安装:

conda install pytorch torchvision -c pytorch

我们将使用经典的MNIST手写数字数据集作为示例。这个数据集包含60,000张28x28像素的手写数字图像,非常适合用来理解分类任务的基本原理。

import torch from torchvision import datasets, transforms # 定义数据转换 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 加载训练集和测试集 train_dataset = datasets.MNIST('./data', train=True, download=True, transform=transform) test_dataset = datasets.MNIST('./data', train=False, transform=transform) # 创建数据加载器 train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=64, shuffle=True) test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=1000, shuffle=True)

提示:MNIST数据集中的图像已经被标准化到0-1范围,我们进一步使用均值0.1307和标准差0.3081进行归一化,这有助于模型更快收敛。

2. 构建神经网络模型

接下来,我们定义一个简单的卷积神经网络(CNN)来处理MNIST图像。虽然模型结构不是本文的重点,但理解各层的作用对调试NLLLoss很有帮助。

import torch.nn as nn import torch.nn.functional as F class SimpleCNN(nn.Module): def __init__(self): super(SimpleCNN, self).__init__() self.conv1 = nn.Conv2d(1, 10, kernel_size=5) self.conv2 = nn.Conv2d(10, 20, kernel_size=5) self.fc1 = nn.Linear(320, 50) self.fc2 = nn.Linear(50, 10) def forward(self, x): x = F.relu(F.max_pool2d(self.conv1(x), 2)) x = F.relu(F.max_pool2d(self.conv2(x), 2)) x = x.view(-1, 320) x = F.relu(self.fc1(x)) x = self.fc2(x) return F.log_softmax(x, dim=1)

注意模型最后一层的输出:我们使用了F.log_softmax而不是普通的softmax。这是使用NLLLoss的关键前提——NLLLoss期望接收的是对数概率(log probabilities),而不是原始概率。

3. 理解NLLLoss的工作原理

现在来到核心部分:负对数似然损失函数。在PyTorch中,它由nn.NLLLoss类实现。让我们先看看它的数学本质:

假设我们的模型对某个样本的输出概率分布为[0.1, 0.8, 0.1],真实标签是1(第二类)。那么:

  1. 取正确类别的概率:0.8
  2. 计算其对数:log(0.8) ≈ -0.2231
  3. 取负值:0.2231

这就是NLLLoss的计算过程。当正确类别的预测概率越高,损失值就越小。

在代码中实现这一点非常简单:

model = SimpleCNN() criterion = nn.NLLLoss() optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.5)

注意:常见的错误是忘记在模型最后添加LogSoftmax层,或者错误地使用了普通的softmax。NLLLoss必须与LogSoftmax配合使用,如果使用普通softmax会导致计算错误。

4. 训练循环与结果分析

让我们把前面准备好的组件组合起来,实现完整的训练过程:

def train(epoch): model.train() for batch_idx, (data, target) in enumerate(train_loader): optimizer.zero_grad() output = model(data) loss = criterion(output, target) loss.backward() optimizer.step() if batch_idx % 100 == 0: print(f'Train Epoch: {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)} ' f'({100. * batch_idx / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}') def test(): model.eval() test_loss = 0 correct = 0 with torch.no_grad(): for data, target in test_loader: output = model(data) test_loss += criterion(output, target).item() pred = output.argmax(dim=1, keepdim=True) correct += pred.eq(target.view_as(pred)).sum().item() test_loss /= len(test_loader.dataset) print(f'\nTest set: Average loss: {test_loss:.4f}, ' f'Accuracy: {correct}/{len(test_loader.dataset)} ' f'({100. * correct / len(test_loader.dataset):.0f}%)\n') for epoch in range(1, 10): train(epoch) test()

运行这段代码,你会看到类似下面的输出:

Train Epoch: 1 [0/60000 (0%)] Loss: 2.312423 Train Epoch: 1 [6400/60000 (11%)] Loss: 0.876543 ... Test set: Average loss: 0.0023, Accuracy: 9234/10000 (92%)

5. NLLLoss与CrossEntropyLoss的关系

很多初学者会困惑:为什么PyTorch同时提供了NLLLoss和CrossEntropyLoss?它们之间有什么区别?

实际上,CrossEntropyLoss = LogSoftmax + NLLLoss。也就是说:

# 这两种方式是等价的: loss1 = nn.CrossEntropyLoss()(model_output, target) # 等价于 log_probs = F.log_softmax(model_output, dim=1) loss2 = nn.NLLLoss()(log_probs, target)

那么为什么PyTorch要提供两种实现呢?主要有两个原因:

  1. 灵活性:有时你可能需要在LogSoftmax和NLLLoss之间插入其他操作
  2. 历史原因:这两个概念在数学上是分开的,分开实现更符合理论定义

在实际应用中,如果你只是需要一个标准的分类损失函数,直接使用CrossEntropyLoss更为方便。但理解NLLLoss的工作原理对于调试模型和实现自定义损失函数非常有帮助。

6. 常见问题与调试技巧

在使用NLLLoss时,你可能会遇到以下几个典型问题:

问题1:损失值出现负数

这通常意味着你的模型输出没有经过LogSoftmax处理。NLLLoss期望输入是对数概率,如果直接传入原始分数,可能会计算出无意义的结果。

问题2:损失值下降但准确率不提高

检查你的LogSoftmax是否应用在了正确的维度上。对于分类任务,通常应该在最后一个维度(dim=1)上应用。

问题3:损失值突然变成NaN

这可能是由于数值不稳定导致的。尝试:

  1. 减小学习率
  2. 添加梯度裁剪
  3. 检查数据中是否有异常值
# 梯度裁剪示例 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

7. 扩展应用:自定义NLLLoss

理解了基本原理后,我们可以尝试实现自己的NLLLoss。这不仅能加深理解,还能根据需要添加特殊功能:

class MyNLLLoss(nn.Module): def __init__(self): super(MyNLLLoss, self).__init__() def forward(self, input, target): # input是log probabilities # target是类别索引 loss = -input[range(target.shape[0]), target].mean() return loss

这个自定义实现与PyTorch内置的NLLLoss功能相同,但代码更加透明。你可以在此基础上添加权重、忽略特定类别等功能。

8. 实际项目中的最佳实践

在真实项目中使用NLLLoss时,有几个经验值得分享:

  1. 始终验证输入形状:确保你的log probabilities和targets的形状匹配

    # log_probs形状应为[N, C],targets形状应为[N] assert log_probs.shape[0] == targets.shape[0] assert log_probs.shape[1] == num_classes
  2. 考虑类别不平衡:如果某些类别样本很少,可以使用weight参数

    # 假设类别0的样本是类别1的2倍 weight = torch.tensor([1.0, 2.0]) criterion = nn.NLLLoss(weight=weight)
  3. 与LogSoftmax的配合:确保只在训练时使用LogSoftmax,推理时直接取argmax即可

  4. 学习率调整:NLLLoss对学习率比较敏感,建议使用学习率调度器

    scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)

在图像分类任务中,经过适当调参,使用NLLLoss的简单CNN模型在MNIST上可以达到98%以上的准确率。这证明了即使不依赖复杂的数学推导,通过实践也能很好地理解和应用这一重要概念。

http://www.jsqmd.com/news/567059/

相关文章:

  • 计算机毕业设计:Python 二手车价格预测与特征分析系统 Flask框架 requests爬虫 可视化 数据分析 大数据 机器学习 大模型(建议收藏)✅
  • Node.js版本管理神器NVM:从安装到实战的保姆级教程(Mac版)
  • 用Python和Keras实战LSTM-AutoEncoder:手把手教你搭建室内空气质量异常检测模型
  • 终极指南:在Windows上快速免费安装安卓应用的完整解决方案
  • PyAEDT实战指南:高效电磁仿真的Python自动化方案
  • Ubuntu 22.04 编译内核踩坑记:手把手教你用gcc-9解决 `yylloc` 重定义错误
  • 如何解析B站视频:零门槛的bilibili-parse使用指南
  • 北海穷游必吃的美食攻略
  • 51单片机实战:用快马平台生成智能农业监测系统,从需求到完整代码
  • AI 牛马项圈公司新估值 20 亿美元,亚秒级实时监控;ProactiveVideoQA:首个视频多模态模型主动交互基准丨日报
  • 野火STM32F429与LVGL实战:从CubeMX配置到GUI移植全解析
  • AI智能体在测试自动化中的作用
  • 告别CNN!用Vision Transformer(ViT)和CellViT搞定病理切片细胞分割,附完整代码与避坑指南
  • 【超详细教程】手把手教你部署OpenClaw,两步解锁龙虾AI助理!
  • sqlmap工具超详细使用教程:从零基础到实战攻防(附避坑指南)
  • QMCDecode:突破QQ音乐格式枷锁,高效实现音乐自由播放
  • 什么是BGP协议
  • 内网UOS服务器装不了软件?手把手教你用ISO搭建本地YUM源,离线也能秒装包!
  • 别再用HAL_Delay了!STM32F407用CubeMX配置GPIO实现精准LED闪烁(附源码)
  • 智能手表与手机数据打架?用HHAR数据集实战多设备传感器融合与校准
  • 巴法云MQTT实战:避开ESP8266连接与App Inventor开发的5个常见坑
  • 虚拟人格入殓师:为废弃AI写墓志铭
  • 终极AI编程助手OpenCode:如何5分钟告别传统编码困境
  • SD2026 一轮省集
  • Altium Designer实战:多层PCB布线中的信号完整性与热焊盘设计
  • Whisper-WebUI语音转写工具从部署到优化全指南:解决环境配置与功能实现难题
  • 智能商品标题生成:EcomGPT-7B+Transformer实战
  • 从直流潮流到PTDF:一个电力‘老司机’的MATPOWER避坑指南与效率技巧
  • Graphormer科研效率提升方案:替代传统DFT计算的轻量级AI代理模型
  • 【企业级MCP服务模板首发】:内置JWT鉴权+OpenTelemetry追踪+动态插件热加载——仅限首批200位开发者获取的v3.2.0私有分支