1. 从二分类到多分类:为什么这个问题值得单独拿出来讲
很多人学神经网络,第一个能跑通的 demo 大概率是猫狗二分类——一个 sigmoid 输出,阈值 0.5 一刀切,准确率还挺好看。但真正到了实际项目里,你面对的几乎不可能是“只有两类”的干净问题。手写数字识别是 10 类,垃圾邮件分类可能是几十个标签,商品图像归类动辄几百上千个细分类目。这时候如果还拿二分类那套思路硬套,模型要么训不动,要么训出来一堆“看起来收敛了但实际全预测成同一类”的废模型。
图像多分类,说白了就是让网络从一张图里判断它属于哪一类,而且类别数大于 2。它和二分类的本质区别不在于“多了一个输出”,而在于输出层的表达方式、损失函数的选择、标签的编码形式、评估指标的口径全都变了。我见过太多人卡在“loss 一直不降”“准确率永远等于 1/类别数”这种坑里,根子上就是把多分类当二分类在做。
这篇内容我打算按一个完整项目的推进顺序来讲:先讲清楚多分类的整体设计思路和方案选型,再拆核心细节和实操要点,然后给一套能直接抄的完整实现流程,最后把我自己踩过的坑和排查经验整理成速查表。适合已经跑通过二分类、想往真实多分类任务进阶的人,也适合带学生做项目的老师拿来当参考。读完你应该能独立搭一个图像多分类的完整 pipeline,并且知道每一步为什么这么做。
2. 图像多分类的整体设计与方案选型
2.1 多分类任务的本质:从“是不是”到“是哪一个”
二分类回答的是“是不是”的问题,输出一个概率值,大于 0.5 就是正类。多分类回答的是“是哪一个”的问题,输出的是一个概率分布——所有类别的概率加起来等于 1,谁最大就判给谁。
这个转变带来一个关键约束:各类别之间是互斥的。一张图要么是猫要么是狗,不可能同时是猫又是狗(多标签任务另说,那是另一个话题)。正因为互斥,我们才能用 softmax 把原始得分压成一个合法的概率分布。如果你把多分类当成“每个类别独立做一次二分类”,用一堆 sigmoid,那各类概率加起来可能大于 1,逻辑上就矛盾了,训练信号也会互相打架。
所以方案选型的第一条铁律:互斥多分类用 softmax + 交叉熵,多标签才用 sigmoid + 二元交叉熵。这个判断决定了后面所有代码的写法。
2.2 主干网络怎么选:别一上来就上最深的
图像多分类的主干网络选择,我的经验是按数据量和算力分档:
| 数据规模 | 推荐主干 | 理由 |
|---|---|---|
| 几千张以内 | 预训练 ResNet18 / MobileNetV3 | 参数少、不易过拟合、训练快 |
| 几万到几十万张 | 预训练 ResNet50 / EfficientNet-B0 | 容量够、迁移效果好 |
| 百万级以上 | ConvNeXt / ViT 系列 | 大数据下才能发挥容量优势 |
这里有个反直觉的点:数据少的时候,越大的模型越差。我试过在 3000 张图的数据集上直接上 ResNet152,训练集准确率能到 99%,验证集死活卡在 60% 上下,典型的过拟合。换成 ResNet18 加冻结主干,验证集直接上到 85%。原因很简单,大模型的参数量远超数据能提供的约束,它记住训练集比学会泛化容易得多。
另一个关键决策是要不要用预训练权重。除非你的数据是某种极其特殊的领域图像(比如工业缺陷图、医学影像),否则一律用 ImageNet 预训练权重起步。预训练模型已经学会了边缘、纹理、形状这些通用特征,你只需要微调高层语义部分。实测下来,用预训练比从零训练,收敛速度快 3 到 5 倍,最终精度通常也高 5 到 15 个百分点。
2.3 损失函数与标签编码:交叉熵背后的门道
多分类的标准损失是交叉熵损失(CrossEntropyLoss)。它的形式是:
Loss = -log(softmax(logits)[真实类别])翻译成人话:先看模型给真实类别打了多少概率,概率越高损失越小,取个负对数放大惩罚。如果模型给真实类别只打了 0.1 的概率,损失就是 -log(0.1) ≈ 2.3;如果打了 0.9,损失就是 -log(0.9) ≈ 0.1。这个设计让模型有强烈动机把真实类别的概率往上推。
标签编码有两种方式,必须搞清楚:
- 整数标签(sparse):标签就是 0, 1, 2, ..., N-1。PyTorch 的
nn.CrossEntropyLoss默认吃这种,内部自动做 softmax + log。 - 独热编码(one-hot):标签是一个长度 N 的向量,真实类别位置是 1,其余是 0。用
nn.BCEWithLogitsLoss或nn.NLLLoss配合手动 log_softmax 时用。
注意:PyTorch 的
CrossEntropyLoss已经把 softmax 和 log 合并在一起了,所以你的模型输出层不要再加 softmax。加了就是重复,会导致梯度异常、loss 不降。这是新手最常犯的错之一。
2.4 评估指标:准确率不是万能的
多分类的评估指标比二分类丰富得多,选错了会严重误判模型好坏:
- Top-1 准确率:预测概率最高的类别就是真实类别的比例。最常用。
- Top-5 准确率:真实类别出现在概率前 5 的比例。类别多的时候(比如 1000 类)更有参考价值。
- 混淆矩阵:看哪些类别容易被搞混。比如“哈士奇”和“阿拉斯加”经常互错,说明特征区分度不够。
- 每类精确率/召回率/F1:类别不平衡时必看。整体准确率 90% 但某个小类召回率只有 30%,模型其实是有问题的。
我一般训练时盯 Top-1 准确率,验证时看混淆矩阵和每类 F1,上线前再算一遍 Top-5。这套组合能覆盖绝大多数问题。
3. 核心细节解析与实操要点
3.1 数据准备:多分类的成败一半在这里
多分类任务对数据的要求比二分类苛刻得多,因为类别一多,类别不平衡和类间相似这两个问题会被放大。
先说类别不平衡。假设你有 10 个类,其中 9 个类各有 1000 张图,剩下 1 个类只有 50 张。模型很快会发现“把所有图都预测成那 9 个大类”就能拿到 90% 以上的准确率,于是那个小类永远学不会。解决办法有几个,我按推荐顺序排:
- 重采样:对小类过采样(复制+增强),对大类欠采样。简单粗暴但有效。
- 类别权重:给
CrossEntropyLoss传weight参数,小类权重调高。计算方式是weight_i = 总样本数 / (类别数 * 第i类样本数)。 - 数据增强:对小类做更强的增强(旋转、裁剪、颜色抖动),变相增加样本多样性。
再说类间相似。有些类别天生难分,比如不同型号的螺丝、不同品种的花。这时候光靠数据量堆没用,得靠更细粒度的特征。我的做法是:先用预训练模型跑一版 baseline,看混淆矩阵找出最容易混的类别对,然后针对性地对这些类别做数据补充或引入注意力机制。
数据划分上,训练/验证/测试一般按 7:1.5:1.5 或 8:1:1。验证集和测试集的类别分布必须和真实场景一致,不能随机切完发现某个类全跑到训练集去了。
3.2 输出层设计:类别数决定一切
输出层的维度必须等于类别数 N。假设主干网络输出的特征维度是 512,那最后一层就是nn.Linear(512, N)。这个 N 是硬约束,写错了模型直接报错或者静默出错。
有个细节很多人忽略:输出层的初始化。默认的随机初始化在类别多的时候可能导致初始 logits 方差过大,softmax 后概率极端化,训练初期梯度不稳。我的习惯是用小方差初始化,比如nn.init.normal_(fc.weight, std=0.01),让初始输出接近均匀分布,训练更平滑。
另外,如果用了预训练主干,输出层一定是随机初始化的(因为 ImageNet 是 1000 类,你的类别数大概率不是 1000)。这时候训练策略要分两阶段:先冻结主干只训输出层几个 epoch,让随机初始化的输出层先稳定下来,再解冻主干整体微调。直接一起训容易把预训练学到的特征带崩。
3.3 训练策略:学习率、批次与轮次
多分类的训练策略有几个经验值可以直接参考:
- 学习率:微调阶段主干用 1e-4 到 1e-5,输出层用 1e-3。用 Adam 或 AdamW 优化器。如果从零训练,可以上 1e-2 配 SGD + momentum。
- 批次大小:受显存限制,一般 32 或 64。批次太小(比如 4)会导致 BatchNorm 统计不稳,批次太大泛化可能变差。
- 训练轮次:配合早停(early stopping),验证集准确率连续 5 到 10 个 epoch 不提升就停。别硬训到几百轮,过拟合了精度反而掉。
- 学习率调度:CosineAnnealing 或 ReduceLROnPlateau 都行。我偏好 Cosine,平滑下降,不用调 patience。
这里有个我踩过的坑:冻结主干时,BatchNorm 层的行为。如果主干里有 BatchNorm,冻结时它的 running_mean 和 running_var 还在更新,可能和新数据分布不匹配。稳妥做法是冻结阶段把 BatchNorm 也设成 eval 模式,或者干脆用 GroupNorm 的主干。
3.4 数据增强:多分类的隐形加速器
数据增强在多分类里不是可选项,是必选项。因为类别多,每个类的样本相对就少,不增强很容易过拟合。
基础增强组合(几乎万能):随机水平翻转、随机裁剪(从原图裁 0.8 到 1.0 区域再缩放回原尺寸)、颜色抖动(亮度、对比度、饱和度各 ±0.2)。这套组合对自然图像效果稳定。
进阶增强(数据少时用):RandAugment、Mixup、CutMix。Mixup 是把两张图按比例混合、标签也按比例混合,能显著提升泛化。CutMix 是裁剪一块区域贴到另一张图上,标签按面积比例混合。这两个我实测在 1 万张以下的数据集上能提 3 到 8 个百分点。
注意:验证集和测试集绝对不能做随机增强,只能做确定性的 resize 和中心裁剪。否则每次评估结果都不一样,没法比较。
4. 完整实操流程:从零搭一个图像多分类 pipeline
4.1 环境与依赖准备
我用的是 PyTorch 生态,版本建议 torch 2.x 配 torchvision 0.15+。核心依赖就三个:torch、torchvision、numpy。可视化用 matplotlib,进度条用 tqdm。不需要额外装什么花哨的库,多分类本身不复杂,复杂的是数据和调参。
pip install torch torchvision numpy matplotlib tqdm数据集组织成 ImageFolder 格式最省事:
dataset/ train/ class_0/ xxx.jpg ... class_1/ xxx.jpg ... ... val/ class_0/ ... class_1/ ...每个类别一个文件夹,文件夹名就是类别名。torchvision.datasets.ImageFolder会自动扫描并生成整数标签,省去手写 Dataset 的麻烦。
4.2 数据加载与增强代码
import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms train_tf = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.2, 0.2, 0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) val_tf = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) train_ds = datasets.ImageFolder('dataset/train', transform=train_tf) val_ds = datasets.ImageFolder('dataset/val', transform=val_tf) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4) val_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=4) num_classes = len(train_ds.classes) print('类别数:', num_classes, '类别列表:', train_ds.classes)这里的 mean 和 std 是 ImageNet 的统计值,用预训练模型时必须用这套归一化,否则输入分布和预训练时不匹配,效果会打折。RandomResizedCrop的 scale 设 0.8 到 1.0 是保守做法,数据少可以放宽到 0.6。
4.3 模型构建与两阶段训练
import torch.nn as nn from torchvision import models def build_model(num_classes, freeze_backbone=True): model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) if freeze_backbone: for p in model.parameters(): p.requires_grad = False in_features = model.fc.in_features model.fc = nn.Linear(in_features, num_classes) nn.init.normal_(model.fc.weight, std=0.01) nn.init.zeros_(model.fc.bias) return model model = build_model(num_classes, freeze_backbone=True).cuda() criterion = nn.CrossEntropyLoss()第一阶段只训输出层:
optimizer = torch.optim.AdamW(model.fc.parameters(), lr=1e-3) for epoch in range(5): model.train() for imgs, labels in train_loader: imgs, labels = imgs.cuda(), labels.cuda() optimizer.zero_grad() out = model(imgs) loss = criterion(out, labels) loss.backward() optimizer.step() print(f'[冻结阶段] epoch {epoch} loss {loss.item():.4f}')第二阶段解冻主干,用更小的学习率整体微调:
for p in model.parameters(): p.requires_grad = True optimizer = torch.optim.AdamW([ {'params': model.fc.parameters(), 'lr': 1e-3}, {'params': [p for n, p in model.named_parameters() if 'fc' not in n], 'lr': 1e-4}, ]) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=30) best_acc = 0.0 for epoch in range(30): model.train() for imgs, labels in train_loader: imgs, labels = imgs.cuda(), labels.cuda() optimizer.zero_grad() out = model(imgs) loss = criterion(out, labels) loss.backward() optimizer.step() scheduler.step() model.eval() correct = total = 0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels = imgs.cuda(), labels.cuda() pred = model(imgs).argmax(dim=1) correct += (pred == labels).sum().item() total += labels.size(0) acc = correct / total print(f'[微调阶段] epoch {epoch} val_acc {acc:.4f}') if acc > best_acc: best_acc = acc torch.save(model.state_dict(), 'best.pth')两阶段的关键在于参数分组学习率:输出层是随机初始化的,需要大学习率快速收敛;主干是预训练的,只需要小学习率微调。用同一个学习率要么输出层学得慢,要么主干被带崩。
4.4 评估与混淆矩阵
训练完别只看一个准确率数字,把混淆矩阵画出来:
import numpy as np from sklearn.metrics import confusion_matrix, classification_report model.eval() all_preds, all_labels = [], [] with torch.no_grad(): for imgs, labels in val_loader: imgs = imgs.cuda() pred = model(imgs).argmax(dim=1).cpu().numpy() all_preds.extend(pred) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, target_names=train_ds.classes))classification_report会给出每类的精确率、召回率、F1。如果某个类 F1 明显低,回去看它的样本量和增强策略。混淆矩阵里如果发现两个类互相错得特别多,说明它们的特征在模型眼里太像,要么补数据,要么考虑加细粒度分类头。
5. 常见问题与排查技巧实录
5.1 训练不收敛的典型症状与对策
多分类训练出问题,症状就那么几种,对症下药很快能定位:
| 症状 | 可能原因 | 排查与解决 |
|---|---|---|
| loss 不降,卡在 log(N) 附近 | 输出层加了 softmax 导致重复 | 去掉模型里的 softmax,直接用 logits 喂 CrossEntropyLoss |
| 准确率恒等于 1/N | 模型全预测同一类 | 看混淆矩阵,检查类别是否严重不平衡,加类别权重 |
| loss 震荡剧烈 | 学习率太大 | 降学习率,或换 AdamW 加 warmup |
| 训练准确率高但验证低 | 过拟合 | 加数据增强、加 dropout、减小模型、早停 |
| loss 变 NaN | 学习率过大或数据有脏样本 | 降学习率,检查数据里有没有损坏图片或标签越界 |
我遇到最多的是第一种和第二种。第一种纯粹是代码写错,第二种往往是数据问题。排查顺序永远是:先确认代码逻辑对,再看数据分布,最后才调超参。
5.2 类别不平衡的实战处理
前面提了重采样和类别权重,这里给个具体的权重计算例子。假设 5 个类的样本数分别是 [1000, 1000, 1000, 1000, 50]:
counts = np.array([1000, 1000, 1000, 1000, 50]) weights = len(counts) / (len(counts) * counts) # 每类权重 weights = weights / weights.sum() * len(counts) # 归一化 weights = torch.tensor(weights, dtype=torch.float32).cuda() criterion = nn.CrossEntropyLoss(weight=weights)算下来第 5 类的权重会是其他类的 20 倍左右,模型预测错第 5 类会被重罚,从而被迫去学它。但权重也别调太极端,否则模型会矫枉过正,把其他类都往第 5 类预测。我一般把权重上限控制在 10 倍以内,超出的部分靠数据增强补。
5.3 显存不够与训练太慢的优化
显存不够是多分类常见问题,因为类别多往往意味着数据量大。几个立竿见影的优化:
- 混合精度训练:用
torch.cuda.amp,显存直接省 30% 到 50%,速度还快。代码就几行,收益极高。 - 梯度累积:批次设小一点,累积几个批次的梯度再更新一次,等效于大批次但显存占用小。
- 冻结主干:第一阶段只训输出层,显存占用大幅下降。
- 减小输入尺寸:224 降到 160 或 128,显存和速度都改善,精度损失通常 1 到 3 个点。
混合精度的写法:
scaler = torch.cuda.amp.GradScaler() for imgs, labels in train_loader: imgs, labels = imgs.cuda(), labels.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): out = model(imgs) loss = criterion(out, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()这套下来,同样的显卡能训更大的批次或更大的模型,实测训练速度提升 40% 以上。
5.4 上线前的最后检查清单
模型训完别急着部署,过一遍这个清单:
- 测试集评估:用完全没参与训练和调参的测试集跑一遍,确认精度和验证集接近。差太多说明过拟合了验证集。
- 单张推理测试:拿几张真实场景的图,走一遍预处理+推理+后处理全流程,确认没有维度或归一化错误。
- 类别映射保存:把
train_ds.classes存下来,推理时用同一个映射把输出索引翻译成类别名。顺序错了结果全错。 - 边界情况:全黑图、全白图、尺寸异常的图,看模型输出是否合理,会不会崩。
- 推理速度:测一下单张耗时,确认满足业务要求。ResNet18 在 GPU 上单张通常 5 到 10 毫秒。
我个人在实际操作中的体会是,多分类项目 80% 的时间花在数据上,20% 花在模型和调参上。数据干净、分布合理、增强到位,用一个标准的 ResNet 加交叉熵就能拿到不错的结果。反过来,数据一团糟,再花哨的模型也救不回来。所以每次项目启动,我都会先花半天时间把数据过一遍——看类别分布、看有没有标错的、看图像质量是否一致。这一步偷懒,后面调参调到怀疑人生。
最后再分享一个小技巧:训练时把每个 epoch 的验证集混淆矩阵存成图片,按 epoch 编号。训完之后翻一遍,你能直观看到模型是怎么一步步把容易混的类别区分开的,对理解模型行为特别有帮助。这个习惯我保持了几年,比看 loss 曲线有用得多。