这次我们来看机器学习入门阶段最常遇到的一个概念:softmax 多分类。很多新手在做二分类时用 sigmoid 很顺手,一旦遇到手写数字识别、图像分类、文本多标签这些任务,就不知道输出层该怎么设计、损失函数该选什么、准确率怎么评估。这篇文章就把 softmax 多分类这条线完整串一遍,从数学原理到 Python 代码实现,再到一个可以运行的训练验证流程,一次性讲透。
先给结论:softmax 是深度学习中处理多分类问题的标准输出层方案。它的作用是把模型输出的原始分数(logits)转换成一个概率分布,让每个类别的预测值都在 0 到 1 之间,并且所有类别的概率之和等于 1。配合交叉熵损失函数,模型训练时能获得稳定且有效的梯度信号,从而让分类准确率逐步提升。
这篇文章会带你完成以下内容:
- 搞懂 softmax 的数学定义和计算过程;
- 用 NumPy 从零实现 softmax 与交叉熵损失;
- 写一个完整的基于 softmax 的多分类训练脚本;
- 在经典数据集上验证分类效果;
- 输出准确率、混淆矩阵等评估结果;
- 整理常见报错和排查思路。
无论你是正在学机器学习课程的在校学生,还是刚接触深度学习的开发者,这篇文章都能让你少走弯路。建议收藏备用,需要的时候直接翻到对应章节。
1. 核心知识点速览
| 知识点 | 说明 |
|---|---|
| 适用任务 | 多分类任务,即类别数大于等于 3 |
| 输出形式 | 每个类别的概率分布,所有概率之和为 1 |
| 经典搭配 | softmax 输出层 + 交叉熵损失函数 |
| 实现工具 | Python + NumPy,或 PyTorch / TensorFlow |
| 数据集示例 | 鸢尾花数据集、手写数字 MNIST、CIFAR-10 |
| 评估指标 | 准确率、混淆矩阵、每个类别的精确率与召回率 |
| 门槛要求 | 会 Python 基础语法,理解矩阵乘法即可 |
| 运行环境 | 普通笔记本 CPU 即可完成本文全部实验 |
从表里可以看出,softmax 多分类的门槛并不高。即使你没有独立显卡,也能在 CPU 上完整跑通训练和评估流程。
2. softmax 多分类适用场景与边界
softmax 多分类适合处理"输入一个样本,输出它属于哪个类别"的问题。常见场景包括:
- 图片分类:判断一张图片是猫、狗还是鸟;
- 文本分类:判断一篇文章属于体育、娱乐还是科技;
- 手写数字识别:判断一张 28x28 的灰度图是数字 0 到 9 中的哪一个;
- 医疗辅助诊断:根据检查指标判断疾病类型。
在这些任务中,类别之间是互斥的,也就是说一个样本只能属于一个类别。softmax 天然适合这种设定,因为它会把所有类别的概率归一化,最后取概率最大的类别作为预测结果。
但 softmax 并不是所有分类问题的万能答案。下面这几种情况要特别注意:
- 多标签分类:如果一个样本同时属于多个类别,比如一张图片里既有猫又有狗,那就不能直接用 softmax。此时应该对每个类别单独使用 sigmoid,得到多个独立的概率值。
- 类别极度不平衡:如果某个类别的样本数量远大于其他类别,直接训练 softmax 模型会让模型偏向多数类。这时候需要引入类别权重、过采样或调整损失函数。
- 类别数量极大:当类别数达到数万甚至数百万时,标准的 softmax 计算量会非常大,这时候需要考虑负采样、层次 softmax 等优化方案。
另外还要强调一点:使用公开数据集做实验时,要注意数据集的版权和授权范围。像 MNIST、鸢尾花这类经典数据集通常都可以用于教学和科研,但如果要把训练好的模型商用,需要根据数据集的具体 license 确认是否允许。
3. 环境准备与前置条件
本文所有代码都用 Python 编写,依赖库只有 NumPy 和 Scikit-learn。PyTorch 版本会在后面的进阶部分给出,选装即可。
3.1 Python 版本
建议使用 Python 3.8 及以上版本。如果你已经装了 Anaconda,直接用 base 环境就行。
3.2 安装依赖库
在终端中执行以下命令:
pip install numpy scikit-learn matplotlib如果安装速度慢,可以换成国内镜像源:
pip install numpy scikit-learn matplotlib -i https://pypi.tuna.tsinghua.edu.cn/simple3.3 验证环境
python -c "import numpy, sklearn; print(numpy.__version__); print(sklearn.__version__)"能正常输出版本号,就说明环境准备好了。
3.4 硬件要求
本文的实验非常轻量,普通 CPU 即可完成。整个训练过程通常在几十秒到几分钟之间,具体耗时取决于数据集大小和训练轮数。不需要 GPU,也不需要关注显存占用。
4. softmax 数学原理与代码实现
4.1 为什么二分类的 sigmoid 不够用
二分类问题中,模型只需要输出一个概率值 p,另一个类别的概率就是 1-p。但当类别数变成 3、10、100 时,单个概率值就不够用了。我们需要为每个类别都输出一个概率,并且保证这些概率加起来等于 1。
一种直观的方式是:先让模型为每个类别输出一个分数(logit),分数越高代表属于该类别的可能性越大。然后把这些分数转换成概率。softmax 就是完成这个转换的标准方法。
4.2 softmax 函数定义
给定一个向量 z = [z₁, z₂, ..., zₖ],softmax 对第 i 个分量的计算方式为:
softmax(zᵢ) = exp(zᵢ) / Σⱼ exp(zⱼ)
其中分母对所有类别求和。这样得到的结果天然满足两个性质:
- 每个分量的值都在 0 到 1 之间;
- 所有分量的和等于 1。
因此可以把 softmax 的输出理解为一个概率分布。
4.3 用 NumPy 从零实现
import numpy as np def softmax(logits): """ 将原始分数转换为概率分布 参数: logits: shape 为 (batch_size, num_classes) 的二维数组 返回: 概率数组,shape 与 logits 相同,每行之和为 1 """ # 减去每行的最大值,防止指数计算溢出 shifted = logits - np.max(logits, axis=1, keepdims=True) exp_logits = np.exp(shifted) probs = exp_logits / np.sum(exp_logits, axis=1, keepdims=True) return probs代码中有一个非常关键的细节:减去了每行的最大值。如果不做这一步,当 logits 中存在较大数值时,exp 计算很容易发生数值溢出,导致结果变成 inf 或 nan。这是手写 softmax 时最常见的坑之一。
测试一下这个函数:
logits = np.array([ [2.0, 1.0, 0.1], [1.0, 3.0, 0.5] ]) probs = softmax(logits) print(probs) print("每行之和:", probs.sum(axis=1))预期输出应该显示每行之和都接近 1.0。
5. 损失函数与梯度更新
5.1 交叉熵损失
softmax 输出概率分布后,还需要一个损失函数来衡量预测分布和真实标签之间的差距。多分类任务最常用的是交叉熵损失(Cross Entropy Loss)。
对于一个样本,假设真实类别是 y,模型预测的概率分布是 p,交叉熵损失定义为:
L = -log(p_y)
也就是说,我们只关注真实类别对应的那个概率。真实类别的预测概率越接近 1,损失越小;越接近 0,损失越大。这非常符合直觉。
对整个 batch 的损失取平均,就得到当前训练步的损失值。
5.2 梯度推导
softmax 与交叉熵组合在一起时,梯度的形式非常简洁。假设模型的输出为 z,真实标签的 one-hot 编码为 y,则损失对 z 的梯度为:
∂L/∂z = p - y
这是一个很优雅的结果:预测概率减去真实标签。当预测完全正确时,p 和 y 相等,梯度为 0;当预测错误时,梯度的方向会引导模型参数往正确方向调整。
下面是用 NumPy 实现交叉熵损失和梯度的代码:
def cross_entropy_loss_with_grad(probs, labels): """ 计算交叉熵损失和梯度 参数: probs: shape 为 (batch_size, num_classes) 的预测概率 labels: shape 为 (batch_size,) 的整数标签,每个值在 [0, num_classes) 返回: loss: 标量损失值 grad: 形状与 probs 相同的梯度 """ batch_size = probs.shape[0] # 取出每个样本真实类别对应的概率 true_probs = probs[np.arange(batch_size), labels] # 加一个极小值防止 log(0) loss = -np.mean(np.log(true_probs + 1e-12)) # 构建 one-hot 编码 num_classes = probs.shape[1] one_hot = np.zeros_like(probs) one_hot[np.arange(batch_size), labels] = 1.0 # 梯度 = 预测概率 - 真实标签 grad = (probs - one_hot) / batch_size return loss, grad代码里的1e-12是数值保护项。当预测概率极小时,直接取 log 会得到负无穷,导致损失变成 inf。加上一个极小值可以避免这种情况。
6. 完整训练流程:以鸢尾花数据集为例
6.1 数据准备
鸢尾花数据集是机器学习入门最经典的分类数据集之一。它包含 3 个类别,每个类别 50 个样本,每个样本有 4 个特征。类别数正好是 3,非常适合演示 softmax 多分类。
加载数据并进行划分:
from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler # 加载数据 iris = load_iris() X = iris.data y = iris.target # 划分训练集和测试集 X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y ) # 标准化特征,加速收敛 scaler = StandardScaler() X_train = scaler.fit_transform(X_train) X_test = scaler.transform(X_test) # 在训练特征前加一列 1,作为偏置项 X_train = np.hstack([X_train, np.ones((X_train.shape[0], 1))]) X_test = np.hstack([X_test, np.ones((X_test.shape[0], 1))]) print("训练集大小:", X_train.shape) print("测试集大小:", X_test.shape)标准化这一步很重要。如果不做标准化,数值较大的特征会主导梯度更新,导致模型收敛速度变慢,甚至出现不稳定的情况。
6.2 模型定义与训练循环
我们定义一个线性模型,即 model(x) = x @ W,其中 W 的 shape 是 (特征数, 类别数)。这个模型没有隐藏层,是 logistic 回归的多分类版本,也叫 softmax 回归。虽然简单,但足以说明 softmax 多分类的完整流程。
np.random.seed(42) num_features = X_train.shape[1] num_classes = 3 # 随机初始化权重 W = np.random.randn(num_features, num_classes) * 0.01 learning_rate = 0.1 num_epochs = 500 for epoch in range(num_epochs): # 前向传播 logits = X_train @ W probs = softmax(logits) # 计算损失和梯度 loss, grad = cross_entropy_loss_with_grad(probs, y_train) # 梯度下降更新权重 W -= learning_rate * grad.T @ X_train # 每 50 轮打印一次损失 if (epoch + 1) % 50 == 0: print(f"Epoch {epoch + 1}, Loss: {loss:.6f}")注意这里梯度计算的方式:cross_entropy_loss_with_grad返回的 grad 是对 logits 的梯度,shape 为 (batch_size, num_classes)。要得到对权重 W 的梯度,需要用X_train.T @ grad做一次矩阵乘法。
训练完成后,在测试集上评估效果:
# 测试集预测 test_logits = X_test @ W test_probs = softmax(test_logits) test_pred = np.argmax(test_probs, axis=1) # 计算准确率 accuracy = np.mean(test_pred == y_test) print(f"测试集准确率: {accuracy:.4f}")运行完这段代码,测试集准确率通常在 0.9 以上。如果你的结果波动较大,可以调整学习率或训练轮数,也可以换一个随机种子重新初始化。
7. 效果验证:预测、准确率与混淆矩阵
7.1 手动查看预测结果
只关注准确率还不够,最好能直观看到每个样本的预测概率分布。下面这段代码会打印测试集中前 5 个样本的预测概率和最终类别:
for i in range(5): prob_row = test_probs[i] pred_class = test_pred[i] true_class = y_test[i] print(f"样本 {i}: 真实类别={true_class}, 预测类别={pred_class}, 概率={prob_row}")输出示例:
样本 0: 真实类别=0, 预测类别=0, 概率=[9.87e-01 1.23e-02 4.55e-04] 样本 1: 真实类别=1, 预测类别=1, 概率=[2.11e-03 8.65e-01 1.31e-01]从概率分布可以清楚看到模型对每个样本的置信度。如果某个样本预测错误,概率分布通常会显示两个类别的分数比较接近,这时候就可以进一步分析特征或数据是否有问题。
7.2 混淆矩阵
准确率会掩盖细节。比如在 3 分类问题中,准确率 0.9 可能是所有类别都表现不错,也可能是某一类特别好、另一类特别差。混淆矩阵能更细致地展现每个类别的分类情况。
from sklearn.metrics import confusion_matrix, classification_report cm = confusion_matrix(y_test, test_pred) print("混淆矩阵:") print(cm) print("\n分类报告:") print(classification_report(y_test, test_pred, target_names=iris.target_names))混淆矩阵的每一行代表真实类别,每一列代表预测类别。对角线上的数字表示正确分类的样本数,非对角线上的数字表示被误分类的样本数。
分类报告会给出每个类别的精确率、召回率和 F1 值,帮助我们定位是哪一类容易出错。
7.3 用 Scikit-learn 的 LogisticRegression 交叉验证
如果你不想手写训练循环,Scikit-learn 的LogisticRegression内部就是基于 softmax 多分类实现的。我们可以用它来验证一下自己的实现是否正确:
from sklearn.linear_model import LogisticRegression clf = LogisticRegression(max_iter=500) clf.fit(X_train, y_train) sk_pred = clf.predict(X_test) sk_accuracy = np.mean(sk_pred == y_test) print(f"Scikit-learn 准确率: {sk_accuracy:.4f}") print(f"手写实现准确率: {accuracy:.4f}")两款实现的准确率应该非常接近。如果差异很大,说明手写代码中的梯度计算或参数更新有 bug。
8. 训练效率与资源观察
8.1 CPU 训练耗时
鸢尾花数据集只有 120 个训练样本,训练 500 轮只需要几秒钟。我们可以用 Python 的time模块记录实际耗时:
import time start_time = time.time() # 这里放训练循环代码 end_time = time.time() print(f"训练耗时: {end_time - start_time:.2f} 秒")8.2 如何观察内存占用
虽然这个实验内存占用很小,但如果你换用更大的数据集,比如 MNIST,就有必要关注内存占用情况。可以用下面这段代码查看:
import psutil process = psutil.Process() memory_mb = process.memory_info().rss / 1024 / 1024 print(f"当前进程内存占用: {memory_mb:.2f} MB")如果没有安装 psutil,先执行pip install psutil。
8.3 batch size 对训练的影响
鸢尾花数据集很小,所以上面用了全批次梯度下降,即每轮都用全部训练样本计算梯度。当数据集变大时,需要引入小批量训练(Mini-batch Training),每次随机取一部分样本计算梯度。这样做有两个好处:
- 减少每次计算的开销;
- 增加梯度噪声,有助于跳出局部最优或鞍点。
小批量训练的伪代码如下:
batch_size = 32 for epoch in range(num_epochs): # 每个 epoch 打乱数据顺序 indices = np.random.permutation(X_train.shape[0]) X_shuffled = X_train[indices] y_shuffled = y_train[indices] for start in range(0, X_train.shape[0], batch_size): end = start + batch_size X_batch = X_shuffled[start:end] y_batch = y_shuffled[start:end] logits = X_batch @ W probs = softmax(logits) loss, grad = cross_entropy_loss_with_grad(probs, y_batch) W -= learning_rate * (grad.T @ X_batch)这里只是演示 batch 训练的基本写法。实际使用中,Shuffle 通常在训练数据加载器中完成,不需要手动实现。但理解这个过程对调试和调优非常有帮助。
8.4 使用 PyTorch 的参考实现
如果你准备进入深度学习框架的学习阶段,可以直接看 PyTorch 的实现方式。PyTorch 中的torch.nn.CrossEntropyLoss已经把 softmax 和交叉熵合并到一起,使用时不需要手动调用 softmax:
import torch import torch.nn as nn import torch.optim as optim # 数据转换 X_train_t = torch.tensor(X_train, dtype=torch.float32) y_train_t = torch.tensor(y_train, dtype=torch.long) # 定义线性模型 model = nn.Linear(5, 3) criterion = nn.CrossEntropyLoss() optimizer = optim.SGD(model.parameters(), lr=0.1) # 训练 for epoch in range(500): optimizer.zero_grad() logits = model(X_train_t) loss = criterion(logits, y_train_t) loss.backward() optimizer.step() if (epoch + 1) % 50 == 0: print(f"Epoch {epoch + 1}, Loss: {loss.item():.6f}")注意:这里手动把之前的偏置项列加入到了特征里,所以输入维度是 5。如果你不想手动加偏置,可以直接用nn.Linear(4, 3),让 PyTorch 自动管理偏置。
8.5 CPU 与 GPU 的选择
本文的实验规模不需要 GPU。即使你把数据集换成 MNIST 这种稍大的数据,在 CPU 上训练一个线性模型也只需要一两分钟。只有当模型变为多层神经网络、卷积网络,并且训练轮数大幅增加时,GPU 的加速效果才会明显体现出来。
所以,学习 softmax 多分类阶段,完全不用纠结显卡型号和显存大小。把注意力放在理解原理和代码实现上,比追求硬件配置更重要。
9. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| softmax 输出出现 nan | logits 数值过大,exp 溢出 | 检查 logits 的最大值 | 在 softmax 中减去每行最大值 |
| 损失一直不下降 | 学习率过大或过小 | 打印每轮 loss 观察变化 | 调整学习率,或做数据标准化 |
| 准确率在某一类上特别低 | 类别不平衡或特征区分度不足 | 查看混淆矩阵和分类报告 | 收集更多该类样本,或调整类别权重 |
| 梯度爆炸,权重变成 nan | 学习率过大 | 打印权重变化 | 降低学习率,增加初始化缩放 |
| PyTorch 中最后结果不对 | 手动调用了 softmax 后接 CrossEntropyLoss | 检查是否重复计算 softmax | 使用nn.CrossEntropyLoss时不要在输出层手动加 softmax |
| 训练集准确率很高,测试集很低 | 过拟合 | 比较训练集和测试集准确率 | 增加正则化、减少特征数、增加样本量 |
| 随机种子不同结果差异大 | 初始化权重不同 | 多次运行测试 | 设定固定 random seed,方便复现 |
| 数据集特征范围差异大 | 未做标准化 | 查看特征均值和方差 | 使用 StandardScaler 标准化 |
10. 最佳实践与下一步建议
10.1 Softmax 多分类的工程建议
第一,所有输入特征都要做标准化。softmax 回归本质上还是线性模型,对特征尺度敏感。特征数值范围差异过大会让梯度更新不稳定,直接表现为 loss 震荡或收敛缓慢。
第二,训练过程中要定期打印 loss 并在最终评估时使用独立的测试集。不要用训练集准确率评估模型效果,否则很容易被 "假高分" 迷惑。
第三,评估时至少同时看准确率和混淆矩阵。准确率只能反映整体情况,混淆矩阵能告诉你模型在哪些类别之间容易混淆,这对后续优化方向有直接指导意义。
第四,数值稳定性问题要在一开始就处理。手写 softmax 时务必要做 max 偏移,理解背后的原理比直接调库更有价值,但调库时也要清楚框架内部做了什么。
10.2 接下来可以继续深入的方向
如果你已经能独立完成上面所有代码,说明 softmax 多分类这条线基本掌握了。下一步可以根据自己的方向选择以下扩展:
- 从线性模型升级到多层感知机,在 softmax 前面加一个隐藏层,感受特征表示能力的提升;
- 在 MNIST 或 Fashion-MNIST 数据集上复现完整的图像分类流程,理解数据归一化、batch 训练和 epoch 等概念;
- 学习 PyTorch 的分类任务标准流程,包括 Dataset、DataLoader、模型定义和训练循环;
- 尝试在 CIFAR-10 上使用卷积神经网络做多分类,对比 softmax 回归和卷积模型的效果差异;
- 如果手头有真实业务数据,可以尝试将这套流程迁移到自己的多分类场景中。需要注意的是,商用前要确认数据来源合法、标注合规,涉及用户信息的数据要做好隐私保护。
10.3 最容易踩的坑
根据很多初学者反馈,softmax 多分类阶段最容易出问题的不是数学公式,而是下面三个地方:
- 忘记做数据标准化,导致 loss 不降或下降极慢;
- 在 PyTorch 中手动调用了 softmax 后又使用 CrossEntropyLoss,结果训练出的模型效果异常;
- 混淆矩阵和准确率结果对不上,往往是因为在某个环节把标签顺序搞乱了。
这些坑在调试时容易被忽略,但只要你按文章中的流程一步一步来,基本都能避开。
最后给出一句实用建议:不要急着把代码写到最复杂,先把线性模型下的 softmax 多分类完整跑通、看懂每一步的矩阵形状变化和梯度流向,再谈深度网络。这一步扎不扎实,直接决定你后面学卷积网络、注意力模型时是否顺利。