ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

多分类实战:从Softmax到交叉熵的PyTorch实现与评估

2026/9/11 13:45:41 拓冰建站 浏览量
多分类实战:从Softmax到交叉熵的PyTorch实现与评估 今天是我在浙大跟着学习小组推进机器学习的第18天。前17天从线性回归、逻辑回归一路做到简单的MLP二分类已经玩得比较顺手了所以我一度以为多分类问题不过就是换个标签、加几个输出节点的事。结果真正动手做才发现多分类不只是把输出维度从1改成K这么简单它牵扯到输出层的激活方式、损失函数的设计、评估指标的选取还有一堆训练时才暴露出来的隐藏坑点。这篇文章就把我Day18这天在疏锦行项目里学习和实战多分类问题的完整过程记录下来。内容包括多分类与二分类的本质差异、Softmax与交叉熵的配合逻辑、评估指标的选取原则、一个完整的PyTorch多分类实战流程以及我当天踩过的几个比较典型的坑。如果你正在学习机器学习分类任务或者从二分类往多分类迈进时觉得哪里差一口气这篇文章应该能帮你少走不少弯路。1. 为什么多分类值得单独花一天二分类顺手多分类翻车很多教程把逻辑回归讲完之后紧接着就说一句多分类就是把逻辑回归扩展成Softmax回归然后这页PPT就翻过去了。我以前也是这么觉得的直到自己动手训练一个10分类模型才发现其中的细节比想象中多得多。这一节先不聊代码我把多分类问题的认知框架梳理清楚。1.1 二分类到多分类模型输出结构发生了什么变化二分类问题里面模型输出的是一个标量或者一个二维概率向量。以逻辑回归为例输出经过Sigmoid之后落在0到1之间这个值可以理解为属于正类的概率负类的概率就是1减去这个值。但到了多分类场景假设有K个类别模型的输出就不再是一个数而是一个K维向量。绝大多数情况下我们要求这K个维度上的值加起来等于1并且每一维的值都在0到1之间这样它才能被解释成模型对这个样本属于第k个类别的置信概率。如果你只是简单地把输出维度改成K然后仍然用Sigmoid对每个维度单独做激活这样得到的每一维虽然都在0到1之间但它们的总和并不等于1。这个结果其实是K个独立的二分类概率不是真正意义上的多分类概率分布。Day18这天我最开始就犯了这个错误后面在损失计算时怎么调都觉得数值怪怪的。这也是为什么多分类必须有专门的Softmax层或者类似机制来保证概率归一化。1.2 多分类的多种落地路径OvR、OvO与Softmax回归多分类可以走的路不止一条。最容易理解的是拆解法把K分类问题拆成若干个二分类问题。一是一对多OvROne-vs-Rest训练K个二分类器第k个分类器负责区分属于第k类和不属于第k类预测时把所有分类器跑一遍取置信度最高的那个作为最终类别。二是一对一OvOOne-vs-One对任意两类都训练一个分类器总共训练K(K-1)/2个预测时用投票法决定类别这是SVM在处理多分类时常用的手段。第三种就是Softmax回归也就是Logistic回归的直接推广。它不拆问题而是直接用一个模型输出K维概率分布让这K个概率之间互相竞争、彼此约束。神经网络处理多分类问题时几乎都走这条路线因为它端到端训练、梯度和模型结构都更自然。我在做疏锦行项目时选择的是Softmax回归方案。坦白说OvR和OvO更容易理解但放在深度学习框架里Softmax的写法最简洁、训练效率也最高而且后续不管是加正则化还是换更复杂的网络结构都是在这一套框架上扩展。1.3 多分类问题的难度来源类别间的边界竞争多分类比二分类难本质上是因为类别之间的边界变复杂了。二分类只需要找一条决策边界把空间切成两半而K分类需要找到K个区域每个区域对应一个类别区域与区域之间的边界可能是多条甚至可能是非线性的。更麻烦的是真实数据里面类别之间往往存在相似区域。比如CIFAR-10里面鸟和飞机都有翅膀猫和狗都是四条腿加毛茸茸这些类别特征重叠的地方就是模型最容易犯错的地方。二分类任务里通常只关心是与不是而多分类任务里的错误是有远近亲疏关系的——把猫认成狗和把猫认成卡车虽然都是错但前者在语义上显然更接近。这种错误结构在评估模型时也需要额外关心。理解了这一点你就会明白为什么多分类要单独做评估、单独调损失函数你不能只看正确率这一个数字你需要知道模型在哪些类别上犯糊涂它们是在把哪个类别误认成哪个类别。2. Softmax与交叉熵多分类模型的两个核心齿轮多分类深度学习模型的标准配置是最后一层输出K维向量 Softmax归一化 交叉熵损失。这套组合拳不是凭空来的每一步都有它存在的理由。Day18我把这两个东西的数学原理和代码实现都过了一遍下面是我认为最关键的几块拼图。2.1 Softmax的数学本质把得分变成概率分布假设模型最后一层全连接层输出的原始分数是 z [z_1, z_2, ..., z_K]Softmax做的事情就是P(yi|x) exp(z_i) / Σ_{j1}^{K} exp(z_j)也就是说对每个分数取指数再除以所有指数之和。这样做的效果有两个第一所有输出都是正数第二所有输出加起来等于1满足概率分布的定义。为什么要用指数而不是直接用 z_i 除以 z 的和因为原始分数可能是负的直接归一化会出问题。而且指数运算会放大分数之间的差异让原本得分最高的那个类别的概率更显著这符合我们分类时的直觉得分稍微高一点就应该有更大概率被选中。但Softmax还有一个副作用值得注意它会把差距拉得特别大。假设三个类的得分是 [2, 1, 0]经过Softmax之后概率大约是 [0.665, 0.245, 0.090]。如果得分变成 [2, 1, 0.1]看起来差别不大但概率已经变成 [0.659, 0.242, 0.099]说明Softmax对小数点后的变化也比较敏感。这个特性在模型训练后期会体现为模型的预测概率越来越高哪怕它其实没那么确定。2.2 数值稳定性问题exp一不小心就溢出Softmax里面的 exp(z_i) 在 z_i 比较大的时候会爆炸。比如 z_i 1000exp(1000) 直接就是无穷大程序里会变成NaN。这在小模型里可能不常见但一旦没做归一化就直接输入网络经常见。解决办法很简单先找出 z 里面的最大值 m然后算 exp(z_i - m)。因为 Softmax 的分子分母同时除以 exp(m)结果不变但数值范围被控制住了。这就是代码里常见的那行z z - torch.max(z, dim-1, keepdimTrue).values p torch.exp(z) / torch.sum(torch.exp(z), dim-1, keepdimTrue)我在Day18的实践里就把这一步写进了自定义逻辑中。当然如果你直接用PyTorch的torch.nn.CrossEntropyLoss框架内部已经处理好了数值稳定性不用自己操心。但理解这个处理方式仍然很重要因为当你去读别人的代码、或者自己写损失函数时这点小细节往往决定了训练能不能稳定跑下去。2.3 交叉熵损失为什么它是多分类的首选有了概率分布之后我们需要一个损失函数来衡量模型给出的概率分布与真实标签之间的差距。交叉熵的公式是L -Σ_{k1}^{K} y_k * log(p_k)其中 y 是真实标签的独热编码p 是模型预测的概率。因为 y 是独热编码只有真实类别那一维是1其他都是0所以公式可以简化为L -log(p_c)其中 c 是样本的真实类别。直观理解就是模型给真实类别分配的概率越高loss越小给真实类别分配的概率越低loss越大。这里有一个很关键的梯度性质Softmax之后接交叉熵它的梯度是∂L/∂z_i p_i - y_i这个公式非常漂亮。它意味着梯度的计算就是预测概率减去真实标签不需要链式法则一层层地算。如果模型觉得某个类别的概率是0.8而真实类别其实是那个类别梯度就是负的0.2推动模型往降低该类别得分的方向更新。这个性质让Softmax交叉熵在数值上非常稳定也是它们在分类任务里成为黄金组合的根本原因。2.4 在实际代码里CrossEntropyLoss替你做了什么PyTorch的torch.nn.CrossEntropyLoss是一个封装了三个功能的类LogSoftmax、负对数似然损失NLLLoss、以及一些内部优化。也就是说你在调用它的时候不需要在最后一层额外加Softmax直接把网络输出的logits丢进去就行。如果你在最后一层手动加了Softmax再把结果传给CrossEntropyLoss等于算了两次Softmax数值会变差训练也可能出问题。我自己的教训是网络最后的输出层应该保持裸的K维向量训练阶段用CrossEntropyLoss计算loss只有到了推理阶段需要看概率分布时才在模型输出后手动包一层Softmax。这个习惯从那之后一直没变过。3. 多分类评估指标准确率之外我更该看什么多分类任务的评估是最容易被忽视的环节。初学者一般只看一个数字——Accuracy准确率——只要模型在测试集上达到90%就觉得万事大吉。但Day18做完CIFAR-10那个例子之后我发现单单看准确率会漏掉很多信息尤其是当类别分布不均匀或者模型在特定类别上有系统性问题的时候。3.1 混淆矩阵看清每一类到底被认成了什么混淆矩阵是多分类评估的第一站。它是一个K×K的矩阵第 i 行第 j 列的含义是真实类别为 i、但被模型预测为 j的样本数量。对角线上的数字越大越好非对角线上的数字就是具体的错误模式。CIFAR-10上训练完成后我打印出混淆矩阵发现模型很容易把狗预测成猫把鹿预测成马。这些错误本身有高度的结构相似性——四足动物之间互相混淆。这种信息在准确率数字里完全体现不出来但它恰恰指导着我们下一步改进的方向是需要给某些类别加更多训练样本还是需要设计更好的特征提取器来区分相似类别。在代码层面可以用sklearn.metrics.confusion_matrix一行计算也可以用PyTorch在测试循环里自己累加confusion torch.zeros(num_classes, num_classes, dtypetorch.long) for x, y in test_loader: pred model(x).argmax(dim1) for t, p in zip(y.view(-1), pred.view(-1)): confusion[t, p] 1有了混淆矩阵你可以非常直观地定位病根。3.2 Macro-F1、Micro-F1和Weighted-F1到底该选哪个由于准确率在多类不均衡时缺乏参考价值更合理的做法是看F1。但F1在多分类环境里有三套计算方式我第一次看的时候也容易晕这里用大白话梳理一下。Macro-F1宏平均先对每个类别单独计算精确率和召回率得到一个F1然后把所有类别的F1取平均。它对每个类一视同仁不会因为某个类样本多就占更大权重。如果你特别关心小类别能不能被识别出来Macro-F1是更严格的指标。Micro-F1微平均把所有类别的预测结果汇总到一起全局统计TP、FP、FN然后计算整体的精确率和召回率最后算F1。当所有类别的样本量差不多时Micro-F1和Accuracy数值上会非常接近当类别不均衡时大类别会主导Micro-F1。Weighted-F1加权平均还是先对每个类算F1但按每个类的真实样本占比给它加权然后求和。这样既保留了逐类的F1信息又反映了类别在数据中的实际重要性是把Macro和Micro各取一半的做法。样本不均衡时我优先推荐它。3.3 不均衡多分类场景的处理思路现实中很多多分类任务不同类别的样本数量差异很大比如故障诊断里正常样本远多于故障样本图像识别里某些罕见物种几乎没啥训练数据。这时候直接训练出来的模型会倾向于把所有样本都预测成大类准确率看着挺高实际一点用都没有。处理思路通常有几种第一种是对损失函数加权给样本少的类别更大的权重第二种是过采样复制小类别的样本让它们多出现几次第三种是改用Focal Loss把焦点放在那些难以分类的样本上第四种是在评估时坚持看Macro-F1而不是Accuracy。Day18我的CIFAR-10数据还算均衡所以没有走这些复杂路子但我把Focal Loss的原理搞清楚了就是给交叉熵加一个调制因子 (1-p_t)^γ当模型已经能把某个样本分得很好时让这个样本产生的梯度权重降低迫使模型更多关注那些老大难样本。后续如果遇到不均衡数据这套思路可以直接平移过去。4. PyTorch多分类实战用CIFAR-10把概念跑通理论说得再多不如一个完整案例直接。Day18下午我用PyTorch搭了一个CIFAR-10图像分类模型从数据准备到训练再到测试把多分类的完整链路跑了一遍。下面把步骤和关键代码贴出来附带每一步的解释方便你照着复现。4.1 数据准备CIFAR-10与数据增强的取舍CIFAR-10是一个10类、每类6000张32×32彩色图像的经典数据集类别包括飞机、汽车、鸟、猫、鹿、狗、青蛙、马、船、卡车。它规模不大训练集5万张、测试集1万张非常适合在学习阶段把多分类的流程跑通。import torch import torchvision import torchvision.transforms as transforms transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) trainset torchvision.datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) testset torchvision.datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_test) trainloader torch.utils.data.DataLoader(trainset, batch_size64, shuffleTrue, num_workers2) testloader torch.utils.data.DataLoader(testset, batch_size64, shuffleFalse, num_workers2)这里有个细节值得多说一句数据增强只加在训练集测试集只用ToTensor和Normalize。RandomCrop和RandomHorizontalFlip是在训练时给模型变着花样看数据增强泛化能力但测试时必须保证数据的真实性和一致性否则测试结果会失真。Normalize那三个数分别是CIFAR-10数据集在RGB三个通道的均值和标准差。归一化不是可有可无的体操动作它能让输入特征的数值范围稳定在0附近帮助模型更快收敛也能避免某些特征数值过大导致梯度不稳定。4.2 网络结构一个足够应付CIFAR-10的简单CNN我没有一开始就上ResNet而是先用一个三层卷积加全连接的小网络把多分类流程跑通再说。网络结构是这样的import torch.nn as nn import torch.nn.functional as F class SimpleCNN(nn.Module): def __init__(self, num_classes10): super(SimpleCNN, self).__init__() self.conv1 nn.Conv2d(3, 32, kernel_size3, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) self.conv3 nn.Conv2d(64, 128, kernel_size3, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(128 * 4 * 4, 256) self.fc2 nn.Linear(256, num_classes) self.dropout nn.Dropout(0.3) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x self.pool(F.relu(self.conv2(x))) x self.pool(F.relu(self.conv3(x))) x x.view(x.size(0), -1) x F.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x32×32的输入经过三次2倍池化之后空间尺寸变成4×4所以全连接层第一层的输入维度是128×4×42048。fc2的输出是10维对应CIFAR-10的10个类别。注意forward的最后没有Softmax训练时直接把原始logits丢给CrossEntropyLoss。在动手写网络前算一算特征图尺寸是个好习惯这样能减少调试维度不匹配的时间。如果懒得算也可以先在代码里打印一句print(x.shape)确认一下。4.3 训练循环与优化器选择优化器我选了Adam学习率设成0.001这是很多小模型的稳妥起点。如果你选的SGD往往需要配合动量并且手动调学习率Adam对这种快速验证的场景更友好。import torch.optim as optim model SimpleCNN(num_classes10) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.1) epochs 30 for epoch in range(epochs): model.train() running_loss 0.0 correct 0 total 0 for images, labels in trainloader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() scheduler.step() train_acc 100.0 * correct / total print(fEpoch {epoch1}/{epochs}, Loss: {running_loss/len(trainloader):.4f}, Acc: {train_acc:.2f}%) print(Training finished.)每训练完一个epoch我还会顺便在测试集上跑一遍记录测试准确率方便观察有没有过拟合。StepLR让学习率每10个epoch衰减10倍后期收敛更稳。有个小建议训练循环里最好保留model.train()和model.eval()的显式切换。因为Dropout和BatchNorm在训练和推理时的行为不一样忘了切会导致测试集上的指标偏低或波动。4.4 测试评估从准确率延伸到混淆矩阵和F1测试阶段我把模型切到eval模式关闭梯度计算然后统计准确率、逐类精确率/召回率/F1并生成混淆矩阵。from sklearn.metrics import classification_report, confusion_matrix, f1_score model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in testloader: outputs model(images) _, predicted torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) print(fTest Accuracy: {sum(1 for p, l in zip(all_preds, all_labels) if p l) / len(all_preds):.4f}) print(classification_report(all_labels, all_preds, digits4)) cm confusion_matrix(all_labels, all_preds) print(cm)classification_report会直接给你每一类的precision、recall、f1-score以及Macro和Weighted的平均值省去自己算的功夫。有了这些数字我就能判断模型在哪些类别上弱再回头去找原因。跑完30个epoch这个简单CNN在测试集上大概能到78%左右的准确率。这跟当下动辄95%以上的大模型没法比但作为理解多分类全流程的载体已经足够了。真正有价值的不是数字而是从数据加载到评估指标这条链路的完整性和各个环节的处理逻辑。5. 多分类训练中踩过的坑Day18实测的真实教训这部分我特别想写因为Day18训练的多数时间其实不是在写网络而是在处理各种莫名其妙的报错和反常现象。我把当天踩过的坑按症状—原因—解法的方式整理出来希望你能绕开。5.1 标签格式独热编码 vs 类别索引第一次写多分类代码时我习惯性地把标签做成独热编码再送进模型。然后发现CrossEntropyLoss报错提示目标类别范围不对。这里要特别强调PyTorch的CrossEntropyLoss要求的目标是整数索引也就是0到K-1之间的整数张量形状一般是(B,)或(B, 1)而不是(B, K)的独热编码。这是很多从Keras或其他框架转过来的人常踩的坑——TensorFlow的categorical_crossentropy经常搭配独热编码而PyTorch默认用整数索引。如果你确实已经生成了独热编码需要转回整数索引用torch.argmax(y_onehot, dim1)即可。如果你更习惯用独热编码的那种写法也可以自己调用torch.nn.functional.binary_cross_entropy_with_logits单独设计损失但那样又回到多个独立二分类的套路不是标准多分类了。5.2 Loss直接变成NaN学习率过高和数值溢出我在跑一个实验时把学习率调到了0.01结果不到3个epoch损失值就一路狂飙变成NaN。原因很简单梯度更新步长太大参数一下跳到了损失函数曲面非常陡峭的区域梯度进一步爆炸最后数值直接溢出。排查这类问题我的经验是先从这几个方向下手。第一步把学习率先降回0.001看看loss是否恢复正常第二步检查输入数据里是否有NaN或异常大值归一化是否正确第三步确认最后一层输出没有手动加Softmax再送给CrossEntropyLoss第四步如果网络非常深可以考虑加梯度裁剪比如torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)。Loss变成NaN通常不是你写得不够好而是哪个数学操作在数值上不够稳。按顺序排查不要瞎试。5.3 输出层到底要不要加激活函数多分类网络的最后一层最常见的写法是不加任何激活函数直接返回logits。很多人初学时会觉得多分类需要概率输出所以最后一层要Softmax于是把Softmax放进网络forward函数的最后。如果模型只是用来做推理这没问题但如果训练时还继续用CrossEntropyLoss就会出问题。因为刚才说过CrossEntropyLoss内部自带LogSoftmax和NLLLoss你在外面已经做过一次Softmax等于先算了一次概率分布又被取了一次对数这个对数再被当作logits放到内层的LogSoftmax里数值全都乱了。我的建议是训练用的model只输出logits不输出概率。推理时再单独对模型输出执行Softmax或者直接torch.argmax(outputs, dim1)取预测类别连Softmax都可以省——因为Softmax是单调的不会改变argmax的结果。大多数情况下你需要的只是预测类别而不是精确的概率值。5.4 类别不平衡模型变成复读机我在一个小型自定义数据集上做过测试其中类别A占了80%类别B和C各占10%。模型训练完后准确率高达78%但一看混淆矩阵类别B和C几乎全军覆没A类准确率95%以上。如果只看准确率你甚至会以为模型还不错但它实际上只会无脑复读A类完全失去了分类的意义。遇到类别不平衡的多分类优先做这几件事计算每个类别的样本数打印出来让自己心里有数把loss的weight参数设置成更重视小类别的值评估时看Macro-F1而不是Accuracy如果数据量允许对小类别做过采样或数据增强。PyTorch的CrossEntropyLoss自带weight参数只需要传入一个长度等于类别数的张量class_weights torch.tensor([0.8, 1.0, 2.0, 1.0, 1.5, 1.0, 2.0, 1.0, 1.0, 1.2]) criterion nn.CrossEntropyLoss(weightclass_weights)这样一来错分小类别的惩罚更大模型就有动力去学习小类别的特征了。5.5 每个类别的样本数量、验证集与随机种子还有一个特别不起眼、但影响结果稳定性的点随机种子。多分类模型的初始化权重、数据加载顺序、数据增强的随机性都会影响最终指标。如果你跑两遍结果差很多第一件事就是固定随机种子。Day18结束时我在训练代码开头加了这几行import random import numpy as np def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) set_seed(42)别小看这个动作它能让你的实验结果可复现在调参的时候避免被随机性误导。6. 从Day18往后看多分类进阶的三个方向路跑通之后我并没有急着继续前进而是花了一点时间梳理多分类问题还能往哪些方向深入。这里把三个我认为最值得关注的方向列出来供同样学到这里的朋友参考。6.1 Focal Loss与难样本挖掘Day18用标准交叉熵跑CIFAR-10时模型最后卡在78%左右上不去很大一部分原因是那些容易混淆的样本没有被特别照顾。标准交叉熵对已经分类正确且高置信度的样本也会产生梯度导致模型把大量精力浪费在简单样本上。Focal Loss做的事情就是压低简单样本的贡献权重让模型集中火力处理困难样本。它在交叉熵前面乘了一个调制因子FL -α_t * (1 - p_t)^γ * log(p_t)当 p_t 接近1时候系数接近0该样本对loss几乎没有贡献当 p_t 偏低时系数较大贡献被保留。γ一般取2α_t用来调节正负样本不平衡。这个损失函数最初出现在目标检测领域但现在已经被广泛用在各种不均衡分类任务中。6.2 标签平滑让模型别那么自信交叉熵损失为了最小化loss会逼着模型把真实类别对应的概率推向1。但在训练数据本身存在噪声或标注错误时这种过度自信会让模型记忆噪声反而降低泛化能力。标签平滑的做法是把真实的独热编码修改为y_smooth(k) 1 - ε当 k c 时y_smooth(k) ε / (K - 1)当 k ≠ c 时这里的 ε 是一个很小的超参数通常取0.1。它告诉模型真实类别不一定是绝对正确的其他类别也可能有一点概率等于给模型施加了正则化限制了输出概率过于极端。在很多图像分类竞赛和业务模型里标签平滑都能稳定提升泛化性能是我比较推荐尝试的进阶技巧。6.3 从多分类到多标签Sigmoid与Binary Cross Entropy的过渡多分类的另一条进阶路线是多标签分类。多分类里一个样本只能属于一个类别但很多真实场景里一个样本可能同时拥有多个属性标签比如一张图片里面同时有猫和狗或者一篇文章同时涉及科技、经济两个主题。多标签的做法和标准多分类有个关键差异输出层不再用Softmax归一化而是用Sigmoid对每个类别独立激活每个维度的输出表示属于该类别的概率彼此之间不竞争。损失函数也从CrossEntropyLoss换成torch.nn.BCEWithLogitsLoss。这个转换说难不难但思维方式需要转一个弯多分类是互相排斥的K选1多标签是相互独立的K个二分类。理解了这两个问题的区别你的分类知识体系就会更完整以后再遇到各种业务场景也能快速判断用哪种方案。Day18这一天我自己最大的收获倒不是记住了多少公式而是真正建立了多分类是一个完整闭环的意识从模型输出层的设计到损失函数的选择再到评估指标的解读任何一个环节拿捏不准都会让整个任务的质量打折扣。如果你现在也卡在二分类往多分类过渡的阶段建议你动手写一个完整的小项目把今天文章里提到的每个环节都亲自跑一遍。纸上得来终觉浅等你真的看到自己在CIFAR-10上训练出了第一个多分类模型那种通了的感觉会比读十篇笔记都管用。