ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

手写Softmax回归:从数学推导到PyTorch自定义反向传播

手写Softmax回归:从数学推导到PyTorch自定义反向传播 我先说个反直觉的事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_isoftmax把这个分数向量转成概率分布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是1000exp(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(dim1, keepdimTrue).values logits_shifted logits - max_logits exp_logits torch.exp(logits_shifted) probs exp_logits / exp_logits.sum(dim1, keepdimTrue) return probs注意logits.max(dim1)默认返回的是一个命名元组(values, indices)如果直接拿去运算会报错或者得到错误结果。记得加.values或者用两个变量接收max_logits, _ logits.max(dim1)。2.2 比NaN更隐蔽的坑keepdim没写导致的形状错乱这个坑我当年踩过一次之后到现在都很敏感。logits.max(dim1)返回的形状是[N]而不是[N, 1]。如果你直接拿它去减logits形状[N, K]PyTorch的广播机制会“隐式地”把[N]当成[1, N]处理实际效果变成每个样本减的max完全错位。代码长这样时loss不会报错但结果全错# 错误示范max_logits 形状是 [N]不是 [N, 1] max_logits logits.max(dim1).values logits_shifted logits - max_logits # 广播规则不是按行减而是按列对齐正确做法是加keepdimTrue让max结果保持[N, 1]这样才能按行对每个样本的logits做平移。max_logits logits.max(dim1, keepdimTrue).values这种“形状细节”在手动实现里到处都是因为API不会替你操心。用keepdimTrue养成习惯能帮你省下大量Debug时间。2.3 顺带提一句temperature参数是干什么的如果你将来读到知识蒸馏、对比学习之类的论文会看到softmax带一个温度系数T的版本p_i exp(z_i / T) / Σ_j exp(z_j / T)T1就是标准softmaxT越大输出分布越平滑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, dim1) return -log_probs.gather(1, labels.view(-1, 1)).mean()这是推荐写法简洁且稳定。但既然标题是“手动实现”我习惯再往前拆一步把log_softmax的底层逻辑也写出来让你看到它到底做了什么。第三种写法完全手动log_softmax。利用前面提到的数值稳定技巧先减max再算exp最后取logdef manual_log_softmax(logits): max_logits logits.max(dim1, keepdimTrue).values logits_shifted logits - max_logits exp_logits torch.exp(logits_shifted) sum_exp exp_logits.sum(dim1, keepdimTrue) log_sum_exp torch.log(sum_exp) return logits_shifted - log_sum_exp仔细观察这个式子的含义logits_shifted是减完max的logitslog_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(dim1).mean()这么写不是不行但有几个明显的缺点内存浪费。类别数一多one-hot矩阵比索引向量大得多需要额外转换一步代码更啰嗦必须加epsilon防止log(0)但加了之后数值不精确。所以在PyTorch中分类任务的实际标签基本都是索引形式配合gather或nn.CrossEntropyLoss内部的索引查找来定位真实类别。理解两种标签形式的区别能帮你避开很多老教程里的坑。4. 反向传播softmax那看似吓人的Jacobian矩阵最终化简成p减去one-hot4.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_jj≠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(dim1, keepdimTrue).values logits_shifted logits - max_logits exp_logits torch.exp(logits_shifted) probs exp_logits / exp_logits.sum(dim1, keepdimTrue) 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(dim0) 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, dtypetorch.float64, requires_gradTrue) W torch.randn(2, 3, dtypetorch.float64, requires_gradTrue) b torch.randn(3, dtypetorch.float64, requires_gradTrue) 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×28784维输入训练完了也没法直接画图观察模型学到了什么。这里我用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_samples1500, n_features2, n_informative2, n_redundant0, n_clusters_per_class1, n_classes3, random_state42 ) X torch.tensor(X_np, dtypetorch.float32) y torch.tensor(y_np, dtypetorch.long)生成之后记得做标准化。虽然make_classification生成的数据量不大不标准化也大概率能收敛但梯度下降对特征的尺度敏感尺度不一致会导致某些维度更新过快、某些维度更新过慢收敛速度明显变慢。标准化之后训练会更稳定X_mean X.mean(dim0, keepdimTrue) X_std X.std(dim0, keepdimTrue) X (X - X_mean) / X_std5.2 参数初始化、训练循环和损失打印softmax回归的参数就两个权重W和偏置b。W形状是[特征数, 类别数]b形状是[类别数]。初始化方面我习惯把W设成randn乘个小系数0.1b设成全零。为什么不全部初始化为零因为如果所有参数都一样同一层内所有神经元的学习完全对称梯度更新也完全一样模型退化成一个没有区分能力的线性函数。虽然softmax回归只有一个权重矩阵全零初始化不会像深层网络那样“对称坍缩”但随机初始化依然是更稳妥的做法它能让不同类别对应的权重方向在训练一开始就有所不同。W torch.randn(2, 3, requires_gradTrue) * 0.1 b torch.zeros(3, requires_gradTrue) 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(dim1) y).float().mean().item() print(fepoch {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), indexingij ) grid torch.stack([xx.ravel(), yy.ravel()], dim1) pred (grid W.detach()).argmax(dim1).reshape(xx.shape) plt.figure(figsize(8, 6)) plt.contourf(xx.numpy(), yy.numpy(), pred.numpy(), alpha0.6) plt.scatter(X[:, 0].numpy(), X[:, 1].numpy(), cy.numpy(), edgecolork, s20) plt.title(softmax regression decision boundary) plt.show()从图上能明显看到三个类别的边界是三条直线两两相交把平面分成三个扇形区域。这正是线性分类器的特征——它只能画直线边界没法拟合弯曲的分界线。如果你把数据集换成环形或者螺旋形softmax回归的准确率会掉得很难看这时候就需要上多层感知机了。6. 从NaN到shape错位这次实现中最值得记住的工程细节6.1 一个亲测有效的排查清单手动实现过程中我遇到或预判到的问题不少这里整理成一份排查清单按出现频率排序问题现象常见原因解决办法loss变成NaN训练几步后loss突然变nanexp溢出、学习率过大、log(0)softmax前减max学习率调小用log_softmax准确率一直不涨训练多轮准确率还在随机水平数据未标准化、初始化不当、梯度没更新特征标准化随机初始化检查W.grad是否为Noneshape报错gather报dim不匹配或广播结果错乱keepdim没写、索引形状不对统一用keepdimTruelabels.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 sizeT是序列长度V是词表大小。这时候softmax应该作用在V这个维度上也就是dim-1或者dim2。把前面二维版本改成三维版本只需要把dim1全部换成dim-1其余逻辑几乎不变def log_softmax_3d(logits): max_logits logits.max(dim-1, keepdimTrue).values logits_shifted logits - max_logits exp_logits torch.exp(logits_shifted) sum_exp exp_logits.sum(dim-1, keepdimTrue) log_sum_exp torch.log(sum_exp) return logits_shifted - log_sum_exp标签形状从[N]变成[B, T]gather时也要注意对齐。这个扩展很简单但如果你从一开始就只记“softmax要用dim1”到三维场景就会懵。建议现在就记住一个更通用的判断方式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这三步能过滤掉百分之七八十的数值问题。很多人觉得手动实现这种基础模型不过是“把公式翻译成代码”真正动手以后才发现每一个看似简单的公式背后都藏着数值稳定性、形状对齐、梯度链路这些实际工程问题。希望这篇文章能帮你少走一点弯路。
返回列表