我先说个反直觉的事:softmax回归这个名字里带着“回归”二字,但它从头到尾都在做分类。更反直觉的是,用PyTorch实现它,最方便的方式其实是一行F.cross_entropy直接搞定,根本轮不到你动手。那为什么还要“手动实现”?因为我发现,很多把API用得滚瓜烂熟的人,遇到loss突然变成NaN、梯度不再更新、自定义算子时梯度算错这类问题,完全没有排查方向。softmax为什么要先在指数运算前减掉最大值?交叉熵对logits的梯度为什么最终会化简成p − y这样的简单形式?这些细节,一行API不会告诉你,只有亲手从零写一遍,才会真正踩到那些坑、也才能真正明白设计者的用意。
这篇文章记录的就是我的完整实现过程:从数学原理推导、数值稳定性处理,到自定义反向传播函数,再训练一个三分类模型并画出决策边界。全程不依赖nn.Linear和nn.CrossEntropyLoss,只靠PyTorch最基础的计算算子。代码量不大,但里面藏着大量值得抠的细节。
1. 自带“回归”二字,却都在做分类——先搞懂softmax要解决什么问题
1.1 为什么线性回归的套路直接套不到分类上
传统线性回归做的是拟合连续值,模型输出一个实数,比如预测房价、温度。它的损失函数常用均方误差,因为输出和目标都是连续数值,差的平方天然可以作为“偏离程度”的度量。
但分类任务完全不同:目标是一个离散的类别标签。拿三分类来说,标签可能是0、1、2。如果直接训练一个线性模型y = XW + b去拟合这个数值标签,问题立刻就来了:
- 数值大小没有意义。类别
2并不比类别1“大两倍”,它们是并列关系; - 模型输出可以落在任意区间,比如
-3.7或15.2,怎么解释成“属于某个类别的置信度”?
所以分类任务需要把模型输出从“任意实数”压缩成一个合法概率分布:每个类别得到一个0~1之间的数值,且所有类别数值加起来等于1。softmax干的就是这件事。
1.2 从二分类的sigmoid推广到多分类的softmax
如果你熟悉逻辑回归,其实softmax回归是它的直接推广。二分类时,逻辑回归用sigmoid函数输出一个0~1之间的值,代表“属于正类”的概率。sigmoid的数学形式是:
σ(z) = 1 / (1 + exp(−z))它其实可以看成是在处理两个类别:正类的概率是σ(z),负类的概率是1 − σ(z)。两者加起来正好是1。
多分类时,假设我们有K个类别,模型对每个类别算出一个分数z_i,softmax把这个分数向量转成概率分布:
p_i = exp(z_i) / Σ_j exp(z_j)每个p_i都在0~1之间,所有p_i求和为1。这里的z_i通常叫logits,可以理解为“未经归一化的对数概率”。
那为什么一定要用exp,不用别的函数?这是初学者最爱问的问题。我的理解是这样:
exp是单调增函数,原始logits谁大,softmax之后概率谁就大,排名顺序不会变;exp会把差距“拉大”。比如原始logits是[1.0, 2.0],softmax后概率约[0.269, 0.731],两个类别差异被放大,模型训练时梯度信号更明显;- 指数形式和交叉熵损失天然搭配,后面推导梯度时你会看到,这种组合能让反向传播的梯度形式变得极其简洁。
另外有个容易被忽略的点:softmax回归本质上还是线性模型,决策边界是线性的。它没有隐藏层、没有非线性激活,所以它解决不了线性不可分的问题。但它作为理解多分类、理解损失函数、理解反向传播的入门载体,再合适不过。
2. 手写softmax最容易翻车的地方:数值稳定与形状细节
2.1 一个会让loss直接变成NaN的细节
数学公式p_i = exp(z_i) / Σ_j exp(z_j)看起来清爽,但直接照着写代码,十有八九会在某个时刻得到NaN。
原因很简单:exp增长极快。如果某个logit是1000,exp(1000)在浮点数里直接溢出为infinity,分子分母都变成无穷,结果自然变成NaN。这种问题在训练初期特别容易出现——权重一初始化,某些样本的logits可能就跑到几百上千,loss瞬间崩掉。
解决方法是利用softmax的一个数学性质:给输入向量整体加上或减去一个常数,结果不变。
p_i = exp(z_i − c) / Σ_j exp(z_j − c)通常取c为当前向量中的最大值。这样最大的logits变成0,指数运算永远不会溢出。这个操作通常叫log-sum-exp技巧的一部分。
对应到PyTorch代码:
def softmax(logits): max_logits = logits.max(dim=1, keepdim=True).values logits_shifted = logits - max_logits exp_logits = torch.exp(logits_shifted) probs = exp_logits / exp_logits.sum(dim=1, keepdim=True) return probs注意:
logits.max(dim=1)默认返回的是一个命名元组(values, indices),如果直接拿去运算会报错或者得到错误结果。记得加.values,或者用两个变量接收:max_logits, _ = logits.max(dim=1)。
2.2 比NaN更隐蔽的坑:keepdim没写导致的形状错乱
这个坑我当年踩过一次之后,到现在都很敏感。logits.max(dim=1)返回的形状是[N],而不是[N, 1]。如果你直接拿它去减logits(形状[N, K]),PyTorch的广播机制会“隐式地”把[N]当成[1, N]处理,实际效果变成每个样本减的max完全错位。
代码长这样时,loss不会报错,但结果全错:
# 错误示范:max_logits 形状是 [N],不是 [N, 1] max_logits = logits.max(dim=1).values logits_shifted = logits - max_logits # 广播规则不是按行减,而是按列对齐正确做法是加keepdim=True,让max结果保持[N, 1],这样才能按行对每个样本的logits做平移。
max_logits = logits.max(dim=1, keepdim=True).values这种“形状细节”在手动实现里到处都是,因为API不会替你操心。用keepdim=True养成习惯,能帮你省下大量Debug时间。
2.3 顺带提一句:temperature参数是干什么的
如果你将来读到知识蒸馏、对比学习之类的论文,会看到softmax带一个温度系数T的版本:
p_i = exp(z_i / T) / Σ_j exp(z_j / T)T=1就是标准softmax;T越大,输出分布越平滑;T越小,分布越尖锐,接近one-hot。它的本质是调节概率分布的“锐度”,改变类别间置信度的差距。手动实现时只需要在exp前对logits除以T即可,这里不展开,但既然写了数值稳定性,多提一嘴这个扩展总没坏处。
3. 交叉熵损失不是“套个公式”那么简单
3.1 从最大似然的角度理解交叉熵
softmax给出概率分布之后,需要设计一个损失函数衡量“预测分布”和“真实标签”的差距。这里最自然的选择是最大似然:我们希望模型给真实类别的概率尽可能高。
对单个样本来说,假设真实类别是y,模型给出的对应概率是p_y。极大化似然就是极大化p_y,等价于极小化−log(p_y)。对所有样本取平均,就是交叉熵损失。
如果写成信息论里的完整形式,是:
L = −Σ_k y_k * log(p_k)y_k是one-hot标签。因为one-hot向量里只有真实类别那一项是1,其余全是0,所以这个求和最终只剩一项:
L = −log(p_y)这就是为什么交叉熵的代码实现常常是“取真实类别的log概率,取负,求平均”。
3.2 手动实现交叉熵的三种写法,以及为什么前两种都有隐患
第一种写法:先算softmax,再取log,再取负。看着直观,但log(0)的问题来了。如果某个类别的概率因为浮点下溢变成0(这在训练后期很常见),log(0)直接是负无穷,loss变成NaN。
于是很多人会加一个小常数:torch.log(probs + 1e-8)。这能暂时避免NaN,但引入的小常数会让梯度变得不精确。别人问起来为什么加1e-8,自己也说不清楚。
第二种写法:用PyTorch的torch.log_softmax。它内部会把log和softmax合并起来算,避免先算概率再取log带来的数值问题。用它的时候,我们只需要把logits传给log_softmax,然后在真实类别位置取值取负:
def manual_cross_entropy(logits, labels): log_probs = torch.log_softmax(logits, dim=1) return -log_probs.gather(1, labels.view(-1, 1)).mean()这是推荐写法,简洁且稳定。但既然标题是“手动实现”,我习惯再往前拆一步,把log_softmax的底层逻辑也写出来,让你看到它到底做了什么。
第三种写法:完全手动log_softmax。利用前面提到的数值稳定技巧,先减max再算exp,最后取log:
def manual_log_softmax(logits): max_logits = logits.max(dim=1, keepdim=True).values logits_shifted = logits - max_logits exp_logits = torch.exp(logits_shifted) sum_exp = exp_logits.sum(dim=1, keepdim=True) log_sum_exp = torch.log(sum_exp) return logits_shifted - log_sum_exp仔细观察这个式子的含义:logits_shifted是减完max的logits,log_sum_exp是所有样本在对应类别上的“归一化分母的对数”。两者相减,得到的就是数值稳定的log_softmax。它本质上就是z − max − log(Σ exp(z − max))。
用这个实现交叉熵:
def manual_cross_entropy(logits, labels): log_probs = manual_log_softmax(logits) return -log_probs.gather(1, labels.view(-1, 1)).mean()这里用到了gather——按索引从每行中把真实类别对应的log概率取出来。标签y是长整型索引,形状[N],需要view(-1, 1)变成[N, 1]才能和[N, K]对齐。
3.3 one-hot标签与索引标签,到底有什么区别
很多教程里手写交叉熵用的是one-hot标签:
loss = -(one_hot * torch.log(probs + 1e-8)).sum(dim=1).mean()这么写不是不行,但有几个明显的缺点:
- 内存浪费。类别数一多,one-hot矩阵比索引向量大得多;
- 需要额外转换一步,代码更啰嗦;
- 必须加epsilon防止log(0),但加了之后数值不精确。
所以在PyTorch中,分类任务的实际标签基本都是索引形式,配合gather或nn.CrossEntropyLoss内部的索引查找来定位真实类别。理解两种标签形式的区别,能帮你避开很多老教程里的坑。
4. 反向传播:softmax那看似吓人的Jacobian矩阵,最终化简成p减去one-hot
4.1 先推导:为什么梯度是 p − y
手动实现softmax回归,最核心也是最劝退的部分是反向传播。很多人在这一步被softmax的Jacobian矩阵吓住——确实,直接对p = softmax(z)求∂p/∂z,需要算一个K×K的矩阵,对角线元素是p_i(1−p_i),非对角线元素是−p_i p_j。但别忘了我们最终关心的是损失L对z的梯度,而不是p对z的梯度。
对单个样本L = −log(p_y),展开来看:
L = −z_y + log(Σ_j exp(z_j))这个形式干净多了。直接对z_y求导:
∂L/∂z_y = −1 + exp(z_y) / Σ_j exp(z_j) = p_y − 1对非真实类别z_j(j≠y)求导:
∂L/∂z_j = exp(z_j) / Σ_k exp(z_k) = p_j把两种情况合起来,用向量表示就是:
∂L/∂z = p − one_hot(y)也就是说,交叉熵损失对logits的梯度,等于“softmax输出减去真实类别的one-hot向量”。这个结果极其简化,训练时只需要一次softmax前向计算,就能得到优雅的梯度形式。这也是softmax配交叉熵如此流行的核心原因——它让反向传播变得异常简洁。
4.2 自定义autograd.Function,完整手动实现forward和backward
手动实现反向传播,最“彻底”的方式是继承torch.autograd.Function,自己写forward和backward。这里以X @ W + b得到logits为例,写一个完整的SoftmaxCrossEntropy算子:
import torch class SoftmaxCrossEntropy(torch.autograd.Function): @staticmethod def forward(ctx, X, W, b, labels): # 前向:算 logits、softmax、交叉熵 logits = X @ W + b max_logits = logits.max(dim=1, keepdim=True).values logits_shifted = logits - max_logits exp_logits = torch.exp(logits_shifted) probs = exp_logits / exp_logits.sum(dim=1, keepdim=True) loss = -torch.log(probs.gather(1, labels.view(-1, 1))).mean() # 保存反向传播需要的中间变量 ctx.save_for_backward(X, W, probs, labels) return loss @staticmethod def backward(ctx, grad_output): X, W, probs, labels = ctx.saved_tensors N = labels.shape[0] # 构造 one-hot 标签 one_hot = torch.zeros_like(probs) one_hot.scatter_(1, labels.view(-1, 1), 1.0) # 核心梯度公式:p - one_hot,再除以N是因为loss取了mean grad_logits = (probs - one_hot) / N # 链式法则回传到各参数 grad_W = X.T @ grad_logits grad_b = grad_logits.sum(dim=0) grad_X = grad_logits @ W.T return grad_X, grad_W, grad_b, None使用方式和普通PyTorch函数一样:
loss = SoftmaxCrossEntropy.apply(X, W, b, y) loss.backward()这个实现里值得注意的细节有三个:
- 为什么用
ctx.save_for_backward而不是普通变量保存?因为PyTorch的自动求图机制在forward之后可能释放中间变量,只有通过save_for_backward保存的张量才能真正在backward阶段被取回,且它是专门为自定义Function设计的安全通道。 - 为什么要除以
N?因为loss用的是mean而不是sum,梯度回传时也要对应地取平均,否则梯度方向没问题但步长会随batch大小变化,导致学习率不稳定。 - 返回的梯度顺序必须和forward输入参数顺序一致。forward是
(X, W, b, labels),backward就返回(grad_X, grad_W, grad_b, None)。labels是整数索引,不需要梯度,所以返回None。
4.3 别以为自己推导对了就是对的:用gradcheck验证
手写backward最怕的就是公式推导没错但代码写错。PyTorch提供了torch.autograd.gradcheck,可以数值化地验证自定义Function的梯度是否正确。这个工具强烈建议用,每次写完自定义算子都跑一遍,能过滤掉绝大多数低级错误。
from torch.autograd import gradcheck X = torch.randn(8, 2, dtype=torch.float64, requires_grad=True) W = torch.randn(2, 3, dtype=torch.float64, requires_grad=True) b = torch.randn(3, dtype=torch.float64, requires_grad=True) y = torch.randint(0, 3, (8,)) gradcheck(SoftmaxCrossEntropy.apply, (X, W, b, y))注意:gradcheck要求输入张量是float64类型,否则数值扰动精度不够,即使梯度是对的也可能报错。这是一个非常容易忽略的细节,我最初用float32跑了半天都是False,换成double立刻通过。
5. 完整训练流程:从数据生成到决策边界可视化
5.1 为什么我选了二维数据而不是MNIST
实现完整训练循环之前,先回答一个很多人会问的问题:为什么不用MNIST?因为MNIST是28×28=784维输入,训练完了也没法直接画图观察模型学到了什么。这里我用make_classification生成一个二维三分类数据集,特征只有两列,训练完之后可以把整个平面的预测结果画出来,直观看到线性决策边界长什么样。如果是图像任务,只需要把输入维度改成784,代码逻辑完全一样。
import torch import matplotlib.pyplot as plt from sklearn.datasets import make_classification torch.manual_seed(42) X_np, y_np = make_classification( n_samples=1500, n_features=2, n_informative=2, n_redundant=0, n_clusters_per_class=1, n_classes=3, random_state=42 ) X = torch.tensor(X_np, dtype=torch.float32) y = torch.tensor(y_np, dtype=torch.long)生成之后记得做标准化。虽然make_classification生成的数据量不大,不标准化也大概率能收敛,但梯度下降对特征的尺度敏感,尺度不一致会导致某些维度更新过快、某些维度更新过慢,收敛速度明显变慢。标准化之后训练会更稳定:
X_mean = X.mean(dim=0, keepdim=True) X_std = X.std(dim=0, keepdim=True) X = (X - X_mean) / X_std5.2 参数初始化、训练循环和损失打印
softmax回归的参数就两个:权重W和偏置b。W形状是[特征数, 类别数],b形状是[类别数]。
初始化方面,我习惯把W设成randn乘个小系数0.1,b设成全零。为什么不全部初始化为零?因为如果所有参数都一样,同一层内所有神经元的学习完全对称,梯度更新也完全一样,模型退化成一个没有区分能力的线性函数。虽然softmax回归只有一个权重矩阵,全零初始化不会像深层网络那样“对称坍缩”,但随机初始化依然是更稳妥的做法,它能让不同类别对应的权重方向在训练一开始就有所不同。
W = torch.randn(2, 3, requires_grad=True) * 0.1 b = torch.zeros(3, requires_grad=True) lr = 0.1 epochs = 400 for epoch in range(epochs): logits = X @ W + b loss = manual_cross_entropy(logits, y) loss.backward() with torch.no_grad(): W -= lr * W.grad b -= lr * b.grad W.grad.zero_() b.grad.zero_() if epoch % 50 == 0: acc = (logits.argmax(dim=1) == y).float().mean().item() print(f"epoch {epoch:3d} | loss {loss.item():.4f} | acc {acc:.4f}")这段代码是最简的SGD手动更新。关键点在于:W -= lr * W.grad必须在torch.no_grad()下执行,否则PyTorch会把梯度更新操作本身也记进计算图,导致内存膨胀甚至报错。更新完之后要手动调用grad.zero_()清零,否则下一轮梯度会和上一轮累积。
我实际跑的时候,前50个epoch损失下降非常快,准确率从随机水平一路冲高;后面曲线逐渐变平缓,400轮之后损失大概稳定在0.25左右,训练准确率能到90%上下。这里额外说一句:学习率0.1是我试过之后选的值。如果调到0.5,前期损失下降更快但容易震荡;调到0.01,收敛会慢不少但不能说错。学习率的选择在简单线性模型上没那么敏感,但训练过程会暴露出来。
5.3 画出决策边界,直观理解线性分类器
训练完之后,把二维平面上每个点都喂进模型,取argmax得到预测类别,然后画等高线图:
xx, yy = torch.meshgrid( torch.linspace(-3, 3, 300), torch.linspace(-3, 3, 300), indexing="ij" ) grid = torch.stack([xx.ravel(), yy.ravel()], dim=1) pred = (grid @ W.detach()).argmax(dim=1).reshape(xx.shape) plt.figure(figsize=(8, 6)) plt.contourf(xx.numpy(), yy.numpy(), pred.numpy(), alpha=0.6) plt.scatter(X[:, 0].numpy(), X[:, 1].numpy(), c=y.numpy(), edgecolor="k", s=20) plt.title("softmax regression decision boundary") plt.show()从图上能明显看到,三个类别的边界是三条直线,两两相交把平面分成三个扇形区域。这正是线性分类器的特征——它只能画直线边界,没法拟合弯曲的分界线。如果你把数据集换成环形或者螺旋形,softmax回归的准确率会掉得很难看,这时候就需要上多层感知机了。
6. 从NaN到shape错位:这次实现中最值得记住的工程细节
6.1 一个亲测有效的排查清单
手动实现过程中,我遇到或预判到的问题不少,这里整理成一份排查清单,按出现频率排序:
| 问题 | 现象 | 常见原因 | 解决办法 |
|---|---|---|---|
| loss变成NaN | 训练几步后loss突然变nan | exp溢出、学习率过大、log(0) | softmax前减max,学习率调小,用log_softmax |
| 准确率一直不涨 | 训练多轮准确率还在随机水平 | 数据未标准化、初始化不当、梯度没更新 | 特征标准化,随机初始化,检查W.grad是否为None |
| shape报错 | gather报dim不匹配或广播结果错乱 | keepdim没写、索引形状不对 | 统一用keepdim=True,labels.view(-1, 1) |
| 自定义backward梯度不对 | gradcheck返回False | 中间变量保存方式错误、忘记除以N | 用ctx.save_for_backward,核对返回顺序 |
| 可复现性差 | 同样的代码每次结果不同 | 没有固定随机种子 | 代码开头设torch.manual_seed与random.seed |
排查梯度更新是否生效,最快捷的方法是在训练循环里打印W.grad的范数。如果范数一直是0,要么前向计算根本没有连接到W,要么在no_grad块里误用了W.grad.zero_()之外的操作把计算图断开了。
6.2 从二维向多维batch演进:softmax的dim到底该选哪个
很多人把二维版本跑通之后,以为就算完事了。但实际工程里,logits绝大多数时候不是[N, K]这么简单。举个NLP里的例子:语言模型的logits形状是[B, T, V]——B是batch size,T是序列长度,V是词表大小。这时候softmax应该作用在V这个维度上,也就是dim=-1或者dim=2。
把前面二维版本改成三维版本,只需要把dim=1全部换成dim=-1,其余逻辑几乎不变:
def log_softmax_3d(logits): max_logits = logits.max(dim=-1, keepdim=True).values logits_shifted = logits - max_logits exp_logits = torch.exp(logits_shifted) sum_exp = exp_logits.sum(dim=-1, keepdim=True) log_sum_exp = torch.log(sum_exp) return logits_shifted - log_sum_exp标签形状从[N]变成[B, T],gather时也要注意对齐。这个扩展很简单,但如果你从一开始就只记“softmax要用dim=1”,到三维场景就会懵。建议现在就记住一个更通用的判断方式:softmax永远作用在“类别分数所在的维度”上,不是固定的第几个维度。
6.3 手动实现与nn.CrossEntropyLoss的差距在哪
最后聊一个很实际的问题:我们手动写的这个版本,和PyTorch官方的nn.CrossEntropyLoss到底差在哪?
差异主要有两点:
第一,官方实现把log_softmax和nll_loss封装在一起,同时支持ignore_index、label_smoothing、class_weight等高级参数。比如类别不平衡时,给少数类更高的权重,手动实现就需要自己乘上权重,比较麻烦。
第二,官方实现用的是torch._C._nn.cross_entropy_loss底层算子,性能做了高度优化;自定义Function走的是Python层面的autograd,虽然功能正确,训练速度会有差距。对softmax回归这种小模型影响不大,但放到大模型场景,性能差距就会暴露出来。
不过差距归差距,我依然强烈推荐每个人至少手写一次。因为用nn.CrossEntropyLoss时,你只是一个“调用者”;手写一遍之后,你才变成“理解者”。遇到loss异常、梯度异常、自定义算子报错时,理解者和调用者的排查速度完全是两个量级。
最后分享一个我自己的排查习惯:以后不管你是自己实现了softmax回归,还是自定义其他算子,先别急着跑完整训练。拿一个小batch、打开torch.autograd.set_detect_anomaly(True)、再配合gradcheck,这三步能过滤掉百分之七八十的数值问题。很多人觉得手动实现这种基础模型不过是“把公式翻译成代码”,真正动手以后才发现,每一个看似简单的公式背后都藏着数值稳定性、形状对齐、梯度链路这些实际工程问题。希望这篇文章能帮你少走一点弯路。