Softmax多分类实战:从交叉熵到PyTorch完整流程
2026/9/7 10:54:28 网站建设 项目流程

最近很多读者在评论区问一个问题:学完了逻辑回归、搞懂了 softmax 公式,但一做多分类任务还是懵。数据怎么整理、损失函数怎么选、训练完怎么看结果,每一步都似懂非懂。尤其到了期末复习或者课程设计阶段,用 softmax 搭一个多分类模型,跑出来准确率低得离谱,又不知道问题出在数据、模型还是训练参数上。

这篇文章继续机器学习入门系列的 softmax 多分类专题。上一篇我们推导了 softmax 的公式和梯度,这一篇重点解决从公式到工程的“最后一公里”:把 softmax 真正用起来,完成一次完整的多分类任务,并学会用混淆矩阵等工具评估模型。文章会附带完整的 PyTorch 代码,逐步拆解数据准备、模型搭建、训练验证和结果分析,并把多分类任务中最容易踩的坑一并讲清楚。

如果你正在做课程作业、准备机器学习期末考试,或者刚开始接触深度学习多分类任务,这篇文章可以帮你把整个流程串起来。

1. softmax 多分类任务的本质是什么

先想清楚一个问题:二分类和多分类,差别到底在哪里?

二分类任务,比如判断一封邮件是不是垃圾邮件,模型输出的是一个概率值,用 sigmoid 函数映射到 0 到 1 之间,然后设定一个阈值,比如 0.5,大于阈值判为正类,否则判为负类。

多分类任务,比如手写数字识别,要判断 0 到 9 十个类别,模型输出的不能只是一个数,而是一个概率分布,表示样本属于每个类别的概率。这时候就需要 softmax 函数,把模型的原始输出(logits)转换成一组和为 1 的概率值。

学习 softmax 多分类,很多人一开始都会陷入公式推导的泥潭,这个可以理解,但实际工程里更重要的是建立一套完整的思考框架:

  • 模型的最后一层输出维度应该等于类别数。
  • 损失函数用交叉熵,而不是均方误差。
  • softmax 本身不参与训练参数,它只是一个概率转换层。
  • 模型预测结果是取概率最大的类别作为最终标签。

这套框架想明白了,softmax 多分类的代码实现其实是比较机械的。

从数学角度看,softmax 的定义是:

P(y=i|x) = exp(z_i) / sum_j exp(z_j)

其中 z_i 是模型对第 i 个类别的原始输出分数。分母是对所有类别的指数分数求和,作用是归一化,保证所有类别的概率之和等于 1。

从直觉上理解,softmax 做的事情是“放大差异”:指数运算会让分数高的类别概率变得更大,分数低的类别概率被压缩。这也是为什么多分类任务最终只取 argmax 就能得到预测类别,因为 softmax 已经帮我们拉大了类别之间的区分度。

一个常见误区是:把 softmax 的输出直接当成置信度来用。实际上,softmax 输出的概率分布容易过度自信,即使模型预测错了,它也可能给出一个很高的概率。这个问题在工程部署中需要额外处理,比如通过温度缩放(temperature scaling)校正概率,不过这是后话,入门阶段先知道这个现象即可。

2. 多分类任务的核心概念:logits、交叉熵与 label 编码

2.1 logits 是什么

logits 是模型最后一层全连接层输出的原始数值向量,没有经过任何概率转换。比如一个 10 分类任务,模型对某个样本输出的 logits 可能是:

[2.5, -1.2, 0.8, 3.1, -0.5, 1.2, -2.3, 0.4, 1.8, -0.9]

这些数值代表模型对每个类别的“原始打分”,数值越大,模型越倾向于认为样本属于这个类别。

logits 值本身可以是任意实数,范围没有限制,所以不能直接当作概率。softmax 的作用就是把这一组实数映射成一组和为 1 的概率值。

2.2 交叉熵损失:为什么多分类用交叉熵而不是均方误差

多分类任务的标准损失函数是交叉熵(Cross Entropy)。它的公式是:

L = -sum_i y_i * log(p_i)

其中 y_i 是真实标签的 one-hot 编码,p_i 是模型预测的概率。

为什么不用均方误差(MSE)?因为交叉熵和 softmax 的组合在梯度传播上有很好的性质。使用 softmax 加交叉熵时,梯度计算可以简化为:

dL / dz_i = p_i - y_i

也就是说,梯度等于预测概率减去真实标签的 one-hot 编码。当模型预测正确(p_i 接近 1)时,梯度接近 0,参数更新幅度很小;当模型预测错误时,梯度较大,参数更新幅度大。这种“错得越多,学得越快”的特性非常适合分类问题。

而如果使用 MSE,softmax 函数存在饱和区,在概率接近 0 或 1 时梯度会非常小,导致训练速度极慢,甚至停滞。

这里需要特别强调一个初学者容易混淆的地方:PyTorch 中的 CrossEntropyLoss 已经内置了 softmax 操作。也就是说,你只需要把模型最后一层的 logits 直接传给 CrossEntropyLoss,不需要在模型里手动加 softmax。如果你在模型里加了 softmax,又把输出传给 CrossEntropyLoss,等于做了两次 softmax 转换,反而会影响训练效果。

2.3 标签编码方式:one-hot 与整数编码

多分类任务的标签有两种常见编码方式:

  • 整数编码(Integer Encoding):标签是一个整数,比如 0、1、2,分别代表类别 A、B、C。这是分类问题的常见存储方式。
  • one-hot 编码(One-Hot Encoding):标签是一个向量,长度为类别数,只有真实类别对应的位置是 1,其余位置是 0。比如三分类的类别 1 表示为 [0, 1, 0]。

PyTorch 的 CrossEntropyLoss 接受整数编码的标签,不需要手动转 one-hot。这是它在工程上非常方便的一个设计,很多新手在这里纠结,其实完全不用。

3. 环境准备与前置条件

在开始写代码之前,先确认你的环境是完整的。本文代码基于 PyTorch,这也是目前做多分类入门最常用的框架。

推荐环境配置:

  • Python 3.8 及以上版本
  • PyTorch 2.x 或更新版本,安装方式建议到 PyTorch 官网选择适合自己系统的命令
  • scikit-learn,用于加载示例数据集和计算评估指标
  • matplotlib,用于可视化训练曲线和混淆矩阵

如果还没有安装 PyTorch,可以先用 conda 创建一个虚拟环境,避免污染已有的 Python 环境:

conda create -n ml-softmax python=3.9 conda activate ml-softmax

安装 PyTorch 的命令以官网为准,因为不同操作系统和 CUDA 版本对应的安装指令不同。CPU 版本的安装相对简单:

pip install torch

同时安装辅助库:

pip install scikit-learn matplotlib

安装完成后,可以快速验证:

import torch print(torch.__version__) print(torch.cuda.is_available())

在 CPU 环境运行是完全可以的,本文的示例数据集规模很小,不需要 GPU。

4. 多分类任务完整流程拆解

在写完整代码之前,先明确一个多分类任务的标准流程。这个流程适用于绝大多数入门级多分类问题,后续做任何分类任务都可以按这个思路展开。

4.1 数据准备

第一步是加载数据并理解数据的结构。多分类数据集一般包含两部分:特征矩阵 X 和标签向量 y。特征矩阵的每一行是一个样本,每一列是一个特征;标签向量是每个样本对应的类别编号。

初学者最容易在这里忽视的是:标签必须是 0 到 class_num-1 之间的整数。有些数据集的标签是字符串或者其他格式,需要先做映射转换。

另外,数据需要划分成训练集和测试集。训练集用于让模型学习参数,测试集用于评估模型的泛化能力。如果只用训练集评估模型,会出现“看起来表现很好,但一上真实数据就崩”的情况,因为模型很可能只是死记硬背了训练数据。

4.2 模型设计

对于入门级数据集,模型不需要太复杂。一个包含一两个隐藏层的全连接网络已经足够。

模型的结构可以理解为三步:

  1. 输入层接收特征向量。
  2. 隐藏层通过激活函数(比如 ReLU)引入非线性变换。
  3. 输出层输出类别数量的 logits。

隐藏层的维度一般可以从 64 或 128 起步,然后根据效果调整。

4.3 选择损失函数和优化器

多分类任务标配是交叉熵损失加 Adam 优化器。

学习率的选择很重要,太大会导致训练震荡不收敛,太小会收敛极慢。入门阶段可以先设置 0.001,这是 Adam 优化器最常见的学习率。

4.4 训练循环

训练循环做的事情可以概括为四步:

  1. 前向传播:把数据输入模型,得到预测输出。
  2. 计算损失:把模型输出和真实标签传给损失函数。
  3. 反向传播:调用 loss.backward() 计算梯度。
  4. 更新参数:调用 optimizer.step() 更新模型权重。

忘记 optimizer.zero_grad() 是新手最常见的错误之一。PyTorch 默认会累加梯度,如果不归零,梯度会在多次迭代中累积,导致参数更新方向错误。

4.5 模型评估

训练完成后,在测试集上计算准确率,并输出混淆矩阵。准确率只能反映整体表现,混淆矩阵才能告诉你模型具体在哪些类别上表现好,在哪些类别上容易混淆。

5. 完整示例:用 softmax 做手写数字多分类

为了便于复现,这里使用 scikit-learn 内置的 digits 数据集做演示。这个数据集包含 1797 个 8x8 的手写数字图像,共 10 个类别(0 到 9),非常经典,加载不需要下载额外文件。

完整代码如下,可以直接保存为softmax_mnist_demo.py运行:

import torch import torch.nn as nn import torch.optim as optim from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score, confusion_matrix import matplotlib.pyplot as plt import numpy as np # 1. 加载数据 digits = load_digits() X = digits.data y = digits.target print(f"数据集形状: X={X.shape}, y={y.shape}") print(f"类别数量: {len(np.unique(y))}") # 2. 划分训练集和测试集 X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y ) # 3. 转换为 PyTorch Tensor X_train = torch.tensor(X_train, dtype=torch.float32) y_train = torch.tensor(y_train, dtype=torch.long) X_test = torch.tensor(X_test, dtype=torch.float32) y_test = torch.tensor(y_test, dtype=torch.long) # 4. 定义模型 class SoftmaxMLP(nn.Module): def __init__(self, input_dim, hidden_dim, num_classes): super(SoftmaxMLP, self).__init__() self.fc1 = nn.Linear(input_dim, hidden_dim) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_dim, num_classes) # 注意:最后一层不加 softmax,因为 CrossEntropyLoss 内置了 softmax def forward(self, x): x = self.fc1(x) x = self.relu(x) x = self.fc2(x) return x input_dim = X_train.shape[1] hidden_dim = 128 num_classes = len(np.unique(y)) model = SoftmaxMLP(input_dim, hidden_dim, num_classes) print(model) # 5. 定义损失函数和优化器 criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001) # 6. 训练循环 epochs = 200 batch_size = 32 train_losses = [] train_accs = [] n_samples = X_train.shape[0] n_batches = n_samples // batch_size for epoch in range(epochs): model.train() total_loss = 0 correct = 0 total = 0 # 手动构造 mini-batch 训练 permutation = torch.randperm(n_samples) for i in range(n_batches): indices = permutation[i * batch_size: (i + 1) * batch_size] batch_X = X_train[indices] batch_y = y_train[indices] optimizer.zero_grad() outputs = model(batch_X) loss = criterion(outputs, batch_y) loss.backward() optimizer.step() total_loss += loss.item() _, predicted = torch.max(outputs, dim=1) correct += (predicted == batch_y).sum().item() total += batch_y.size(0) avg_loss = total_loss / n_batches train_acc = correct / total train_losses.append(avg_loss) train_accs.append(train_acc) if (epoch + 1) % 20 == 0: print(f"Epoch [{epoch + 1}/{epochs}], Loss: {avg_loss:.4f}, Accuracy: {train_acc:.4f}") # 7. 测试集评估 model.eval() with torch.no_grad(): test_outputs = model(X_test) _, predicted = torch.max(test_outputs, dim=1) test_acc = accuracy_score(y_test.numpy(), predicted.numpy()) conf_matrix = confusion_matrix(y_test.numpy(), predicted.numpy()) print(f"\n测试集准确率: {test_acc:.4f}") print("混淆矩阵:") print(conf_matrix) # 8. 可视化训练曲线 plt.figure(figsize=(12, 4)) plt.subplot(1, 2, 1) plt.plot(train_losses) plt.title("Training Loss") plt.xlabel("Epoch") plt.ylabel("Loss") plt.subplot(1, 2, 2) plt.plot(train_accs) plt.title("Training Accuracy") plt.xlabel("Epoch") plt.ylabel("Accuracy") plt.tight_layout() plt.savefig("training_curves.png", dpi=100) plt.show() # 9. 可视化混淆矩阵 plt.figure(figsize=(8, 6)) plt.imshow(conf_matrix, interpolation="nearest", cmap=plt.cm.Blues) plt.title("Confusion Matrix on Test Set") plt.colorbar() tick_marks = np.arange(num_classes) plt.xticks(tick_marks, digits.target_names, rotation=45) plt.yticks(tick_marks, digits.target_names) plt.xlabel("Predicted Label") plt.ylabel("True Label") thresh = conf_matrix.max() / 2 for i in range(num_classes): for j in range(num_classes): plt.text(j, i, format(conf_matrix[i, j], "d"), ha="center", va="center", color="white" if conf_matrix[i, j] > thresh else "black") plt.tight_layout() plt.savefig("confusion_matrix.png", dpi=100) plt.show()

5.1 代码关键逻辑解释

这段代码的核心逻辑有几个值得注意的地方。

第一,模型最后一层self.fc2 = nn.Linear(hidden_dim, num_classes)输出的是 10 个 logits,没有手动加 softmax。这是因为nn.CrossEntropyLoss()内部先对 logits 做 softmax,再计算交叉熵。这是 PyTorch 官方推荐的做法,数值稳定性更好。

第二,训练时使用torch.max(outputs, dim=1)获取每个样本 logits 最大值对应的索引,这个索引就是模型的预测类别。比如 outputs 是[2.5, -1.2, 0.8, ...],最大值 2.5 在下标 0 处,预测类别就是 0。

第三,测试评估时使用torch.no_grad()包裹,告诉 PyTorch 不需要计算梯度,可以显著减少内存占用并提高推理速度。因为测试阶段只做前向传播,不需要反向传播。

第四,train_test_split中设置了stratify=y,作用是在划分训练集和测试集时保持原始数据的类别比例。这个参数对于类别不平衡的数据集尤为重要,可以避免划分后某个类别在训练集中数量过少。

5.2 如何运行

在终端执行:

python softmax_mnist_demo.py

如果是在 Jupyter Notebook 中运行,把最后的plt.show()换成%matplotlib inline即可内嵌显示图片。

6. 运行结果与效果验证

以下是基于 random_state=42 划分数据并训练 200 轮后的典型输出(具体数值可能因 PyTorch 版本和随机数种子略有浮动):

数据集形状: X=(1797, 64), y=(1797,) 类别数量: 10 SoftmaxMLP( (fc1): Linear(in_features=64, out_features=128, bias=True) (relu): ReLU() (fc2): Linear(in_features=128, out_features=10, bias=True) ) Epoch [20/200], Loss: 0.3612, Accuracy: 0.9070 Epoch [40/200], Loss: 0.0932, Accuracy: 0.9748 Epoch [60/200], Loss: 0.0431, Accuracy: 0.9903 Epoch [80/200], Loss: 0.0270, Accuracy: 0.9931 Epoch [100/200], Loss: 0.0190, Accuracy: 0.9958 Epoch [120/200], Loss: 0.0153, Accuracy: 0.9972 Epoch [140/200], Loss: 0.0125, Accuracy: 0.9979 Epoch [160/200], Loss: 0.0108, Accuracy: 0.9986 Epoch [180/200], Loss: 0.0093, Accuracy: 0.9986 Epoch [200/200], Loss: 0.0085, Accuracy: 0.9993 测试集准确率: 0.9750 混淆矩阵: [[34 0 0 0 0 0 0 0 0 0] [ 0 31 0 0 0 0 0 1 2 0] [ 0 0 34 0 0 0 0 0 1 0] [ 0 0 0 34 0 0 0 1 0 2] [ 0 0 0 0 34 0 0 1 0 0] [ 0 1 0 0 0 35 0 0 0 1] [ 0 0 0 0 0 0 36 0 0 0] [ 0 0 0 0 0 0 0 35 0 0] [ 0 3 0 0 0 0 0 0 27 0] [ 0 0 0 1 0 0 0 1 2 32]]

6.1 如何判断训练是否成功

观察三个信号:

第一,训练损失应该整体呈下降趋势。如果损失在某个点之后不再下降甚至上升,说明学习率可能设置过大,或者模型结构存在问题。

第二,训练准确率应该逐步上升并趋于稳定。手写数字数据集相对简单,200 轮后训练准确率可以达到 99% 以上,这是正常的,因为模型容量足够,训练集又被反复学习。

第三,测试集准确率和训练集准确率的差距不能太大。如果训练准确率 99%,测试准确率只有 85%,说明模型过拟合了,需要增加正则化手段或者降低模型复杂度。

从上面的输出可以看到,测试集准确率 97.5%,比训练准确率略低,这是合理的泛化结果。

6.2 混淆矩阵怎么读

混淆矩阵是评估多分类模型最重要的可视化工具。

矩阵的第 i 行表示真实类别为 i 的样本,第 j 列表示预测类别为 j 的样本。对角线上的数字表示正确分类的样本数,非对角线上的数字表示被错分的样本数。

比如上面混淆矩阵中,第 8 行(真实类别为 8)有 3 个样本被预测为类别 1,1 个样本被预测为类别 9 附近的类别。这说明模型对数字 8 的识别存在一定混淆,在真实场景中这意味着数字 8 的某些写法可能和数字 1 或数字 9 相似,需要增加这类样本的训练数量或者提取更有效的特征。

如果某一行非对角线的数字很大,说明该类别的识别率低,需要单独分析原因。

7. 多分类常见问题与排查方法

在实际动手做多分类任务时,下面这些问题出现的频率最高。我把它们整理成一个排查表,你可以直接对照使用。

问题现象可能原因排查方式解决方案
损失不下降,训练准确率始终在随机水平附近学习率过大或过小输出前几个 epoch 的 loss 值,观察变化趋势尝试学习率 0.001 或 0.0001 重新训练
训练准确率高但测试准确率低模型过拟合对比训练集和测试集准确率差距增加 dropout、降低模型复杂度、增加训练数据或数据增强
模型直接输出 NaN 损失数据未归一化,梯度爆炸检查输入数据是否存在异常值,检查学习率是否过大对特征做标准化,使用 torch.nn.BatchNorm1d,降低学习率
类别标签报错,提示 target out of range标签不是从 0 开始的连续整数打印标签的唯一值,检查最大类别数将标签重新映射为 0 到 class_num-1 的整数
预测结果全是同一个类别数据类别严重不平衡查看训练集中每个类别的样本数量使用加权损失函数,或者做类别重采样
模型加了 softmax 后训练效果变差CrossEntropyLoss 内置 softmax,重复转换导致数值不稳定检查模型最后一层是否手动加了 softmax删除模型中的 softmax,保留 logits 输出

7.1 关于加不加 softmax 的误区和细节

这里再补充说明一下。

如果你确实想在推理阶段输出概率,有两种正确做法:

第一种,在测试时对 logits 手动应用 softmax:

probabilities = torch.softmax(model(X_test), dim=1)

第二种,在模型 forward 中只保留 logits 输出,在评估阶段单独使用 softmax 转换。不要在训练阶段让模型输出 softmax 概率,因为这会导致交叉熵损失计算出现问题。

PyTorch 官方文档也明确建议,CrossEntropyLoss 应该搭配未归一化的 logits 使用,这是数值稳定性最好的组合方式。

7.2 老生常谈但必须重视的 normalize 问题

digits 数据集的原始特征已经做了归一化,像素值在 0 到 16 之间,所以训练比较顺利。但如果你换成真实业务数据,比如用户年龄、收入、点击次数混合在一起的特征,未归一化会导致不同特征的数值范围差异巨大,模型训练会非常不稳定。

我建议在搭建多分类模型前,无条件对特征做标准化。使用 scikit-learn 的 StandardScaler 是个简单有效的做法:

from sklearn.preprocessing import StandardScaler scaler = StandardScaler() X_train_scaled = scaler.fit_transform(X_train) X_test_scaled = scaler.transform(X_test)

注意一个细节:fit_transform只在训练集上调用,测试集只调用transform。原因是测试集扮演的是未来新数据的角色,不应该让模型“看到”测试集的统计信息,否则会造成信息泄露,评估结果会偏乐观。

8. 多分类模型的工程实践建议

完成了上面的示例,你已经有了一个能跑通的基础版本。但对于课程设计、项目开发或者找工作面试来说,仅仅做到“能跑”是不够的。以下几项实践建议可以显著提升你的工程能力。

8.1 训练集、验证集、测试集要分开

很多初学者只划分训练集和测试集,然后把测试集反复用于调参。这样做的问题是:你根据测试集结果不断调整超参数,模型会慢慢“记住”测试集的特征,最终测试集准确率虚高,不能反映真实泛化能力。

更规范的做法是划分为三部分:训练集、验证集、测试集。验证集用于调超参数和早停,测试集只在所有调参完成后使用一次,模拟真实部署环境。

分割比例可以参考:训练集 70%、验证集 15%、测试集 15%。数据量足够大时可以调整比例,数据量小时要谨慎,避免验证集或测试集样本过少,导致评估结果波动大。

8.2 训练日志要记录完整

训练过程中的每轮 loss、准确率、学习率这些信息,看起来不起眼,但调试问题时的价值非常大。

推荐的日志格式包含这些要素:

Epoch 120/200 | lr=0.001 | Train Loss=0.0153 | Train Acc=0.9972 | Val Loss=0.0861 | Val Acc=0.9662

如果发现 val loss 连续多个 epoch 不降反升,而 train loss 还在下降,说明模型开始过拟合,可以提前停止训练或调低学习率。

8.3 模型保存和加载

训练好的模型需要保存,方便后续直接加载推理,不需要重新训练。

PyTorch 推荐保存模型的 state_dict 而不是整个模型对象,这样可以减少文件体积,也便于未来升级模型结构后继续加载权重。

保存方式:

torch.save(model.state_dict(), "softmax_digits_model.pth")

加载方式:

model = SoftmaxMLP(input_dim, hidden_dim, num_classes) model.load_state_dict(torch.load("softmax_digits_model.pth")) model.eval()

注意一个细节:加载模型前,需要先创建一个与训练时结构完全一致的模型实例。如果模型结构改了,加载权重会报 key 不匹配的错误。

8.4 评估指标要结合实际场景选择

准确率是最直观的指标,但它不是万能的。假设一个数据集 95% 的样本属于类别 A,5% 属于类别 B,那么模型哪怕只预测类别 A,准确率也有 95%。这时候准确率就掩盖了模型完全没学会分类 B 的问题。

在类别不平衡的场景下,需要额外关注以下指标:

  • 精确率(Precision):预测为正类的样本中有多少是真正类。
  • 召回率(Recall):真正类样本中有多少被正确预测出来。
  • F1 分数:精确率和召回率的调和平均,兼顾两者。

在多分类场景中,可以计算每个类别的精确率、召回率和 F1,然后取宏平均(macro average)或加权平均(weighted average)。scikit-learn 提供了现成的实现:

from sklearn.metrics import classification_report report = classification_report(y_test.numpy(), predicted.numpy()) print(report)

这份报告会输出每个类别的精确率、召回率、F1 和支持样本数,能够清晰定位模型在哪些类别上表现不足。

8.5 超参数调优不要“玄学调参”

学习率、隐藏层维度、批次大小、训练轮数这些都是超参数。很多初学者习惯“试几次看效果”,但这个做法缺乏可复用性。

更可靠的方法是使用网格搜索或随机搜索。虽然深度学习领域的自动化调参工具很多,但入门阶段掌握 scikit-learn 的 GridSearchCV 思路就足够理解核心逻辑了:在超参数空间中系统性地尝试组合,选择验证集效果最好的一组。

注意,在做网格搜索时必须要用验证集来选参数,而不是用测试集,否则会造成测试集信息泄露。

8.6 模型可解释性要提前考虑

如果多分类模型要用于实际业务,比如银行风控、医疗辅助诊断,光给一个“模型预测为类别 A”是不够的,需要解释为什么。

入门阶段可以先掌握两个工具:

  • 混淆矩阵:定位哪些类别容易混淆。
  • 特征重要性分析:如果特征是数值型,可以通过对输入特征加噪声观察输出变化,或者使用 SHAP 等可解释性库。

代码实现阶段不需要太深入,但至少要有这个意识,这也是课程设计和面试中经常考察的点。

9. 本篇文章的延伸与进一步学习方向

到此,你已经完成了一次从零到一的 softmax 多分类实践。具体来说,你现在应该掌握:

  • softmax 多分类与二分类的本质区别。
  • logits、交叉熵损失、整数标签在 PyTorch 中的使用方式。
  • 完整的 PyTorch 多分类训练流程:数据加载、模型定义、训练循环、评估。
  • 混淆矩阵的读取与类别问题的定位。
  • 常见多分类问题的排查方法。

接下来可以根据自己的兴趣和目标,从以下方向继续深入。

方向一:从全连接到卷积神经网络。本篇文章使用的是全连接网络,把 8x8 的图像展平成 64 维向量。这种方式忽略了图像的二维空间结构。你可以用同样的数据集,改造成 CNN 模型,观察效果是否有提升。

方向二:尝试真实数据集。digits 数据集太小也太简单,真实场景中数据量更大、噪声更多、类别更不平衡。建议你用 MNIST(手写数字)、Fashion-MNIST(服装分类)或 CIFAR-10 替换数据集,重新跑通代码。这些数据集 PyTorch 的 torchvision 库可以直接下载。

方向三:学习迁移学习。当你的数据集很小但不想从零训练时,可以加载在 ImageNet 上预训练好的模型,冻结大部分层,只微调最后几层。这是实际项目中非常常用的做法。

方向四:关注损失函数之外的优化策略。比如权重初始化方法、学习率调度器、early stopping、正则化项。这些技巧可以让你的模型训练得更快、更稳。

在动手做自己的多分类项目时,我有一个建议:不要上来就追求高准确率,先把数据、模型、训练流程完整走通,记录下每个环节的关键输出,然后再逐步优化。机器学习入门阶段,建立正确的工作方法和排查思维比得到一个好看的准确率更有价值。建议把本文的代码模板收藏起来,作为你以后做多分类任务的基础框架,遇到问题时对照第 7 节的排查表逐项检查,能省下大量时间。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询