Softmax多分类实战:从交叉熵到PyTorch完整流程
最近很多读者在评论区问一个问题学完了逻辑回归、搞懂了 softmax 公式但一做多分类任务还是懵。数据怎么整理、损失函数怎么选、训练完怎么看结果每一步都似懂非懂。尤其到了期末复习或者课程设计阶段用 softmax 搭一个多分类模型跑出来准确率低得离谱又不知道问题出在数据、模型还是训练参数上。这篇文章继续机器学习入门系列的 softmax 多分类专题。上一篇我们推导了 softmax 的公式和梯度这一篇重点解决从公式到工程的“最后一公里”把 softmax 真正用起来完成一次完整的多分类任务并学会用混淆矩阵等工具评估模型。文章会附带完整的 PyTorch 代码逐步拆解数据准备、模型搭建、训练验证和结果分析并把多分类任务中最容易踩的坑一并讲清楚。如果你正在做课程作业、准备机器学习期末考试或者刚开始接触深度学习多分类任务这篇文章可以帮你把整个流程串起来。1. softmax 多分类任务的本质是什么先想清楚一个问题二分类和多分类差别到底在哪里二分类任务比如判断一封邮件是不是垃圾邮件模型输出的是一个概率值用 sigmoid 函数映射到 0 到 1 之间然后设定一个阈值比如 0.5大于阈值判为正类否则判为负类。多分类任务比如手写数字识别要判断 0 到 9 十个类别模型输出的不能只是一个数而是一个概率分布表示样本属于每个类别的概率。这时候就需要 softmax 函数把模型的原始输出logits转换成一组和为 1 的概率值。学习 softmax 多分类很多人一开始都会陷入公式推导的泥潭这个可以理解但实际工程里更重要的是建立一套完整的思考框架模型的最后一层输出维度应该等于类别数。损失函数用交叉熵而不是均方误差。softmax 本身不参与训练参数它只是一个概率转换层。模型预测结果是取概率最大的类别作为最终标签。这套框架想明白了softmax 多分类的代码实现其实是比较机械的。从数学角度看softmax 的定义是P(yi|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参数更新幅度很小当模型预测错误时梯度较大参数更新幅度大。这种“错得越多学得越快”的特性非常适合分类问题。而如果使用 MSEsoftmax 函数存在饱和区在概率接近 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 python3.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 模型设计对于入门级数据集模型不需要太复杂。一个包含一两个隐藏层的全连接网络已经足够。模型的结构可以理解为三步输入层接收特征向量。隐藏层通过激活函数比如 ReLU引入非线性变换。输出层输出类别数量的 logits。隐藏层的维度一般可以从 64 或 128 起步然后根据效果调整。4.3 选择损失函数和优化器多分类任务标配是交叉熵损失加 Adam 优化器。学习率的选择很重要太大会导致训练震荡不收敛太小会收敛极慢。入门阶段可以先设置 0.001这是 Adam 优化器最常见的学习率。4.4 训练循环训练循环做的事情可以概括为四步前向传播把数据输入模型得到预测输出。计算损失把模型输出和真实标签传给损失函数。反向传播调用 loss.backward() 计算梯度。更新参数调用 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_size0.2, random_state42, stratifyy ) # 3. 转换为 PyTorch Tensor X_train torch.tensor(X_train, dtypetorch.float32) y_train torch.tensor(y_train, dtypetorch.long) X_test torch.tensor(X_test, dtypetorch.float32) y_test torch.tensor(y_test, dtypetorch.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(), lr0.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, dim1) 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(fEpoch [{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, dim1) 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, dpi100) plt.show() # 9. 可视化混淆矩阵 plt.figure(figsize(8, 6)) plt.imshow(conf_matrix, interpolationnearest, cmapplt.cm.Blues) plt.title(Confusion Matrix on Test Set) plt.colorbar() tick_marks np.arange(num_classes) plt.xticks(tick_marks, digits.target_names, rotation45) 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), hacenter, vacenter, colorwhite if conf_matrix[i, j] thresh else black) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi100) plt.show()5.1 代码关键逻辑解释这段代码的核心逻辑有几个值得注意的地方。第一模型最后一层self.fc2 nn.Linear(hidden_dim, num_classes)输出的是 10 个 logits没有手动加 softmax。这是因为nn.CrossEntropyLoss()内部先对 logits 做 softmax再计算交叉熵。这是 PyTorch 官方推荐的做法数值稳定性更好。第二训练时使用torch.max(outputs, dim1)获取每个样本 logits 最大值对应的索引这个索引就是模型的预测类别。比如 outputs 是[2.5, -1.2, 0.8, ...]最大值 2.5 在下标 0 处预测类别就是 0。第三测试评估时使用torch.no_grad()包裹告诉 PyTorch 不需要计算梯度可以显著减少内存占用并提高推理速度。因为测试阶段只做前向传播不需要反向传播。第四train_test_split中设置了stratifyy作用是在划分训练集和测试集时保持原始数据的类别比例。这个参数对于类别不平衡的数据集尤为重要可以避免划分后某个类别在训练集中数量过少。5.2 如何运行在终端执行python softmax_mnist_demo.py如果是在 Jupyter Notebook 中运行把最后的plt.show()换成%matplotlib inline即可内嵌显示图片。6. 运行结果与效果验证以下是基于 random_state42 划分数据并训练 200 轮后的典型输出具体数值可能因 PyTorch 版本和随机数种子略有浮动数据集形状: X(1797, 64), y(1797,) 类别数量: 10 SoftmaxMLP( (fc1): Linear(in_features64, out_features128, biasTrue) (relu): ReLU() (fc2): Linear(in_features128, out_features10, biasTrue) ) 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 个样本被预测为类别 11 个样本被预测为类别 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 手动应用 softmaxprobabilities torch.softmax(model(X_test), dim1)第二种在模型 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 | lr0.001 | Train Loss0.0153 | Train Acc0.9972 | Val Loss0.0861 | Val Acc0.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% 的样本属于类别 A5% 属于类别 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 节的排查表逐项检查能省下大量时间。