尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

多分类问题核心机制与PyTorch实践:Softmax与交叉熵

多分类问题核心机制与PyTorch实践:Softmax与交叉熵 1. 多分类的定位从单一答案到概率分布1.1 为什么多分类不是二分类的简单扩展Day18这节课的标题看起来很简单就叫“多分类问题”但我跟完整个课程后发现它其实是很多入门选手第一次真正触碰到模型如何表达不确定的地方。先说说它和二分类的本质区别。二分类输出一个值通过Sigmoid把它压到0到1之间大于0.5判为正类小于0.5判为负类。模型只需要回答是还是否决策边界在概念上是一维的。但多分类问题中模型面对的是10个、100个、甚至1000个候选类别它必须输出一个完整分布告诉你不只是哪个类别最有可能还包括每个类别分别有多大的概率。课程里举了一个很贴近生活的例子手写数字识别。一张图可能是0到9中的任何一个数字模型要输出10个得分这10个得分经过处理后变成一个概率分布。如果你写了一个有点歪的7模型可能会给7打出0.82的概率给1打出0.13的概率剩下0.05散落在其他数字上。这不是模型不自信而是它合理表达了自己的判断。在二分类中你很难感受到这种分布的质感一旦到了多分类整个输出的结构就完全变了。1.2 理解多分类是往下走所有复杂任务的前提我当时听这门课的时候有个体会多分类不只是图像识别的基础它几乎串联了后续所有任务形态。比如目标检测里的分类分支每个候选框要去判断它属于哪一类物体文本情感分析里模型要在正面、负面、中立之间做选择语音指令识别本质上是把音频段匹配到有限个指令槽位中。只要任务不是是/否这种二值判断你就逃不开多分类的设计逻辑。另外一个容易被忽略的点是多分类的误区理解会直接影响你对回归任务的认识。有些人把分类和回归完全对立但多分类在输出层做的事——把连续得分映射成概率——其实和回归里的归一化处理非常像。区别在于分类的监督目标是离散的标签回归的监督目标是连续的数值。Day18课程在开篇花了大段时间说清楚这个边界我觉得核心目的就是让你别把任务类型搞混因为任务类型决定了你选什么损失函数而损失函数选择错了后面的训练过程就会非常痛苦。适合来看这部分内容的人我大致归为三类一是学过基本线性模型和Sigmoid二分类正要往多分类过渡的入门者二是已经用现成框架跑过分类任务但没搞清楚Softmax和交叉熵为什么是标准搭配的初学者三是想从头捋一遍分类模型设计细节方便自己改模型、调试loss的进阶同学。这门课的内容并不局限于某个框架但它用PyTorch做演示的时候非常接地气后面的实操部分我会把具体实现也一并拆开讲。2. 核心机制拆解Softmax和交叉熵的配合逻辑2.1 模型最后一层到底在输出什么多分类模型的结构前面无论用多少层卷积、Transformer、全连接真正到了最后一层通常就是接一个线性层把隐藏状态映射成类别数量的实数向量。比如Fashion-MNIST有10个类别所以最后一层通常是nn.Linear(hidden_dim, 10)输出的10个数字被称为logits。一开始我不太理解为什么要管这个中间结果叫logits。后来课程里补了一段说明logits可以简单理解为未经归一化的得分这个得分取值范围是负无穷到正无穷。它本身没有任何概率含义因为它的数值可能很大也可能很小不同样本之间的得分差异也可能巨大。模型在学习过程中真正要学的就是怎么让正确类别的logit尽可能大、错误类别的logit尽可能小。从损失函数角度看模型优化的是这个得分分布的差距而不是直接优化一个所谓的置信度。所以你在观察一个模型输出的时候如果只盯着最大logit对应的类别其实会失去很多信息。Day18课上老师展示了一张表格把同一个batch里的logits打印出来你会发现有些样本的最大得分和次大得分差距很小这种样本往往是容易混淆的类别比如数字4和9或者T恤和衬衫。这类样本如果只记录预测对了没就完全丢失了训练过程中最有价值的信号。2.2 Softmax如何把得分变成概率要让logits变成可解释的概率分布最常用的方法就是Softmax。它的公式写出来很简洁softmax(z_i) exp(z_i) / sum_j exp(z_j)也就是说对每个logit取指数再除以所有指数和。这么做有两个好处第一指数函数保证所有结果都是正数第二除以总和保证所有类别的输出加起来等于1。这样一来模型输出的就是一个合法的概率分布。但这里有个很容易踩的数值稳定性问题。如果某个logit特别大比如到了20exp(20)会变成一个非常大的天文数字计算时可能溢出。Day18课程专门讲了这个细节实际操作中会对每个logits向量先减去它的最大值再做Softmax。因为减去同一个常数Softmax的结果是不变的但数值的绝对值被拉回到安全区间。这件事在PyTorch里其实已经被封装好了你直接调torch.nn.functional.softmax(logits, dim1)不会遇到溢出但了解原理能让你在手动实现或者换框架排查问题时心里有底。从直觉上理解Softmax可以想象成把一班学生的原始考试成绩转换成排名概率。原始分数差距可能很大但转换后分数最高的那个学生获得最高概率其他人的概率则按相对差距指数级递减。如果两个学生的原始分非常接近那他们的概率也会很接近。这正好对应分类模型里的类别间分差越小判断越犹豫。2.3 为什么损失函数选择交叉熵而不是均方误差多分类最常用的损失函数是交叉熵原因值得展开说。如果从公式看对于单个样本假设真实类别是y模型输出的概率分布是p交叉熵损失就是loss -log(p[y])这里面有意思的点是因为真实标签是one-hot形式的其他正确类别的项全都乘了0最后只留下了模型给正确类别预测概率的负对数。所以交叉熵本质上是在惩罚模型对正确类别不够自信。如果模型给正确类别的概率是0.9损失是-log(0.9)非常小如果概率是0.1损失变成-log(0.1)非常大。那为什么不沿用回归任务里的均方误差MSE呢课程里给了一个非常直白的解释MSE配合Sigmoid或者Softmax会出现梯度消失的问题。当Softmax输出的概率接近0或1时Sigmoid曲线的导数会接近0误差通过链式法则传回前面时梯度几乎消失了训练会变得非常慢。而交叉熵配合Softmax求导之后形式非常干净误差直接和模型预测概率减真实标签成正比也就是说模型错得越离谱梯度越大学得越快。我个人的体会是这个点如果你只背结论永远感受不到差别有多大。我自己试过一次把交叉熵换成MSE去训练一个简单的多分类模型结果训练loss下降得极其缓慢跑了十几个epoch准确率还在原地打转。那一刻才真正明白为什么标准方案是Softmax在输出层 交叉熵做损失这不是历史遗留的偏好是数学上配合出来的最优解。2.4 一个稳妥的数值处理方案课程里还专门提了一个训练时的细节PyTorch里的nn.CrossEntropyLoss在接收输入时期望的是原始logits而不是已经过Softmax的概率。这个设计很容易让刚上手的人搞混。如果你手动先做了一个Softmax再把结果丢进nn.CrossEntropyLoss损失并不会报错但结果和你预期的不一样因为CrossEntropyLoss内部本身已经包含了一个LogSoftmax的操作你多算一次等于把分布又扭曲了一遍。所以标准做法是模型最后一层直接输出原始logits计算loss时传logits和标签计算预测结果时才用torch.softmax或者torch.max取索引。这套设计的另一个好处是在数值上更加稳定因为LogSoftmax和NLLLoss组合时候可以做各种数学化简避免了先算Softmax再算log产生的中间数值误差。3. 动手实现用PyTorch搭一个完整的图片多分类流程3.1 数据集准备与数据加载细节Day18的课程例子用的是Fashion-MNIST我觉得这个选择很聪明。MNIST数字识别被用得太泛滥了大家熟到已经没法产生真实的分类难度感知。Fashion-MNIST里有T恤、裤子、套头衫、裙子、外套、凉鞋、衬衫、运动鞋、包、短靴类别同样是10个但某些类比之间长得确实像比如衬衫和T恤或者凉鞋和运动鞋这样训练过程中你能更真实地感受到有些错误是模型必然会犯的。加载数据的代码可以直接用torchvisionimport torch import torchvision import torchvision.transforms as transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.2860,), (0.3530,)) ]) train_dataset torchvision.datasets.FashionMNIST( root./data, trainTrue, transformtransform, downloadTrue ) test_dataset torchvision.datasets.FashionMNIST( root./data, trainFalse, transformtransform, downloadTrue ) train_loader torch.utils.data.DataLoader( train_dataset, batch_size64, shuffleTrue ) test_loader torch.utils.data.DataLoader( test_dataset, batch_size64, shuffleFalse )这里Normalize的参数是Fashion-MNIST数据集的全局均值和标准差torchvision官方文档或者很多开源项目里都能查到。要注意的是用错均值和标准差不会让程序报错但模型收敛速度会变差因为输入特征的分布没有被拉到零附近。我见过有人直接照搬MNIST的(0.1307, 0.3081)训练出来效果也能用但这属于能用但不够好的状态。3.2 搭建一个足够用的小型CNN课程并没有一上来就堆ResNet那种大模型而是用了两层卷积加全连接的小网络。目的是让你把注意力放在多分类的整体流程上而不是被复杂的网络结构带偏。这个选择放在入门阶段非常正确因为多分类要解决的核心问题——输出层设计、损失计算、评估逻辑——和网络的深度没有直接关系。import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 16, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(16, 32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(32 * 7 * 7, 128), nn.ReLU(inplaceTrue), nn.Linear(128, num_classes), ) def forward(self, x): x self.features(x) x self.classifier(x) return x注意看forward的返回值没有经过任何Softmax它就是原始的logits。如果你在最后一层后面加了nn.Softmax之后又用nn.CrossEntropyLoss算损失那就会出前面说的重复Softmax问题。为了避免这种混乱最安全的做法是模型只负责输出logits概率转换放到后续评估阶段做。3.3 训练循环中的损失计算与精确率统计训练一个epoch的流程在外观上跟二分类很像但是有一个容易忽略的关键点你计算准确率的时候是在logits上取最大值索引不是取概率。因为Softmax不会改变logits的相对顺序最大值在转换前后是同一个位置所以直接对比logits的argmax和标签就可以了。def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss 0.0 correct 0 total 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits model(images) loss criterion(logits, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) preds logits.argmax(dim1) correct (preds labels).sum().item() total images.size(0) avg_loss total_loss / total accuracy correct / total return avg_loss, accuracy这里我习惯用loss.item()乘上batch大小来累加总损失最后算平均这样能防止最后一个不完整batch对平均值的打扰。如果你直接用每次loss.item()求平均当数据总数不能被batch_size整除时最后一个batch样本数少但权重和别人一样统计出来的均值就有轻微偏差。虽然影响不大但在对比实验里这种微小的统计口径不一致容易让你误判模型好坏。模型定义、优化器和损失函数的选择如下device torch.device(cuda if torch.cuda.is_available() else cpu) model SimpleCNN(num_classes10).to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) epochs 10 for epoch in range(1, epochs 1): train_loss, train_acc train_one_epoch( model, train_loader, optimizer, criterion, device ) print(fEpoch {epoch:02d}, Loss: {train_loss:.4f}, Acc: {train_acc:.4f})3.4 额外需要探索的关键超参如果你想自己重复实验并观察不同设置的影响我会建议你像做对照实验一样固定其他变量只改一个参数。课程里花了相当篇幅讲学习率的影响因为这个参数可以说是所有超参数里最核心的。用1e-3的Adam优化器在Fashion-MNIST上通常收敛得不错但如果把学习率调到1e-1你会发现loss在训练初期疯狂震荡甚至直接变成nan调到1e-5训练倒是稳定可十个epoch过去loss还在原地踏步这时候你就该明白学习率不是越大越好也不是越小越稳而是需要在学得动和学得稳之间找平衡。批量大小同样值得控制。常见的选择是32、64、128。在同样学习率下大batch会让梯度估计更准收敛更平滑但每步更新次数变少小batch会有更多噪声有时候这种噪声反而帮助模型跳出局部极小。不要盲目迷信某个最佳默认值多跑几次对比实验记录下训练曲线比看单个准确率数字有用得多。3.5 验证集评估是必要的检查动作训练集上的准确率不能说明模型的真实能力因为模型可能把训练数据背下来了。Day18课程里专门强调了验证集/测试集评估的重要性。Fashion-MNIST把train和test划分好了训练结束之后应该在test_loader上做一次完整的预测统计。def evaluate(model, loader, criterion, device): model.eval() total_loss 0.0 correct 0 total 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) logits model(images) loss criterion(logits, labels) total_loss loss.item() * images.size(0) preds logits.argmax(dim1) correct (preds labels).sum().item() total images.size(0) return total_loss / total, correct / total模型训练和评估要切换modemodel.train()和model.eval()对没有Dropout和BatchNorm的简单模型影响不大但如果你后面用到了BatchNorm忘记切到eval模式会导致评估时仍然用当前batch的统计量做归一化结果会和真实推理情况不一致。养成习惯训练和评估各归各的mode避免后面踩这个暗坑。4. 常见问题排查与调参建议4.1 训练loss不下降的几类原因如果你按照上面的流程跑下来一切顺利那自然是好事。但很多人第一次接触多分类时都会在某个环节卡住。我把自己遇到过和帮别人排查过的问题整理成一个速查表基本覆盖了大部分情况。现象可能原因检查方式与解决方法loss一开始就很高而且不下降学习率设置过大或过小打印每个batch的loss观察是否震荡试试1e-2、1e-3、1e-4的对比loss快速降到某个平台后停滞模型容量不足或特征提取不够增加卷积层通道数或全连接层宽度loss是nan学习率过大导致梯度爆炸调低学习率检查输入数据是否包含异常值训练准确率很高但测试准确率很低过拟合增加数据增强、增加Dropout、缩小模型容量准确率一直在10%附近不动训练流程存在根本性bug用一小批数据先跑一个batch看loss有没有下降检查标签是否从0开始编号、类别数是否匹配课程里最强调的一个排查思路是拿到一个新模型新数据先别急着跑全量epoch而是先取一小部分数据比如一个batch看看模型能不能过拟合。如果连几个batch的数据都学不进去那说明模型实现或者loss计算有问题而不是数据量不够的问题。这是一个非常高效的冒烟测试手段能帮你把代码有bug和模型能力不足快速分开。4.2 logits输出分布异常的含义有时候你会发现模型输出的概率分布非常平所有类别都大约在0.1附近。这通常意味着模型还处在欠拟合状态或者特征根本没学到有区分性的信息。直观理解类别完全分不开时每个类别概率都一样。如果你在训练到一半时看到这种分布别急先看看loss是否还在下降。如果loss在降说明模型正在学习只需要更多epoch如果loss不动说明学习率不合适或者模型结构有问题。另一种常见情况是模型对某一个类别给出了接近1的概率无论输入什么样本都这样。这时候往往不是模型太自信而是数据有问题——比如训练集中某个类别占了90%以上模型学到了把所有样本都判成这个类别这种偷懒解。这种问题的根源不在损失函数而在数据不均衡。简单的缓解思路是调整类别权重把样本少的类别在loss里的权重调高。PyTorch里nn.CrossEntropyLoss有一个weight参数可以传一个和类别数相等的权重向量。4.3 使用交叉熵时容易犯的隐蔽错误前面说了CrossEntropyLoss接收logits而不是Softmax概率这是最容易被忽略的坑。但还有一个更隐蔽的坑是标签类型和取值范围不匹配。PyTorch的CrossEntropyLoss要求标签是整数张量每个值在0到类别数减1之间。如果你用了one-hot编码的标签需要先用argmax把one-hot转成整数标签如果你标签从1开始计数记得减1。这些错误不像类型不匹配那样一定会崩溃有可能只是静默地产生一个很大的loss然后整个训练就这么跑飞了。我自己排查这类bug的经验是在训练循环开始前先单独打印一条logits的shape、标签的shape、标签的最小值和最大值以及标签的dtype。这几个信息足够帮你排除大部分低级错误。4.4 类别不均衡时怎么办Fashion-MNIST是均衡数据集每个类别样本数一样这其实掩盖了多分类实战中最大的一类问题。你想一下如果某类只占总样本的1%就算模型把这一类的所有样本都判错它的整体准确率仍然可能高达99%但显然这个模型并不好用。应对不均衡的常见策略从数据层面可以做重采样把样本少的类别复制多份或者用over-sampling策略。在loss层面就是给不同类别加权增大少数类的惩罚力度让模型不敢忽视它们。还有一种更细化的指标用法不要只看accuracy而是看每个类别的precision、recall和F1。课程里提了一句非常重要的话——accuracy只告诉你平均表现不适合告诉你哪一类出了问题。实际情况确实如此工程上你特别需要知道某个细分类别是不是在拖后腿这时候类别级别的统计量比单一准确率可靠得多。5. 和回归任务、多标签任务的边界区分5.1 不要把多分类问题做成多个二分类有些任务看起来类别很多但本质上真的需要输出一个非此即彼的分布。比如手写数字一张图不可能同时是3和5所以这么做多分类是对的。但有些任务比如一张图片里同时有人、车、树它不是一个多分类问题而是一个多标签问题因为多个类别可以同时成立。如果你把多标签问题错误地包装成多分类问题让模型在人、车、树几个标签里只能选一个那在有人又有车的图片上无论模型选择人还是车输出都对一半错一半。正确的做法是对每个标签单独做二分类输出多个独立的Sigmoid值或者用其他为多标签设计的结构。这个差别在Day18课程里用一张对比图讲得特别透彻我想强调一下输出层用Softmax还是多个Sigmoid取决于是独占还是可共存。如果你准备做YOLO这类目标检测你真的同时需要分类头来做独占类别的判断、还需要回归头来框位置理解多分类和多标签的边界会很有帮助。5.2 分类和回归的任务界限回归和分类的界限也要想清楚。你预测明天的气温是25度这是一个回归问题你预测天气是晴、多云、下雨、下雪这是一个多分类问题。两者最核心的区别是目标空间是否离散。如果目标值是连续的用MSE或者MAE这类回归损失如果目标值是离散的类别用交叉熵。但现实世界里有些任务处在中间地带比如预测一个影视评分可能是0到10的连续值但最终展示给用户的时候常常被离散化成星级。你可以把它当回归做再取整也可以把它当分类做再映射到分数区间。这两种方案各有优劣前者更细粒度但可能预测出9.7这种不好展示的值后者好展示但会丢失分数内的细微差别。遇到这种任务选择哪种方式取决于你的应用场景更在乎什么。这不只是一个模型问题更是一个工程取舍问题。Day18把这种边界说得挺清楚的它是多分类内容里最容易被人忽略的下半场。如果你只记得怎么写模型、怎么算loss却不知道自己面对的任务到底属于哪一类很可能在真实项目里把方案设计错了方向。工具用得再熟练方向错了也是白搭。6. 实际操作中的一点体会6.1 可视化是理解模型的捷径前面说了一堆理论和代码最后一个个人觉得非常有用的技巧是可视化。具体做法很简单模型预测完一批测试集样本你把那些预测错的图片连同正确标签、预测标签一起打印出来人眼扫一遍比看任何指标都能更快地发现规律。我当时用Fashion-MNIST做实验时连续打印了三十多张被分类错误的图片发现一个特别显著的共性很多错误发生在形态相近的类别之间比如衬衫被识别成T恤因为Fashion-MNIST图片分辨率只有28x28很多衬衫和T恤在这么低的分辨率下确实极其相似。这个观察让我明白这些错误不一定是模型笨而是原始图像信息本身存在歧义。如果你是个工程师此时应该考虑换更高分辨率的数据集或者增加特征输入渠道如果你是个竞赛选手可以考虑把相似类别的区分作为模型结构设计的重点。6.2 从loss曲线里读出训练状态我还想多说一句关于记录loss的习惯。跑实验的时候我通常每个epoch结束都会把训练loss和测试loss放在同一张图里画出来。这两个曲线的相对关系能告诉你很多东西如果训练loss还在下降但测试loss已经开始回升说明模型开始过拟合了你该考虑加正则化或者提前停止如果两个loss一直在下降但速度很慢可能是学习率太小可以试试用学习率调度器在中间阶段加速或者换用更大的初始学习率。这里推荐一个简单的做法训练时把loss值追加到一个list里不要只print。等训练完用matplotlib画一张loss曲线图既方便自己回顾也方便写博客记录或做报告。这个习惯在只有几十行代码的实验里显得不值一提但一旦模型变大、训练变长你就会感激自己当初记录了这些曲线。调参这件事说到底就是和loss曲线打交道——能正确读懂loss曲线就具备了独立调优的基本能力。6.3 多做实验、少背结论Day18的课程看下来所有例子都在反复说明一件事多分类问题的标准方案不是背下来的而是在一次次实验、观察和调试中沉淀出来的。你有机会也值得自己动手做几个对比实验——把Softmax换成没有归一化的原始输出、把交叉熵换成MSE、把Sigmoid硬套在多分类输出上去真实感受一下每个组合带来的差异。很多结论只有你自己亲自踩过坑、试过错误方案之后才会真正变成你自己的理解。另一个建议是把所有实验结果做记录。今天改了学习率明天改了网络宽度后天加了数据增强如果全都靠记忆一周后就全乱了。我用一个简单的表格记录每次实验的设置和结果几行字就够了但对后续迭代的帮助是巨大的。这也是我学完Day18之后最深的感受——深度学习入门从来不缺理论资料缺的是一套属于自己的、认真的实验习惯。
返回列表