ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Softmax 函数详解:从数学原理到工程实践的完整指南

Softmax 函数详解:从数学原理到工程实践的完整指南 做分类任务绕不开 Softmax。我接触机器学习后的第一个多分类模型是一个接在卷积网络后面的两层结构全连接层输出 10 个分数后面跟 Softmax然后拿交叉熵算损失。那时候我只把 Softmax 当成“把分数变成概率的开关”直到后来自己动手推梯度、踩数值溢出、在注意力机制和知识蒸馏里反复看见它才明白这个看起来只有一行公式的函数几乎是整个深度学习中概率建模的基石。Softmax 是机器学习里使用频率最高的激活函数之一主要用在多分类任务的输出层。它做的事情很具体把网络输出的 K 个实数这些数通常叫 logits也就是逻辑值转换成 K 个非负、并且加起来等于 1 的概率值。它解决的核心需求就是——当模型面对“这张图是猫、是狗还是鸟”的问题时它不能只丢给你一个硬编码标签而要给出一个带有置信度含义的概率分布。这篇文章会从数学定义出发把公式背后的直觉、工程实现里的数值稳定性、和交叉熵搭配时的原理以及实际项目里的高频问题、调试经验全部串一遍。适合正在补机器学习数学基础的同学、准备算法面试的人以及做分类项目时想把细节搞清楚的从业者。1. 从二分类到多分类Softmax 到底解决了什么问题1.1 为什么 Sigmoid 不够用要理解 Softmax最好先回到 Sigmoid。二分类场景里模型只需要判断“是”与“否”常用 Sigmoid 把一个实数压到 (0,1) 区间作为正类的概率。它天然满足“输出在 0 到 1 之间”这个概率约束并且配合交叉熵使用时梯度形式非常漂亮。但到了多分类问题变复杂了。假设有猫、狗、鸟三类一个最朴素的思路是训练三个独立的 Sigmoid 分类器猫分类器输出“是不是猫”狗分类器输出“是不是狗”鸟分类器输出“是不是鸟”。这样做在数学上会遇到两个麻烦。第一个麻烦是输出不再构成一个概率分布。三个 Sigmoid 的结果分别是 0.8、0.4、0.6加起来是 1.8你没法直接解释成“模型认为猫的概率是 0.8”。你当然可以事后做一次归一化比如除以它们的和变成 0.44、0.22、0.33但这种后处理没有任何理论依据——训练目标里并没有要求三个 Sigmoid 的输出之间互相约束。第二个麻烦是梯度更新不一致。Sigmoid 对每个类别单独做判断相当于把多分类问题强行拆成了多个二分类问题。这个做法并不是完全不能用多标签分类就是这么干的但对于“互斥的单标签分类”它忽略了“选猫就不能选狗”这个关键结构。模型可能会同时把猫和狗的分数都推得很高而真正的多分类器需要学会在类别之间做竞争和取舍。所以我们需要一个新函数它接收 K 个任意实数输出 K 个非负的数并且这 K 个数之和恒为 1。这就是 Softmax 的设计目标。1.2 从 max 到 softmax一次“软化”的巧妙设计先想一想如果我们足够“硬”多分类的判决规则就是 argmax把最大分数对应的类别选出来其他全部归零。这种方法当然有效argmax 也是很多模型最后做预测时的最终步骤。但它有两个致命问题不可导或者说几乎处处导数为 0没法通过反向传播来训练而且“赢家通吃”完全丢掉了模型对其他类别的判断信息。Softmax 就是在 max 和 argmax 之间做了一个连续化平滑。它用指数的形式把每个类别都保留了一点权重最大的一支获得最高概率但其他类别依然分到少量概率而不是被彻底清零。这种“既突出最大、又不放弃其余”的特性就是 “soft” 这个词的含义。数学上Softmax 还有一种更深刻的解读它其实是对带噪声的 argmax 的一种吉布斯分布温度参数 T 控制噪声尺度。当 T 趋近于 0 时Softmax 趋近于 argmax输出近似 one-hot当 T 趋近于无穷大时输出趋近于均匀分布。这个视角在后面讲知识蒸馏时会非常有用。我们通常写 Softmax 时不显式写出 T相当于默认 T1但它内在地蕴含着“软化程度”的可调维度。2. Softmax 的数学原理公式拆解与直觉理解2.1 公式、计算步骤与数值示例Softmax 的公式长这样[ p_i \frac{\exp(z_i)}{\sum_{j1}^{K} \exp(z_j)}, \quad i 1, 2, \dots, K ]其中 z 是模型最后一层输出的 logits 向量K 是类别数。给定一个输入样本我们先用网络算出每个类别的得分再对得分做指数运算最后把所有指数结果加起来做分母每个指数结果除以分母就得到第 i 类的概率。这里有个细节Softmax 通常接在模型最后一层线性层之后它本身没有可学习参数。它只是一个确定性的归一化算子把“分数”翻译成“概率”。为了直观我们算一个最简单的例子。假设三分类模型对某张图片输出了 logits 为 [2.0, 1.0, 0.0]。第一步算 exp类别logits zexp(z)概率 p猫2.07.3890.665狗1.02.7180.245鸟0.01.0000.090求和得到 S ≈ 11.107。然后分别做除法p1 7.389 / 11.107 ≈ 0.665p2 2.718 / 11.107 ≈ 0.245p3 1.000 / 11.107 ≈ 0.090。所以模型认为这张图有 66.5% 概率属于第一类。注意即使第三类的 logits 只有 0它的概率也不是 0因为 softmax 永远会给每个类别分配一个正概率只是大小不同而已。这一点在“不设未知类”的闭集分类里没什么问题但在开放场景下要注意下面第 5 节会专门展开。2.2 指数运算在做什么放大差异但保留全貌为什么用指数而不是直接对 logits 做归一化比如 z_i / sum(z_j)这个问题我在面试别人时经常问也是理解 Softmax 的关键点。直接线性归一化的最大问题是当 logits 出现负值时输出会包含负数或者变成无意义的比例而且线性归一化对“类别间差距”不敏感。Softmax 使用指数函数有两个明显的作用。第一指数函数把实数轴映射到正半轴 (0, ∞)这样保证了每个输出都为正再配合分母归一化就天然满足“非负且和为 1”的概率公理。第二指数运算放大了 logits 之间的相对差异。还是用刚才的例子原始 logits 差异是 2.0、1.0、0.0经过指数之后变成了 7.39、2.72、1.00原来的倍数差异被急剧拉大最终的输出概率差距也更明显。这相当于让模型在类别之间做出更“自信”的选择。这种放大能力在分类边界附近特别重要训练前期 logits 可能都在 0 附近徘徊如果直接用线性归一化概率分布非常平梯度也不容易拉开指数运算能尽早把正类别和负类别之间的差距暴露出来帮助模型加速收敛。但需要注意的是指数放大的能力也意味着 Softmax 对 logits 的绝对值大小很敏感。如果 logits 整体被放大 10 倍输出概率会迅速逼近 one-hot 分布趋近于“硬分类”。这在某些场景下是好事但在另一些场景下会导致过度自信。这也是为什么我们在做模型校准时会给 Softmax 引入一个温度参数 T用 z/T 来控制它的“锐度”后面第 5 节还会再提到。3. 从理论到代码数值稳定性与工程实现3.1 数值溢出最经典的坑理论上的公式很简单但直接照搬到代码里就会踩到数值坑。最大问题是 exp 的爆炸性增长。当 logits 中有一个数是 1000 时exp(1000) 在双精度浮点数里就变成了 inf再往下算分母也全是 inf最终得到的结果就是 nan。这个问题的标准解法是“减去最大值”。具体做法是先算出 logits 向量里的最大值 m然后把 Softmax 的公式改写为[ p_i \frac{\exp(z_i - m)}{\sum_{j1}^{K} \exp(z_j - m)} ]为什么要减去最大值因为 Softmax 对输入向量的整体平移是不变的分子分母同时乘以 exp(-m)结果完全一样。从数学上可以严格证明把 z_i 替换成 z_i - m本质上就是在分子分母上都乘了同一个非零常数 exp(-m)约掉之后与原来的 Softmax 完全等价。所以减去最大值不会改变输出概率却能把指数运算的输入限制在 (-∞, 0] 范围内避免 exp 溢出。这也是很多初学 PyTorch 的人容易忽略的一步。如果你自己手写 Softmax 并且没有做减最大值的处理训练的时候网络数值一旦变大loss 立刻变成 nan而且这种 nan 不是随机出现的一旦出现基本就要重训。3.2 PyTorch 里的三种实现方式工程上我们有三种常见写法从“手写教程版”到“生产级”我按顺序列一下。第一种是手动实现适合写作业、做推导、验证理解import torch def softmax_manual(x, dim-1): x_max x.max(dimdim, keepdimTrue).values exp_x torch.exp(x - x_max) return exp_x / exp_x.sum(dimdim, keepdimTrue)这里必须注意 dim 参数。如果你处理的是 (batch, num_classes) 形状的 logits那么 dim-1 或者 dim1 都是对的但如果输入是 (batch, seq_len, num_classes) 这种三维向量一定要想清楚在哪一维上做归一化通常是对最后一维类别维做。第二种是使用 log_softmax典型场景是在自定义损失逻辑时使用因为直接对概率做 log 可能会因为概率为 0 而得到负无穷。torch.nn.functional.log_softmax 把 exp 和 log 融合起来数值上更稳定log_probs torch.nn.functional.log_softmax(logits, dim-1) # 然后用 log_probs 和标签构造 NLLLoss第三种是直接用交叉熵损失函数。它内部已经完成了 log_softmax 和 NLLLoss 的融合这是训练分类网络时最推荐的写法loss torch.nn.functional.cross_entropy(logits, target)注意这里传入的是 logits而不是 softmax 之后的结果。很多人一开始会犯一个错先对 logits 做 softmax再把结果传给 cross_entropy结果导致模型不收敛或者训练曲线异常。因为 cross_entropy 内部做 log_softmax 时会先假设输入是 logits如果你传进来的已经是概率相当于对概率又做了一次 log_softmax数值和信息都被破坏了。4. 训练中的黄金搭档Softmax 与交叉熵损失函数4.1 为什么分类任务通常配交叉熵Softmax 输出的是一组概率也就是一个离散概率分布。训练时我们希望模型输出的分布尽量接近真实标签的分布。对于单标签分类真实标签通常表示成一个 one-hot 向量比如第二类是正确类别那么真实分布就是 [0, 1, 0, 0, ...]。衡量两个分布之间的距离最常用的指标是交叉熵Cross Entropy它和 KL 散度只差一个常数项。从最大似然估计的角度看最小化交叉熵就是在最大化正确类别的对数似然这在统计上是最自然的准则。交叉熵还有一个工程上的优点它的梯度形态非常好。如果换成均方误差MSE来约束 softmax 的输出你会发现损失函数对 logits 的梯度并不直接并且当 softmax 输出接近 0 或 1 时梯度会因为 sigmoid 类的饱和效应而变得很小。交叉熵则完全不是这样这就是下一节要说的“简单到令人惊讶的梯度”。4.2 简单到令人惊讶的梯度我们推一下这个推导这也是“机器学习中的数学”系列最值得沉淀的部分。假设交叉熵损失为[ L -\log p_y ]其中 p_y 是正确类别的预测概率。把 Softmax 定义代入可以写成[ L -z_y \log\left(\sum_{j1}^{K} \exp(z_j)\right) ]对任意的第 i 类 logits z_i 求偏导。分两种情况。如果 i y即对正确类别的得分求偏导链式法则得到[ \frac{\partial L}{\partial z_i} -1 \frac{\exp(z_i)}{\sum_j \exp(z_j)} p_i - 1 ]如果 i ≠ y即对其他类别的得分求偏导[ \frac{\partial L}{\partial z_i} \frac{\exp(z_i)}{\sum_j \exp(z_j)} p_i ]合起来就是[ \frac{\partial L}{\partial z_i} p_i - y_i ]这个结果漂亮得不像话交叉熵 Softmax 对 logits 的梯度就是“模型当前预测的概率 p_i 减去真实 one-hot 标签 y_i”。也就是说梯度幅度直接等于模型犯错的幅度。这意味着什么呢当模型对正确类别不太自信、p_y 较小比如 0.3 时梯度是 0.3 - 1 -0.7幅度很大模型会大步调整当模型已经猜对了且非常自信p_y 接近 0.95 时梯度只有 0.95 - 1 -0.05收敛放缓。正类、负类的梯度方向也完全正确正类得分高了要持续推高一点负类得分高了要往下降一点。这种“错误越大、更新越大”的性质让训练非常稳定高效。这也是为什么在分类问题里大家几乎不会用 MSE 去配 Softmax 的原因——不是不可以而是效果明显差一截。5. Softmax 的应用场景从分类器到注意力机制5.1 分类任务中的标准角色最经典的使用场景就是图像分类、文本分类、语音分类等单标签分类任务。模型骨干网络提取特征后接一个线性层把特征映射成类别分数再上 Softmax 得到概率最后用交叉熵训练。这里有个很容易混淆的点多标签分类任务比如一张图里有猫也有狗不能用 Softmax而应该用多个 Sigmoid因为每个标签是否出现是独立事件类别之间并不是互斥的。当类别数量变得非常大时比如面对上百万个类别的商品识别场景直接对全类别做 Softmax 的计算开销会很大。这时一般会用采样 SoftmaxSampled Softmax比如随机采样出一部分负类来近似计算梯度。这种情况下 Softmax 公式本身没变只是训练时用近似代替了全量计算。5.2 注意力机制与知识蒸馏Softmax 的身影远不止输出层。Transformer 里的自注意力机制核心公式是 attention weights softmax(QK^T / sqrt(d))它做的事情是把“查询”和“键”之间的匹配分数归一化成一堆权重然后对“值”做加权求和。这里 Softmax 的作用和分类任务里完全一致把任意实数的相似度分数转换成非负且和为 1 的权重分布。在知识蒸馏中Softmax 还会被刻意加上温度 T变成 softmax(z/T)。温度越高输出的概率分布越平坦小类别也能获得少量概率温度越低输出越接近 one-hot。蒸馏时用较高的温度让教师模型暴露出“软标签”——比如猫的图片在狗类别上也有一点概率——这是类间结构的重要信息学生模型通过学习这些软标签能学会教师模型的泛化能力。这个技术能跑通靠的就是 Softmax 的“软化”特性。5.3 别把 Softmax 概率当绝对置信度Softmax 输出的概率虽然语义上是概率但在实际应用中需要谨慎对待。很多分类模型在训练充分之后对训练集里的样本会输出接近 1 的概率甚至对分布外OOD的样本也一样自信。这种现象叫“过度自信”over-confidence本质原因是交叉熵训练只会鼓励正确类别的概率变高并不保证其他输入的概率校准得很好。如果模型要用在真实业务里尤其需要拿概率来做阈值判断、风险控制时建议做一次“温度缩放”Temperature Scaling校准。做法很简单在验证集上搜索一个温度 T让 softmax(logits / T) 的置信度与真实准确率更贴近。这个 T 是在模型训练完成后加的不修改模型参数只调整 softmax 的锐度。另外一个常见技巧是标签平滑Label Smoothing在构造 one-hot 标签时给正确类别留一点概率给其他类别从而抑制模型过分自信这在很多大模型训练里已经是标配。6. 常见问题与避坑指南6.1 新手最常犯的 3 个错误第一个错误是把 Softmax 的结果送入交叉熵损失。前面讲过PyTorch 的 cross_entropy 内部已经做了 log_softmax直接传入 logits 即可。如果你在训练代码里看到 loss F.cross_entropy(F.softmax(logits), target)立刻改成 loss F.cross_entropy(logits, target)。这个错误最坑的地方在于模型可能也能训出一定的准确率但曲线会很不稳定而且收敛速度明显变慢。第二个错误是 dim 参数传错。Softmax 的归一化维度必须是类别维度。假设张量形状是 (batch, num_classes)请务必用 dim1 或 dim-1而不是 dim0。如果不小心写成 dim0那是在把所有样本、所有类别的分数一起归一化整个概率输出就完全错了而且这种错误在训练中不容易被发现因为 loss 还是能下降。第三个错误是在多标签任务里误用 Softmax。当每个样本可以同时拥有多个标签时Softmax 的“和为 1”约束会强制模型在标签之间做竞争导致同一个样本出现两个高置信度标签时损失极大、模型不敢输出多标签。正确的做法是用 Sigmoid 加 Binary Cross Entropy每个输出维度都独立判断“这个标签是否出现”。6.2 让训练更稳定的调试经验我在实际项目中调试 Softmax 相关问题时一般会先检查三件事。第一看 logits 的数值范围。如果 logits 绝对值非常小比如徘徊在 0.01 附近说明模型对分类没有信心这时可以检查最后一层线性层的初始化是否太小或者特征提取器的学习率是否设置过低。如果 logits 绝对值非常大比如超过 20说明模型过早进入了“过度自信区”概率输出接近 one-hot梯度会很小这时候可以适当调低学习率或者引入标签平滑。第二观察训练初期 loss 的合理范围。对一个 C 类分类问题如果模型初始输出接近均匀分布那么交叉熵 loss 应该接近 log(C)。如果训练一开始 loss 就那么小或者那么大往往意味着初始化有问题。举个例子C10 时初始 loss 如果只有 0.5说明模型已经开始过度自信这时候可以怀疑是不是把 softmax 结果喂给了交叉熵或者学习率设得太大。第三推理时如果需要概率做阈值控制不要直接用分类器给的原始概率先在验证集上做一个温度缩放实验。这个操作成本很低却经常能带来明显的置信度校准效果对后续规则联动、人机协同流程帮助很大。6.3 我的个人心得最后分享几个我在实际操作中养成的习惯。第一个习惯是只要自定义损失函数或自定义模型分支我都会先写一个最小复现用例把 logits 固定成几个已知数值手动算出 softmax 和梯度再去和 PyTorch 的输出做对比。这种“对答案”的方法能帮我快速确认维度、归一化方向、数值稳定性都没问题。第二个习惯是在保存模型做推理时我会把最后的 softmax 和温度参数写进预处理/后处理配置里而不是丢在模型代码里。这样调温度、调阈值时不需要重新加载模型改一行配置就行。第三个习惯是面试或带新人时我会问一个“偏门”问题如果 logits 是 [1000, 999, 998]手写 softmax 会输出什么这看起来简单其实测试的是对数值稳定性、指数放大效应和归一化本质的理解。能一次答对的人通常对 Softmax 的理解已经超过“套公式”的层面了。Softmax 表面上只是一个归一化算子但它连接着概率建模、损失函数设计、数值计算、模型校准一整条知识链。把它的数学本质和工程细节吃透了很多看似玄学的训练问题都会变得清晰很多。这也是我在项目实战中反复验证后的最大体会模型能不能训好往往不在于某个花哨的网络结构而在于这些最基础的组件你是否真的用明白了。
返回列表