多分类任务是绝大多数人从“会调库”走向“真做模型”的第一道坎。二分类做得很顺的人,第一次面对5个、10个、几十个类别时,通常会踩同一个节奏:模型能跑通,准确率也看着不差,但一上混淆矩阵就发现某些类别几乎全军覆没,整体准确率却被那些“好分”的类别拉高了。这篇文章从多分类的建模决策讲到评估环节,把损失函数、标签编码、类别不平衡和混淆矩阵串起来讲。核心代码基于Python和scikit-learn,用的示例是手写数字识别,但同样的思路可以直接迁移到文本分类、图像分类、故障诊断这类任务上。文章不会停在“知道”层面,而是把每一步为什么这样做讲清楚,读完你手里会有一套可以直接改用的代码,也知道怎么从混淆矩阵里定位问题。
1. 多分类问题:概念边界,以及二分类经验为什么不能直接平移
1.1 先分清多分类与多标签
多分类(Multi-class Classification)指类别互斥、且类别数K≥3的问题。一个样本只能落到其中一个类别里:手写数字0到9,一张图只能对应一个数字;新闻标题只能属于财经、体育、娱乐中的一个板块。这种“互斥性”是softmax建模的基础。
我见过不少新人把“多标签”和“多分类”混在一起。多标签问题里,一个样本可以同时属于多个类别,比如一张照片同时有猫和狗,一篇工单同时涉及“退款”和“投诉”。这种任务不能用单个softmax输出层来解决,因为softmax强制所有概率之和为1,样本属于“退款”的概率升高,属于“投诉”的概率必然被压下去,但真实场景里两者完全可以同时成立。
所以接到需求的第一件事,不是急着调包,而是确认:类别之间是否互斥?如果不互斥,后面所有评估逻辑都要换一套。
1.2 类别数增加,复杂度不是线性增长,而是边界数量暴涨
二分类只需要在特征空间里划出一条决策边界,把空间分成两块。三分类需要至少两个边界,四分类、十分类的边界组合会成倍增加。更现实的问题是:在样本总量不变的情况下,类别从2个变成10个,每个类别平均分到的样本量直接砍掉五分之一。类别一多,少数类别的样本就会稀疏,模型很难学到稳定的边界。
我在实际项目里体会最深的一件事是:二分类里“正负样本不均衡”很好处理,无非是换阈值、调权重;多分类里不均衡却是结构性的。十个类别,有的类别样本占40%,有的只占3%,模型天然倾向于把模糊样本判到大类里去,因为这样整体损失最小。这个现象不是调参能完全消除的,必须靠评估环节把它暴露出来。
1.3 “一对一 vs 一对多”的策略选择要心里有数
很多从二分类过来的人,习惯把多分类拆成“某个类别 vs 其余所有类别”来做,这就是一对多(One-vs-Rest)。sklearn里的很多模型默认就是OvR策略。这样做的缺点是:每个分类器只看到“自己类 vs 全世界”,类别之间的细微差异被忽略了。
还有一对一的策略(One-vs-One),每两个类别训练一个分类器,K个类别就要训练K*(K-1)/2个模型。类别数少时还行,20个类别就要训练190个分类器,训练和推理成本都不低。
如果你用的模型本身支持多项分布(Multinomial),比如逻辑回归配lbfgs求解器、神经网络配softmax,它们是在一个模型里同时优化所有类别的边界,类间信息可以共享。这也是为什么深度学习做多分类基本都用softmax,而不是拆成无数个二分类。
2. 建模前的三个关键决策:输出层、损失函数、评估指标
2.1 softmax + 交叉熵:为什么是默认组合
多分类默认的输出层是softmax,损失函数默认是交叉熵。softmax做的事情,是把模型最后一层输出的K个实数,转换成一个概率分布:每个数都被压缩到0到1之间,且所有数加起来等于1。这个“和为1”的性质,就是模型对“互斥类别”的数学表达。
交叉熵衡量的是预测分布和真实分布之间的距离。真实标签如果是类别3,真实分布就是“类别3概率为1,其他为0”,交叉熵会惩罚模型没有给类别3足够高的概率。
我见过有人问:为什么不用均方误差?从数学上看,交叉熵配合softmax,梯度形式更干净,不会出现输出饱和时梯度消失的问题。从直觉上看,分类任务关心的是“概率分对没有”,不是“数值回归得准不准”,交叉熵天然更匹配。所以除非有特殊理由,多分类的默认组合就是softmax加交叉熵,不要自己发明奇怪的组合。
2.2 标签编码:one-hot和稀疏索引的差别
多分类的标签有两种常见编码方式。一种是one-hot:把“类别3”编码成[0,0,1,0,...],长度和类别数一致。另一种是稀疏整数索引,直接存一个整数3。
两种方式在信息上是等价的,但对应的损失函数不同。在TensorFlow/Keras里,one-hot配CategoricalCrossentropy,整数标签配SparseCategoricalCrossentropy;在PyTorch里,CrossEntropyLoss自带softmax,而且要求整数标签。用错编码方式最常见的报错就是维度不匹配,模型输出是(batch, K),标签却是(batch,)或者反过来。
scikit-learn的优势是不需要你自己编码,模型内部会自动处理。但有一点很容易被忽略:model.classes_保存的是训练时见过的类别列表,而且按排序输出。预测时predict_proba的每一列顺序,就是classes_的顺序。如果你自己拼标签、自己画混淆矩阵,一定要用这个顺序对齐,否则画出来的矩阵行列是错位的,但你还看不出来。
2.3 评估指标:准确率、宏平均、微平均怎么选
多分类的评估,单看准确率一定会出问题。假设10个类别里有1个类别占了90%样本,模型全部预测成这个类别,准确率是90%,看起来不错,但这个模型实际上一无是处。
多分类评估里常用的三个平均方式:
- 宏平均(macro):先算每个类别的precision、recall、F1,再取算术平均。每个类别权重相同,不受样本量影响,最能暴露少数类的问题。
- 微平均(micro):把所有类别的TP、FP、FN加起来,再算F1。它等价于全局准确率,受大类影响大。
- 加权平均(weighted):按每个类别的样本量加权,适合业务中类别分布相对固定、希望在现有分布下评估性能的情况。
我在项目里基本是宏平均和加权平均一起看。宏平均低、加权平均高,说明模型牺牲了少数类;两者差距不大,才说明模型在所有类别上还算均衡。
3. 一份可以直接改用的多分类训练与评估代码
3.1 数据准备:stratify切分为什么是必须的
直接开始写代码。先加载数据、划分训练集和测试集:
from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split X, y = load_digits(return_X_y=True) print(X.shape, y.shape) # (1797, 64) (1797,) print(set(y)) # {0,1,2,...,9} X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.25, random_state=42, stratify=y )stratify=y这一步不是可选项。如果不做分层抽样,类别分布会在训练集和测试集之间随机浮动,尤其是少数类,可能训练集里只有三五个,测试集里却有几十个,评估结果完全不靠谱。sklearn的train_test_split默认不做分层,这点必须手动指定。
3.2 训练模型:lbfgs逻辑回归的多分类行为
from sklearn.linear_model import LogisticRegression model = LogisticRegression(max_iter=5000, solver='lbfgs') model.fit(X_train, y_train)逻辑回归默认支持多分类。solver='lbfgs'在多分类场景下走的是multinomial路径,也就是直接用一个softmax模型同时学习所有类别的边界,而不是拆成多个二分类。这样类别之间可以共享信息,对10类数字识别这种任务效果更好。
如果你用的是SVM、随机森林这类模型,思路也完全一样:训练完拿到预测结果,评估流程是一样的。但注意SVM本身不输出概率,或者概率校准需要额外calibrate,这会直接影响你对可信度的判断。
y_pred = model.predict(X_test) y_prob = model.predict_proba(X_test) print(y_prob.shape) # (450, 10) print(model.classes_) # [0 1 2 3 4 5 6 7 8 9] print(y_prob[0].sum()) # 1.0predict_proba返回的是(n_samples, n_classes)矩阵,每行概率和为1。每个位置对应model.classes_里的一个类别,不是按你脑子里的数字顺序,而是按classes_排序后的顺序。如果训练时类别5、7、9,classes_就是[5,7,9],三列分别对应这三个类别。
3.3 混淆矩阵:三行代码出的图和一堆信息量
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt labels = sorted(set(y_test) | set(y_pred)) cm = confusion_matrix(y_test, y_pred, labels=labels) disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=labels) disp.plot(cmap='Blues', values_format='d') plt.title("Multi-class Confusion Matrix") plt.show()这里传入labels参数是为了确保矩阵的行列顺序一致。如果不传,sklearn默认会取y_test和y_pred的并集并排序,结果大概率一样,但如果你手动过滤过某些样本,或者预测结果里恰好缺少某个类别,矩阵维度就会对不上。显式传入labels是最稳的写法。
如果不想用sklearn的绘图,也可以直接用seaborn画,自由度更高:
import seaborn as sns plt.figure(figsize=(10, 8)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=labels, yticklabels=labels) plt.xlabel("Predicted") plt.ylabel("True") plt.show()3.4 分类报告:一次读完所有类别的precision/recall/F1
from sklearn.metrics import classification_report print(classification_report(y_test, y_pred, labels=labels, zero_division=0))输出结果里每一行是一个类别,包含precision、recall、F1和support(该类别的真实样本数)。最后几行是accuracy、macro avg、weighted avg。我看报告的习惯是直接从下往上读:先看macro avg和weighted avg的差距,再回到每一行找那些F1明显偏低的类别。
4. 多分类混淆矩阵的真正用法
4.1 看懂每个格子:行、列、对角线
多分类混淆矩阵是一个K×K的矩阵,行代表真实类别,列代表预测类别。cm[i][j]表示“真实类别是第i类,却被预测成第j类”的样本数量。对角线上的值越大越好,因为那代表预测正确。
这里很容易有一个错觉:对角线数字大就万事大吉。不对。对角线数字大只能说明这个类别绝对错误少,但如果有两个类别形状很像,模型很可能把其中一类大量预测成另一类。这个信息在整体的准确率里看不出来,在混淆矩阵里却一目了然。
每一行所有格子加起来,等于这个类别在测试集里的真实样本数,也就是classification_report里的support。每一列所有格子加起来,等于模型总共预测成这个类别的数量。行求和与列求和之间的差异,能直接看出模型系统性地高估了哪个类别、低估了哪个类别。
4.2 归一化视角:行归一化比绝对值更有用
绝对数量的混淆矩阵受样本量影响太大。10个类别的测试集,样本多的类别可能有100个,少的只有10个。同样是15个错误,“大类的15个错误”和“小类的15个错误”严重程度完全不同。这时候必须归一化。
cm_norm = confusion_matrix(y_test, y_pred, labels=labels, normalize='true')normalize='true'表示按行归一化,每行除以该行总和。这样每个格子的值代表“真实类别i中,有多大比例被预测成了j”,也就是该类别的召回分布。对角线就是每个类别的recall。
如果你的sklearn版本较老,不支持normalize参数,可以手动实现:
cm_norm = cm.astype('float') / cm.sum(axis=1, keepdims=True)注意axis=1是按行求和,keepdims=True是为了保持维度,方便广播运算。没有keepdims的话,除法会得到错误结果,这个细节坑过很多人。
画归一化混淆矩阵时,建议把fmt改成小数点格式:
sns.heatmap(cm_norm, annot=True, fmt='.2f', cmap='Blues', xticklabels=labels, yticklabels=labels) plt.xlabel("Predicted") plt.ylabel("True") plt.show()行归一化适合看“哪个类别的样本被分错了”;如果你想看“预测成某个类别的样本里,究竟有哪些真实类别”,就按列归一化,normalize='pred'。两种视图解决不同的问题,不是随便选一个就行的。
4.3 自动定位最易混淆的类别对
混淆矩阵一大,肉眼扫一遍很容易漏。写段代码自动找出混淆最严重的类别对:
import numpy as np cm_no_diag = cm.copy() np.fill_diagonal(cm_no_diag, 0) row_idx, col_idx = np.unravel_index(np.argmax(cm_no_diag), cm_no_diag.shape) print(f"最大混淆:真实类别 {labels[row_idx]} -> 预测为 {labels[col_idx]},共 {cm_no_diag[row_idx, col_idx]} 个样本")np.argmax返回的是展平后的最大索引,np.unravel_index把它还原成二维坐标。找出来之后,你可以单独抽这两个类别的样本,看看它们到底差在哪。很多时候,你会发现自己提取的特征本身就不足以区分这两个类别。
5. 从评估结果反推优化方向
5.1 用分类报告锁定薄弱类别
classification_report里F1最低的那一行,就是最值得投入精力的类别。先别急着调模型结构,先看support:如果support很小,比如只有5个样本,那F1低可能只是噪声;如果support有50个,F1仍然低,说明模型确实学不会这个类别。
我习惯把宏平均F1和每个类别的F1列成一张表,按F1升序排序,一次看哪些类别排在末尾。多分类优化的优先级永远是“最短的那块板”,不是整体指标。
5.2 用归一化混淆矩阵识别“被吸走”的样本
看行归一化矩阵里,除了对角线之外,哪个格子的值最高。比如第3行第5列是0.4,说明真实类别3的样本有40%被预测成了类别5。“类别5吸走了类别3的样本”,这种系统性的偏移通常有两个原因:一是两个类别在特征空间里确实重叠,二是类别5的样本量远大于类别3,模型偏向大类。
如果是后者,解决办法不是加数据,至少不是只加小类数据,而是适当增加小类样本的权重,或者用class_weight='balanced'重训一版,看看混淆矩阵有没有变化。
5.3 可落地的优化方案:样本、特征、权重、后处理
多分类优化的路径,按投入产出比排序大概是:
- 针对高混淆类别对补充样本,或者做数据增强,让模型看到更多边界样本;
- 增加特征,尤其是能区分高混淆类别对的特征;
- 给少数类加权,或者调整损失函数里的类别权重;
- 训练完走阈值或后处理规则,比如对某些业务类别做二次判断。
我在文本分类项目里试过最有效的一招:找出混淆矩阵中最大的混淆对之后,单独训练一个二分类器,专门判别“是A还是B”,只在模型对这两类的初始概率接近时触发。这样整体准确率提升1到2个百分点,成本却很低。
6. 多分类实战高频踩坑记录
6.1 类别不平衡带来的Accuracy Paradox
准确率很高但模型实际没用,这是多分类最容易踩的坑。测试集里90%是类别A,模型全预测成A,准确率90%,但类别B和C全错。Accuracy Paradox指的就是这个现象:模型越“懒”,准确率反而可能越高,尤其是类别分布极度倾斜的时候。
所以多分类项目里,我几乎不看单独打印的model.score,只看classification_report和混淆矩阵。如果macro avg远低于weighted avg,说明大类在掩盖小类的失败。
6.2 标签顺序不一致导致混淆矩阵错位
这是最隐蔽的bug:代码跑出来混淆矩阵对角线也正常,但就是感觉哪里不对。最常见原因是训练时做了标签编码,比如把字符串类别映射成0、1、2,但映射表的顺序和测试时的映射顺序不一致,导致真实标签和预测标签对不上。
避免方法很简单:所有标签编码、解码、预测、评估,全部使用同一个labels变量,不要在不同的地方各自sort一遍。我写代码时会显式维护一个label_list,训练、预测、画混淆矩阵都传它。
6.3 测试集出现训练集没见过的类别
多分类模型只能输出训练时见过的类别。如果线上新出现了一个训练集里完全没有的新类别,模型不会说“我不知道”,而是会强行把它分到已有的某个类别里,而且往往概率还不低。这在业务中很常见:新工单类型、新故障码、新商品类目。
处理思路有两个:一是定期用线上数据回流重训;二是在预测概率上做卡阈值,所有类别概率都低于阈值时,判为“未知”,进入人工复核。虽然sklearn里的predict没有这个功能,但用predict_proba配合np.max很容易实现。
6.4 多标签问题被当成多分类问题
前面提到过,多标签的样本可以同时属于多个类别。强行用softmax,模型只能选一个概率最大的类别,另一个正确类别会被当错误样本惩罚。训练时模型越努力,业务上越拧巴。
如果确认是多标签,把输出层改成sigmoid,损失函数改成binary cross-entropy,每个类别独立判断“是或不是”,这才是正确路径。在评估上也不能再用混淆矩阵和单标签指标,要改用海明损失、子集准确率这类多标签指标。
我在实际项目里的最后一个习惯是:任何多分类任务,训练结束后第一件事不是看准确率,而是打印两样东西——classification_report和归一化混淆矩阵。有一次文本分类,准确率从0.86提到0.88,看起来不错,但仔细看归一化混淆矩阵才发现,模型把“投诉”类更多地推给了“咨询”类,投诉类召回率反而下降了。后来加了类别权重,牺牲了一点整体准确率,却让业务里最关键的投诉类召回率回到了正常水平。多分类优化,永远先看矩阵,再看数字。