多分类问题,说实在话,很多人在打卡学习时都把它当成“二分类的自然延伸”,觉得只要把输出节点多接几个就完事了。但真正把 Day18 这期内容复盘完,我发现多分类才是一个分水岭:前面学的逻辑回归、BP 网络,在这里第一次要系统性考虑“不止一个类别”时的建模方式、损失函数、评估口径和样本不均衡问题。这篇文章就是把这期分享展开揉碎,把思路、代码和坑一次讲清楚。
所以这篇笔记不是只讲理论,而是从“问题建模”一直讲到“模型评估后的改进方向”,包含一个可直接运行的 PyTorch 多分类示例和一张问题排查表。适合两类人:一类是正在跟着打卡系列学机器学习的初学者,适合用来把概念真正落地;另一类是已经跑过几个模型、但总在类别不平衡和评估指标上踩坑的工程师。相信我,把多分类里面那几个关键选择搞明白,后面做多标签分类、目标检测会顺很多。
1. 从二分类到多分类:思路转换的三个关键点
1.1 多分类不是“多接几个 sigmoid”,而是换一种概率建模方式
二分类模型最常见的输出层是 1 个神经元,后面跟 sigmoid,把任意实数压缩到 (0,1),表示正类概率,1 减去它就是反类概率。这个设计天然满足“两个类别互斥、概率加起来等于 1”的约束。
到了三分类、十分类,很多新手第一反应是:我接 K 个神经元,每个都用 sigmoid 激活,这样每个类别都能输出一个 0 到 1 之间的概率,K 个加起来甚至不用等于 1,取最大的作为预测类别不就行了?
想法没有错,这其实就是“多标签”或者“非互斥多分类”的常见思路。但绝大多数业务场景里,类别是互斥的:一张图片不可能既是猫又是狗,一封邮件只能被标记为“垃圾”“正常”“可疑”中的一种。互斥类别意味着 K 个输出之间不是独立的,而是存在“概率和为 1”的约束关系。这时 Softmax 才是在数学上更合适的选择:先把 K 个 logits 做指数映射,再归一化,得到一组非负、和为 1 的类别概率分布。
我自己的理解是把 Softmax 看成“多个 sigmoid 的升级版”。Sigmoid 只处理“我和 0 比大小”,Softmax 是让一群候选类别互相竞争,谁 logit 大谁就赢得更高概率。这种竞争关系,才是多分类问题真正的建模核心。
1.2 标签不是随便编码的,整数标签和 one-hot 都行,但要和你选的损失函数配对
多分类的标签有两种常见表示方式,一种是对应类别的整数索引,比如 0、1、2;另一种是 one-hot 编码,比如类别 1 编成 [0, 1, 0]。初学者最容易在这里被报错搞晕,因为不同框架的默认要求不一样。
以 PyTorch 为例,nn.CrossEntropyLoss 接收的是整数索引标签,且形状是 [batch_size],不能直接传 one-hot 向量。如果你非要用 one-hot,那就得把网络输出过 Softmax 后取对数,再配合 nn.NLLLoss,等于手工拆成交叉熵的两个步骤。正常情况下直接用 PyTorch 内置的 CrossEntropyLoss 就行,它会自己在内部完成 logits 到概率分布的转换。
TensorFlow/Keras 那边逻辑又不一样,categorical_crossentropy 默认配 one-hot 标签,sparse_categorical_crossentropy 配整数标签。很多从 Keras 跳到 PyTorch 的人,第一周全在跟标签格式搏斗。我的建议是:读框架文档时,重点关注 loss 的输入签名,不要靠记忆硬搬。
1.3 为什么分类问题很少用均方误差,交叉熵的“好”在哪里
我在打卡群里问过一个问题:如果我把三分类标签做成 one-hot,网络输出也过 Softmax 成概率分布,那用均方误差逼它接近 [1,0,0] 不行吗?直观上也说得通。
但实际训练起来,均方误差在分类任务上通常又慢又不稳。原因是多方面的。第一,MSE 是在概率空间算欧氏距离,但概率分布本身有“归一化”约束,两个分布之间用欧氏距离衡量并不合理;第二,MSE 对 Softmax 输出求梯度时,在输出接近 0 或 1 的饱和区梯度会变得很小,模型学不动;第三,分类任务里我们要优化的是“置信度分配是否准确”这件事,不是“数值回归误差是否够小”。
交叉熵为什么适合分类?它只关注正确类别位置上的概率,当正确类预测概率接近 1 时 loss 接近 0,当正确类概率很低时 loss 会迅速变大,这种“惩罚错得离谱”的特性,让模型在早期训练阶段也能得到明显梯度。打个比方,MSE 像是一个老师,只要你离标准答案“距离够近”就给你及格;交叉熵像一个更严格的老师,你只要把高概率分给了错误类别,不管整体距离多近都会重罚。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 模型选型与评估指标:多分类工具箱要成套使用
2.1 传统机器学习里的多分类套路:OvO、OvR 与原生多分类
如果你用 sklearn 的 LogisticRegression 或 SVM 做多分类,内部其实默认走两种策略之一:一对多或者一对一。
一对多会把 K 类问题拆成 K 个二分类器,每个分类器负责判断某个类 vs 其余所有类,预测时选置信度最高的分类器对应的类别。一对一则是在任意两个类别之间都训练一个分类器,总共 K×(K-1)/2 个,预测时投票。两种方案各有取舍:OvR 分类器数量少,训练快一点,但每个二分类器面对的“负类”样本噪声更大;OvO 每个子任务更纯粹,但子模型多,预测慢,在类别很多时开销会明显上升。
不过到了树模型和神经网络时代,情况又不一样了。随机森林这类模型能直接输出类别概率,XGBoost 训练多分类时一般用 softmax 目标函数,树模型内部仍然会构造二分类残差,但对外暴露的是 K 个概率输出。GBDT 这类模型在多分类上也能用,只是类别数太多时训练会比较重。
我的建议是:如果特征是表格型、样本量中等,先跑随机森林或 XGBoost 这类传统模型,把 feature importance、bad case 分布摸清楚,再去上神经网络。多分类问题里“模型怎么选”不是唯一关键,数据长什么样、噪声大不大往往更决定模型上限。
2.2 评估指标别迷信准确率,亲手算一遍 macro-F1 就全懂了
这里我用一个三分类例子把指标细节算一遍,例子不一定来自真实业务,是为了让你看清准确率会怎么骗人。
假设真实分布是:A 类 80 个,B 类 60 个,C 类 40 个,模型预测完成后,混淆矩阵如下:
| 真实\预测 | A | B | C |
|---|---|---|---|
| A | 70 | 8 | 2 |
| B | 10 | 45 | 5 |
| C | 4 | 6 | 30 |
对角线总和是 70+45+30=145,总样本 180,所以准确率 Acc = 145/180 ≈ 0.806。直观看,模型整体判断对了八成,好像还行。
但对 A 类来说,真实 80 个里有 10 个被认错,10 个错里面 8 个去了 B、2 个去了 C。对 C 类来说,真实 40 个里被错认了 10 个,等于丢了四分之一。如果这个 C 类是异常告警或者故障类型,漏掉 25% 是不能接受的。
再看 precision 和 recall。A 类 TP=70,它的列合计是 84,所以 P_A = 70/84 = 0.833;真实类别是 A 的总数 80,所以 R_A = 70/80 = 0.875。同理:
- B 类:P_B = 45/(8+45+6≈59?) 这里看列 B 合计 59,实际 8+45+6=59,P_B=45/59≈0.763;R_B=45/60=0.750。
- C 类:列 C 合计 37,实际 2+5+30=37,P_C=30/37≈0.811;R_C=30/40=0.750。
宏平均就是先算每个类别的指标再取平均:
Macro P = (0.833+0.763+0.811)/3 ≈ 0.802
Macro R = (0.875+0.750+0.750)/3 ≈ 0.792
宏平均对每个类别都一视同仁,不会因为 A 类样本多就偏向 A。如果你想看“每个类别都被平等对待时的平均效果”,Macro-F1 是最常用的。另外还有一个权重平均,会按真实样本占比加权,这里就是 A 类权重大、C 类权重小,算出来的分通常比 Macro 高一些,看起来更好看,但也容易掩盖少数类表现差的事实。
我强烈建议做多分类时至少看一眼“按类别分开的 precision、recall、F1”,而不是只看整体 accuracy。先找到最弱的那一两类,再往下排查,是数据量不足、特征不明显,还是样本质量差,这才是多分类评估真正要做的事情。
2.3 混淆矩阵要可视化出来,光看文字数字发现不了“哪两类容易混”
有一类现象只在混淆矩阵的热力图里能被一眼发现:模型经常把 A 类误判成 B 类,或者 C 类和 D 类来回混淆。如果你只看数字报表,很难形成直觉。
遇到这种规律性误判,常见的改进思路有三条。第一,去检查被误判的样本,看看是不是标注本身就有人为错误;第二,看这两个类别在原始特征上是否存在重叠,比如两个数字的某些笔画区域非常相似,必要时增加能区分它们的前置特征;第三,针对这类样本做难例挖掘,把置信度低或预测错误的样本单独挑出来,加入训练集重点学习。
3. 类别不平衡:多分类实战里被低估最多的坑
3.1 少量类掉点,往往不是模型能力问题,而是数据分布问题
多分类场景里,类别不平衡几乎是常态。故障诊断里正常样本比故障样本多得多,垃圾评论分类里正常评论占主流,工业质检里良品和不良品比例可能相差 100 倍。有些初学者发现 model 整体准确率不低,但一打印每类 F1,发现少数类的 recall 几乎为 0。
原因不复杂:模型在训练过程中为了让 loss 尽量低,只要把所有样本都判成多数类,就能获得一个看似不错的准确率,代价却是少数类完全被忽略。假如 98% 是 A 类,模型把全部样本都预测成 A,准确率都有 98%,但这样的模型在业务里毫无价值。
如果你发现验证集上某一类 recall 特别低,第一反应不是换更大的模型,而是去看这一类在训练集里的样本量。很多情况下就是样本太少了,模型根本没见过足够多的差异模式,学不出稳定的特征边界。
3.2 常用的三类处理手段:数据层面、算法层面、后处理层面
第一类手段是在数据层面做重采样。对少数类做上采样,简单复制样本虽然能让数量平衡,但容易造成过拟合,因为复制出来的样本只是原有样本的重复,没有新增信息量。稍微高级一些的做法是用 SMOTE 在少数类样本之间插值生成新样本,但这类方法对高维图像、文本数据不太合适,更多用在表格数据里。对多数类做欠采样则是随机删掉一部分多数类样本,数据量小的话会比较浪费。
第二类手段是在算法层面给不同类别分配不同权重。最直接的方式是给 CrossEntropyLoss 传入 weight 参数,少数类的 loss 权重调高,让模型在梯度更新时更重视少数类的分类错误。比如二分类时正负样本 1:9,可以设置正类 loss 权重为 9,负类为 1。多分类时可以根据 总样本数 / (类别数 × 每类样本数) 来算权重,也可以直接用 sklearn 的 compute_class_weight 计算。
第三类手段是后处理层面调整预测阈值。二分类里你很清楚“0.5 是阈值”,但多分类的 Softmax 输出其实不能简单说“阈值取多少”,因为各类概率是相互制约的。你可以做的是在验证集上搜索一个温度系数或概率缩放系数,来调节模型整体置信度;如果业务对某几类有特定偏好,还可以在 logits 上给特定类别加偏置,这属于更工程化的后处理技巧。
我个人的经验是:优先尝试 loss weight,因为改动最小、不容易引入数据分布变化带来的副作用。如果少数类还是学不好,再做上采样或难例挖掘。
3.3 Focal Loss 这类“升级版损失函数”什么时候才值得用
2017 年提出的 Focal Loss 最初是为了解决一阶段目标检测里前背景极度不平衡的问题,后来也被用到很多分类任务里。
Focal Loss 的核心是在交叉熵前面乘一个调制因子 (1-p_t)^γ,其中 p_t 是模型对正确类别的预测概率。如果模型对样本判得很有把握,p_t 接近 1,调制因子接近 0,这个样本贡献的 loss 会被压低;如果模型对样本拿不准或判错了,p_t 比较小,调制因子接近 1,loss 会保留下来。配合类别权重,模型会专注在少部分难分样本和少数类样本上。
但要注意,不是所有不平衡问题都适合 Focal Loss。如果少数类本身数据质量差、噪声大,Focal Loss 会让模型更努力地去拟合那些错误标注样本,反而把模型带偏。先用类别权重跑一版 baseline,如果少数类 recall 依然不够,再引入 Focal Loss 对比实验,这是我更推荐的做法。
4. 手把手实操:一个可以直接跑的三分类示例
4.1 数据准备:从现有数据集中裁一个三分类子集
直接拿完整的多分类大作业来讲,容易被细节淹没。这里我以 MNIST 手写数字为例,只取 0、1、2 三个数字,组成一个干净的三分类任务。
python复制import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader, Subset
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
full_train = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
full_test = datasets.MNIST(root="./data", train=False, download=True, transform=transform)
keep_train = [i for i, (x, y) in enumerate(full_train) if y in [0, 1, 2]]
keep_test = [i for i, (x, y) in enumerate(full_test) if y in [0, 1, 2]]
train_subset = Subset(full_train, keep_train)
test_subset = Subset(full_test, keep_test)
# 重新映射标签:0->0, 1->1, 2->2,其实这里刚好不需要改
只取 0、1、2 不算特别费力,但如果你要跑 10 分类全量 MNIST,代码几乎一样,区别只在 keep 里的过滤条件和输出层节点数。实际项目里,数据准备的难点经常在“标签清洗”而不是“代码实现”:类别定义有没有歧义、标注是否一致、哪些样本应该被剔除,这些决策对模型效果的影响远大于你选的网络结构。
4.2 模型搭建:输出层节点数必须等于类别数
多分类模型最后一层的设计很关键。假设我们用一个简单的多层感知机,输入是 28×28 的像素向量,中间接两个全连接层,最后输出的神经元数量必须是 3。
python复制import torch.nn as nn
class SimpleMLP(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Flatten(),
nn.Linear(28 * 28, 128),
nn.ReLU(),
nn.Dropout(0.2),
nn.Linear(128, 64),
nn.ReLU(),
nn.Linear(64, 3)
)
def forward(self, x):
return self.net(x)
这里有个很常见的问题:为什么最后一个线性层后面没有接 Softmax?原因是在 PyTorch 里,nn.CrossEntropyLoss 期望的输入是 logits,也就是网络最后一层出来的原始数值,它会在 loss 函数内部同时完成 Softmax 和交叉熵计算。如果你自己在输出层先加 Softmax,再把结果传给 CrossEntropyLoss,训练时 loss 不仅会更大,而且容易出现数值不稳定的情况,因为两个操作的数值实现没有合并优化。
如果是推理阶段要给别人展示概率,那可以用 torch.softmax(model(x), dim=1) 得到每一类的概率分布,再用 argmax(dim=1) 得到预测类别。训练和推理的写法要分开理解,不要混在一起。
4.3 训练代码与评估:注意 batch 维度和 dtype
训练循环本身不复杂,这里我把关键部分写出来:
python复制import torch.optim as optim
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = SimpleMLP().to(device)
loss_fn = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)
train_loader = DataLoader(train_subset, batch_size=128, shuffle=True)
test_loader = DataLoader(test_subset, batch_size=128, shuffle=False)
for epoch in range(5):
model.train()
total_loss = 0
for images, labels in train_loader:
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
logits = model(images)
loss = loss_fn(logits, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * images.size(0)
# 评估
model.eval()
correct = 0
total = 0
with torch.no_grad():
for images, labels in test_loader:
images, labels = images.to(device), labels.to(device)
logits = model(images)
pred = logits.argmax(dim=1)
correct += (pred == labels).sum().item()
total += labels.size(0)
print(f"epoch {epoch+1}, loss: {total_loss/len(train_subset):.4f}, acc: {correct/total:.4f}")
一个小提醒:CrossEntropyLoss 的 target 必须是整数标签,形状是 [N],不是 [N, 1],也不是 one-hot 的 [N, C]。如果你把 0、1、2 编码成 one-hot 后直接丢进去,会遇到类似“Expected target size [N], got [N, 3]” 的报错。遇到这种问题先别怀疑框架,检查一下标签是否符合 loss 的预期格式。
第二个小提醒:在验证集上算 accuracy 的时候,记得把 model.eval() 打开,并包在 torch.no_grad() 里。前者会关掉 Dropout 和 BatchNorm 的训练行为,后者能省显存并加速计算。没有这两步,你可能会在验证阶段看到忽高忽低的奇怪结果。
4.4 从整体准确率到分指标:打印每类 F1 才能真正定位问题
整体 accuracy 只是第一步。三分类任务里我建议至少打印一份每类 precision、recall、F1 和混淆矩阵:
python复制from sklearn.metrics import classification_report, confusion_matrix
all_preds = []
all_labels = []
model.eval()
with torch.no_grad():
for images, labels in test_loader:
images, labels = images.to(device), labels.to(device)
logits = model(images)
pred = logits.argmax(dim=1)
all_preds.extend(pred.cpu().numpy())
all_labels.extend(labels.cpu().numpy())
print(classification_report(all_labels, all_preds, digits=3))
print(confusion_matrix(all_labels, all_preds))
如果三分类的整体准确率是 0.98,但第二类的 recall 只有 0.90,那说明模型在第二类上有系统性漏判。这时候回去看那些错误样本,往往能发现两类写法在视觉上很相似,或者第二类训练样本数量确实偏少。只有到了这一步,对多分类问题的把控才算完整,而不是“训练完看个 loss 就结束了”。
5. 常见问题与排查技巧实录
5.1 网络输出层节点数和类别数对不上
这是新手最容易犯的问题。模型最后一层输出 2 个神经元,标签却有 3 个类,或者输出 10 个节点但任务只有 3 类。CrossEntropyLoss 在计算 loss 时,会把 logits 的第二个维度和标签的最大值做比较,如果标签里有 2、维度只有 2,会直接报 index out of range。
排查办法很简单:打印 model(x).shape 和 labels.unique(),两者必须满足“logits 第二维 >= 标签最大值 + 1”。输出节点数量不是越大越好,多出来的节点只会增加参数和混淆度。
5.2 输出层用了 Sigmoid,概率总和不为 1,预测却还正常
有些同学在输出层用了 Sigmoid 而不是 Softmax,发现每个输出都落在 0 到 1 之间,用 argmax 也能得到预测类别,整体准确率甚至不低。这会让初学者误以为 Sigmoid 也能做多分类。
问题在于,Sigmoid 对每个输出独立做映射,没有让 K 个输出互相校准。它告诉你的只是“每个类别单独为正的可能性”,而不是“在互斥类别中某一类胜出的概率”。如果类别是互斥的,严格来说还是应该用 Softmax,否则模型输出的概率连基本的归一化约束都不满足,后续要想做阈值调整、概率校准、不确定性估计时,整个口径都是乱的。
5.3 训练 loss 一直降,但验证准确率不涨甚至下降
常见解释是过拟合,但多分类场景里还要考虑是不是标签噪声或者某些类别样本太少。如果模型在训练集上记住了少数异常样本,验证集上遇到真实分布就露馅。
对策方面,我一般会先看训练集和验证集的 loss 曲线差距,如果训练 loss 非常低、验证 loss 明显高,考虑加 Dropout、加正则、做数据增强。如果验证 loss 也降但 accuracy 不涨,要看是不是类别不平衡导致整体 loss 被多数类主导,或者评估指标本身不适合当前任务。
5.4 混淆矩阵展示某个类别错分成“错误类别”非常集中
这不是偶然,通常是两个类别在特征空间中确实存在混淆区域。比如手写数字里 1 和 7 容易互相误判,文本分类里“投诉”和“售后咨询”的边界模糊。
处理这类问题的常用手段:一是增加容易混淆类别的样本,尤其是那些介于两者之间的难例;二是修正标签定义,把难以区分的类别重新梳理,必要时引入“其他”兜底类,在业务上反而更实用;三是增加专门区分这两个类的特征通道,比如图像里裁剪更精细的区域、文本里抽取更贴近业务语义的特征。
5.5 多分类任务最常见的排查速查表
| 症状 | 可能原因 | 处理建议 |
|---|---|---|
| 训练时报 index out of range | 输出层节点数小于类别数 | 检查类别数,保证输出层节点数等于类别数 |
| loss 报错 size mismatch | 标签传了 one-hot 给 CrossEntropyLoss | 转成整数索引标签,或换 NLLLoss |
| 准确率很高但少数类 F1 接近 0 | 类别不平衡且未处理 | 加 class weight,重采样,或换 Focal Loss |
| 验证集 loss 低但指标差 | 评估指标与业务目标不匹配 | 改用 macro-F1、Kappa 或业务自定义代价指标 |
| 两类样本总互相误判 | 特征空间重叠或标签定义模糊 | 检查 bad case,补充难例,必要时调整类别体系 |
| 训练时概率总和大于 1 | 输出层用了 Sigmoid 而非 Softmax | 互斥分类任务输出层改 Softmax |
| 相同代码不同机器上 acc 不稳定 | 随机种子未固定 | 设置 torch.manual_seed、numpy random seed |
这些小问题,我在教学和实际项目里几乎都遇到过。其中印象最深的一次,是某个分类任务看起来效果很好,直到打印每类 F1 才发现跟业务最相关的少数类几乎全是漏检,当时所有人都在看整体准确率,没有一个人去拆混淆矩阵。从那次以后,我给自己定了一个规矩:凡是分类项目,第一版评估必须带混淆矩阵和分指标,不带这两个东西的 report 一律不看。
多分类问题本身并不复杂,但它把“建模、训练、评估、优化”这条链路串得很完整,是值得反复打磨的一块基石。我习惯在跑一个新的表格分类任务时,先用一个简单的 MLP 或逻辑回归做 baseline,从混淆矩阵里找到弱类,再逐步加模型复杂度。顺序反过来的话,花了一堆时间调网络结构,最后发现数据标签本身有问题,心态很容易崩。
复盘完 Day18 这期内容,我最想强调的还是“指标口径”这件事。Softmax 和交叉熵是工具层面的东西,照着文档写就能用;但能不能发现模型在少数类上表现差、能不能识别出两个常被混淆的类别,才真正决定你的多分类模型能不能在业务里扛得住。如果你正被某一个多分类任务卡住,建议先别急着换模型,把混淆矩阵打出来看一眼,答案往往就藏在里面。
