news 2026/9/7 6:05:27

Softmax多分类从零实现:数学原理与代码实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Softmax多分类从零实现:数学原理与代码实战

这次我们来看机器学习入门阶段最常遇到的一个概念:softmax 多分类。很多新手在做二分类时用 sigmoid 很顺手,一旦遇到手写数字识别、图像分类、文本多标签这些任务,就不知道输出层该怎么设计、损失函数该选什么、准确率怎么评估。这篇文章就把 softmax 多分类这条线完整串一遍,从数学原理到 Python 代码实现,再到一个可以运行的训练验证流程,一次性讲透。

先给结论:softmax 是深度学习中处理多分类问题的标准输出层方案。它的作用是把模型输出的原始分数(logits)转换成一个概率分布,让每个类别的预测值都在 0 到 1 之间,并且所有类别的概率之和等于 1。配合交叉熵损失函数,模型训练时能获得稳定且有效的梯度信号,从而让分类准确率逐步提升。

这篇文章会带你完成以下内容:

  1. 搞懂 softmax 的数学定义和计算过程;
  2. 用 NumPy 从零实现 softmax 与交叉熵损失;
  3. 写一个完整的基于 softmax 的多分类训练脚本;
  4. 在经典数据集上验证分类效果;
  5. 输出准确率、混淆矩阵等评估结果;
  6. 整理常见报错和排查思路。

无论你是正在学机器学习课程的在校学生,还是刚接触深度学习的开发者,这篇文章都能让你少走弯路。建议收藏备用,需要的时候直接翻到对应章节。

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/simple

3.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ⱼ)

其中分母对所有类别求和。这样得到的结果天然满足两个性质:

  1. 每个分量的值都在 0 到 1 之间;
  2. 所有分量的和等于 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),每次随机取一部分样本计算梯度。这样做有两个好处:

  1. 减少每次计算的开销;
  2. 增加梯度噪声,有助于跳出局部最优或鞍点。

小批量训练的伪代码如下:

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 输出出现 nanlogits 数值过大,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 多分类这条线基本掌握了。下一步可以根据自己的方向选择以下扩展:

  1. 从线性模型升级到多层感知机,在 softmax 前面加一个隐藏层,感受特征表示能力的提升;
  2. 在 MNIST 或 Fashion-MNIST 数据集上复现完整的图像分类流程,理解数据归一化、batch 训练和 epoch 等概念;
  3. 学习 PyTorch 的分类任务标准流程,包括 Dataset、DataLoader、模型定义和训练循环;
  4. 尝试在 CIFAR-10 上使用卷积神经网络做多分类,对比 softmax 回归和卷积模型的效果差异;
  5. 如果手头有真实业务数据,可以尝试将这套流程迁移到自己的多分类场景中。需要注意的是,商用前要确认数据来源合法、标注合规,涉及用户信息的数据要做好隐私保护。

10.3 最容易踩的坑

根据很多初学者反馈,softmax 多分类阶段最容易出问题的不是数学公式,而是下面三个地方:

  1. 忘记做数据标准化,导致 loss 不降或下降极慢;
  2. 在 PyTorch 中手动调用了 softmax 后又使用 CrossEntropyLoss,结果训练出的模型效果异常;
  3. 混淆矩阵和准确率结果对不上,往往是因为在某个环节把标签顺序搞乱了。

这些坑在调试时容易被忽略,但只要你按文章中的流程一步一步来,基本都能避开。

最后给出一句实用建议:不要急着把代码写到最复杂,先把线性模型下的 softmax 多分类完整跑通、看懂每一步的矩阵形状变化和梯度流向,再谈深度网络。这一步扎不扎实,直接决定你后面学卷积网络、注意力模型时是否顺利。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/7 6:02:29

区间测速技术规范解读:从标准到工程落地实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/7 6:02:18

AI与LLM如何助力引力波搜索:从匹配滤波到智能分类

引力波搜索听起来是纯物理领域的事,但最近这几年,AI/ML 和 LLM 已经实打实地进入了这个方向的分析流程。很多人第一反应是:引力波不是用匹配滤波在做吗?深度学习进来能干嘛?大语言模型又不能算波形。偏偏实际研究里&am…

作者头像 李华
网站建设 2026/9/7 6:02:14

Home Assistant 接入 DeepSeek:OpenAI 兼容协议配置与设备控制实践

在实际的智能家居项目中,Home Assistant 通常已经解决了“设备接入、自动化、统一控制”的问题,但“用自然语言和家里对话”这件事,长期以来只能靠固定的语音指令模板完成。把 ChatGPT、DeepSeek 这类 AI 大模型接入 Home Assistant 之后&…

作者头像 李华
网站建设 2026/9/7 6:01:27

MediaMTX 部署实战:从 Docker 单机到生产级上线

MediaMTX 部署实战:从 Docker 单机到生产级上线 【免费下载链接】mediamtx Ready-to-use Media-over-QUIC / SRT / WebRTC / RTSP / RTMP / LL-HLS / MPEG-TS / RTP live media server and media proxy that allows to read, publish, proxy, record and playback r…

作者头像 李华
网站建设 2026/9/7 6:00:43

Sublime Text 3 配置指南:告别破解版,打造高效开发环境

简介:这是Sublime Text 3的破解安装包资源,面向希望快速获得可用版本的前端开发者、编程初学者及需要离线安装环境的用户。压缩包采用zip格式,共包含2个文件,其中htm格式为安装说明,exe格式为主程序安装文件&#xff0…

作者头像 李华
网站建设 2026/9/7 5:57:39

修仙题材Minecraft服务器搭建指南:从Paper服务端到挂机修炼插件开发

各位朋友好,我是你们熟悉的后端开发博主。今天这篇不是讲 Spring Boot,也不是讲微服务,而是想和大家聊聊一个我最近业余时间一直折腾的话题:Minecraft 服务器,尤其是最近在圈子里非常火的“修仙题材 RPG 服务器”。你会…

作者头像 李华