资讯动态

Softmax在大模型中的原理与工程实践:从注意力到采样

发布时间:2026/9/4 6:53:28 来源:尧图企业网站定制
每当看到“全球大模型都在用他的公式却没人知道他的名字”这类描述时如果把关注点放在技术本身会发现一个反复出现的关键角色Softmax。从 ChatGPT 背后的 Transformer到千问、LLaMA 等开源模型的生成流程再到本地部署大模型时看到的 sampled token几乎每一步都在调用这个公式。它把模型内部一整排任意大小的“打分”变成一组可以比较、可以采样、可以计算损失的概率。很多开发者熟悉torch.softmax的调用方式却未必清楚它在大模型里到底承担了多少工作。注意力权重由它生成自回归 Mask 通过它实现温度参数在它上面调整交叉熵损失也在它上面展开。这篇文章从最小 Python 实现讲起覆盖数值稳定性、注意力 Mask、FlashAttention 中的 Online Softmax、生成采样温度以及部署大模型时的排错思路。看完之后你能写出一个不溢出的 Softmax能看懂 Transformer 注意力里的形状变化也能在模型输出 NaN 或出现重复文本时快速定位问题是否出在这个公式周围。1. 先看清 Softmax 在整个大模型里的位置1.1 大模型为什么要概率化输出自回归大模型预测下一个 token 时神经网络的最后一层通常只输出一个向量这个向量叫 logits。如果词表大小是 32000logits 就是 32000 个实数。每个位置的数字表示当前 token 的“相对得分”数字越大模型越倾向于选择它。但 logits 不能直接当成概率用。原因有两点。第一logits 的范围不受约束可能是正数、负数也可能是很大或很小的值。第二原始 logits 没有“整体为一”的约束。两个不同输入预测出来的 logits 分数量纲可能完全不同难以横向比较。Softmax 提供的标准做法是先对所有 logits 做指数变换再除以指数之和使输出落在(0,1)区间并且所有维度的输出加起来等于 1。在注意力机制中Softmax 更加重要。Transformer 计算注意力权重时Q 和 K 的点击分数同样是任意实数必须经过 Softmax 变成非负且和为 1 的权重才能用于对 V 做加权求和。所以无论生成阶段还是注意力阶段Softmax 都像“概率翻译官”把打分翻译成概率或权重。1.2 一个最小实现从 logits 到概率先写一个最容易理解的朴素版本。输入一组分数输出每个位置的概率。import numpy as np def softmax_naive(logits): exp_logits np.exp(logits) return exp_logits / exp_logits.sum(axis-1, keepdimsTrue) logits np.array([2.0, 1.0, 0.1]) probs softmax_naive(logits) print(probs) print(probs.sum())运行结果类似[0.65900114 0.24243297 0.09856589] 1.0这个例子说明logits 相差越大Softmax 输出的概率差距也越大。2.0和1.0只相差1.0但概率差距不只是 0.1而是接近一倍多。这种“指数放大”能力让模型能够对候选词做显著区分。但由于计算中使用np.exp一旦 logits 里出现1000朴素版本会立刻溢出这就是后面要解决的数值稳定性问题。1.3 这个公式和 Sigmoid、argmax 有什么区别有些开发者容易把 Softmax 和 Sigmoid 混淆。Sigmoid 把一个实数映射到(0,1)常用于二分类或独立多标签分类。Softmax 接收一个向量输出一个向量且所有输出之和为 1更适合互斥类别之间的选择。大模型每个位置只能选一个 token所以用的是 Softmax 而不是一组 Sigmoid。argmax 则完全放弃概率信息。argmax只告诉你哪个位置最大不告诉你第二、第三名的差距。如果只用贪心解码模型好像只需要 argmax但训练阶段无法对 argmax 求导采样阶段也需要概率分布来控制多样性。Softmax 保留了完整的分布信息又保持了可微性这是它成为主流的原因。操作输入输出主要用途Sigmoid一个实数(0,1)的一个数二分类、多标签独立判断Softmax一个向量和为 1 的概率向量互斥类别分类、注意力权重argmax一个向量最大位置下标贪心解码、硬选择LogSoftmax一个向量log 概率向量数值稳定损失计算2. 在注意力机制里Softmax 承担了什么角色2.1 从 Q、K 的点积到注意力权重Transformer 的注意力公式可以写成Attention(Q, K, V) softmax(Q K.T / sqrt(d_k)) V其中 Q 是查询向量来自当前需要“查找信息”的位置K 是键向量来自被查询的上下文位置V 是值向量保存上下文位置真正要携带的信息。模型首先用 Q 和 K 做点积得到一个相似度矩阵。相似度越高说明当前 query 越应该关注对应 key。点积分数并不是权重必须在最后一个维度上做 Softmax 归一化。经过 Softmax 后每行和为 1再去乘 V相当于对 V 做带权求和。某个位置的权重越高最终输出越偏向这个位置的值。这也解释了为什么大模型需要多头注意力。不同头可以做不同位置的 Softmax一个头关注语法关系另一个头关注指代关系。每个头内部都有完整的Q - K - 相似度 - Softmax - 加权 V链路。2.2 为什么缩放因子 sqrt(d_k) 不是可选项公式中Q K.T之后要除以sqrt(d_k)。如果不除点积的方差会随着维度d_k增大而变大。假设 Q 和 K 每个分量是均值为 0、方差为 1 的随机变量那么两个向量点积的方差近似为d_k。维度越大点积数值的波动范围越大。极端情况下很多点积会落在一个很小的区域Softmax 会进入饱和区产生近似 one-hot 的分布。大部分概率被最大的那个位置拿走剩余位置的梯度非常小模型难以学习。除以sqrt(d_k)后点积的方差回到接近 1 的水平Softmax 输入分布更平滑梯度也更容易传导。所以这个缩放因子不是某种理论装饰而是稳定的必要条件。d_k在标准 Transformer 中常等于d_model / num_heads比如 hidden size 是 4096、头数是 32每个头的d_k就是 128。2.3 自回归 Mask 为什么用 -inf 而不是 0自回归模型预测第 t 个 token 时不能看到第 t 个之后的 token。实现这个限制最常用做法是构造一个上三角 Mask把未来位置遮住再送入 Softmax。很多初学者想当然地写出scores scores.masked_fill(future_mask 0, 0)这不会让未来位置失效。Softmax 的分子对所有输入都做了指数变换0经过exp(0)后变成1未来位置仍会被赋予非零权重。正确做法是scores scores.masked_fill(future_mask 0, float(-inf))因为exp(-inf)等于 0未来位置在归一化后贡献正好为 0。这样既保住了当前位置和过去位置的信息又严格防止了标签泄漏。除了因果 MaskPadding Mask 也常用同样思路。批量训练中不同样本长度不同短序列需要在末尾填充 padding token。如果不屏蔽 pad 位置模型会把填充位置也当成有效上下文。Padding Mask 对非有效位置填-inf再送进 Softmax权重就只在真实 token 上分配。2.4 一个可以直接运行的迷你注意力代码下面用 PyTorch 实现一个简化版注意力。它不考虑多头复杂装配只保留 Q、K、V 和可选 Mask。import math import torch def attention(q, k, v, maskNone): # q, k, v 的形状均为 (batch, heads, seq_len, head_dim) d_k q.shape[-1] scores q k.transpose(-2, -1) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) weights torch.softmax(scores, dim-1) output weights v return output, weights torch.manual_seed(0) batch, heads, seq_len, head_dim 1, 2, 4, 8 q torch.randn(batch, heads, seq_len, head_dim) k torch.randn(batch, heads, seq_len, head_dim) v torch.randn(batch, heads, seq_len, head_dim) # 因果 mask只允许位置 i 看到 0..i 的位置 causal_mask torch.tril(torch.ones(seq_len, seq_len)).bool() causal_mask causal_mask.reshape(1, 1, seq_len, seq_len) output, weights attention(q, k, v, maskcausal_mask) print(output shape:, output.shape) print(weights shape:, weights.shape) print(每行权重之和:, weights[0, 0].sum(dim-1))这段代码展示了几个关键点。q k.transpose(-2, -1)会得到形状为(batch, heads, seq_len, seq_len)的注意力分数。mask 0的部分会先填充为-inf再进入softmax被遮挡位置权重为 0。dim-1指定在“被查询的序列位置”维度上归一化这个维度对应每一行的所有 key。运行后应看到每行权重和都接近 1。如果 Mask 设置错误例如使用masked_fill(mask 0, 0)未来行权重不会归零后续 loss 会不正常偏低生成时模型表现却很差。这里是最值得加单测的位置。3. 数值稳定性才是 Softmax 的工程第一课3.1 朴素版本在极端 logits 下会发生什么朴素 Softmax 看起来简单但有明显数学缺陷。import numpy as np logits np.array([1000.0, 1001.0, 999.0]) print(np.exp(logits))在 Python 的 float64 中exp(1000)会直接溢出为inf。在 float16 中问题出现得更早。float16 能表示的最大有限值是 65504所以exp(15)左右就可能触发无穷大。大模型训练和推理中一旦出现极端 logits朴素 Softmax 会产生 NaN反向传播时梯度也会变成 NaN。实际运行时模型输出很少会达到 1000但在混合精度训练、大学习率、特殊权重初始化等场景下logits 可能出现发散。如果代码里每一步都先做np.exp(logits)再求和就等于把数值鲁棒性交给运气。3.2 减去 max 的稳定版本为什么成立Softmax 有一个良好性质对 logits 同时减去同一个常数结果不变。证明思路很简单。设c是任意常数那么第 i 个位置的权重为exp(z_i - c) / sum_j exp(z_j - c)分子和分母同时乘以exp(c)就可以约回原来的exp(z_i) / sum_j exp(z_j)。因此只要让c等于这批 logits 的最大值所有分数都会变成小于等于 0 的数指数运算不会再产生无穷大的exp。稳定版实现如下。import numpy as np def softmax_stable(logits, axis-1): logits np.asarray(logits) shifted logits - np.max(logits, axisaxis, keepdimsTrue) exp_logits np.exp(shifted) return exp_logits / exp_logits.sum(axisaxis, keepdimsTrue) print(softmax_stable([1000.0, 1001.0, 999.0]))这段代码中np.max会找到最后一个维度上的最大值keepdimsTrue保持维度结构以便广播。减去最大值后极大输入变成[ -1, 0, -2 ]指数计算安全。在 PyTorch 中torch.softmax内部已经做了类似优化所以日常调用不需要手动做稳定化。但如果你要自定义注意力 kernel或者在纯 NumPy 里实现模型就必须使用减去 max 的写法。3.3 微调时建议使用 LogSoftmax 与交叉熵大模型微调最常用的损失是交叉熵。给定一个输入 token 序列标准做法是让模型预测下一个 token然后计算预测分布和真实 token 之间的交叉熵。新手容易这样写错probs torch.softmax(logits, dim-1) loss torch.log(probs[target_index] 1e-8)问题是softmax输出的概率可能因为下溢变成 0log(0)会产生-inf。加一个很小的1e-8只是临时止血可能改变梯度方向。更稳定的做法是直接使用log_softmax或者直接使用框架提供的cross_entropy。import torch.nn.functional as F # logits 形状: (batch, seq_len, vocab_size) # labels 形状: (batch, seq_len) loss F.cross_entropy( logits.view(-1, logits.size(-1)), labels.view(-1), ignore_index-100 )F.cross_entropy内部会先计算 LogSoftmax再取标签位置的值整个过程不需要显式构建 one-hot 向量。微调大模型时应该习惯把 logits 和 labels 同时交给 loss 函数尽量不要手动分离softmax和log。3.4 长上下文推理为什么离不开 Online Softmax大模型推理时注意力矩阵形状是(seq_len, seq_len)。如果输入长度是 8192单个注意力头就要产生 8192 乘 8192 的分数。如果直接把整个分数矩阵写回显存再读出来做 Softmax显存带宽会迅速耗尽。FlashAttention 这类优化内核采用分块策略不保存完整分数矩阵。为了在分块过程中仍能获得正确的 Softmax 结果它必须在每一块数据到来时维护两个临时变量当前最大值m和当前指数和l。以下代码演示了 Online LogSumExp 的核心思路一遍扫描分块同时维护 max 和 sum。import numpy as np def online_logsumexp(logits, block_size3): m -np.inf l 0.0 for i in range(0, len(logits), block_size): block logits[i:i block_size] m_new max(m, np.max(block)) # 旧 sum 需要做 scale因为最大概率分量切换了 l l * np.exp(m - m_new) np.sum(np.exp(block - m_new)) m m_new return m np.log(l) scores np.array([2.0, 1.0, 0.1, 10.0, 0.5, -1.0, 5.0]) log_z online_logsumexp(scores) probs np.exp(scores - log_z) print(logsumexp:, log_z) print(probs sum:, probs.sum())这个版本的输出是一个标量对数归一化常数和log(sum(exp(scores)))等价。它不会一次性对所有元素做exp而是逐块处理这正是长序列推理时常用的思路。在实际 FlashAttention 中计算 Softmax 和 V 的加权求和会更复杂因为每一块计算出的权重还不能直接用必须等待后续块出现更大的最大值后对旧结果重新缩放。学习阶段可以先用简单公式理解原理。到了生产部署本地推理引擎或 GPU 推理框架通常会内置这类内核。任务是通过 attention unit test、日志和 profiling 确认内核在正确运行而不是自己在 Python 层重新实现一遍。4. 生成阶段的 Temperature 也是 Softmax 的工程应用4.1 从概率分布到采样 token 的完整链路大模型生成一个 token 的过程可以拆成四步把已生成的 token 序列送入模型。模型前向推理输出词表大小的 logits。对 logits 做 Softmax得到概率分布。根据采样策略从概率分布中选出一个 token。直接把 Softmax 结果和调用 API 这件事联系起来更容易理解 temperature 的作用。OpenAI-compatible 接口、Ollama 本地接口、各类大模型 SDK普遍都会开放temperature、top_p等采样参数。这些参数改变的不是模型权重而是 logits 到最终 token 之间的“概率整形”过程。4.2 用 temperature 改造概率分布带温度参数的 Softmax 公式是p_i exp(z_i / T) / sum_j exp(z_j / T)温度T的作用是控制分布的尖锐程度。当T 1时公式退化为标准 Softmax。当T 1时logits 被放大原来得分高的 token 概率更高分布更集中输出更稳定。当T 1时logits 被缩小概率分布更平缓低分 token 也有机会被采样输出更多样。看一组简单对比。import numpy as np def softmax_temperature(logits, temperature1.0): if temperature 0: raise ValueError(temperature 必须大于 0) logits np.asarray(logits, dtypenp.float32) / temperature logits logits - np.max(logits) exp_logits np.exp(logits) return exp_logits / exp_logits.sum() logits np.array([2.0, 1.0, 0.1]) for t in [0.1, 0.5, 1.0, 2.0, 10.0]: probs softmax_temperature(logits, t) print(fT{t:6} probs{np.round(probs, 6)})输出趋势T0.1 probs[0.999995 0.000005 0.000000] T0.5 probs[0.841494 0.139312 0.019194] T1.0 probs[0.659001 0.242433 0.098566] T2.0 probs[0.549466 0.304951 0.145583] T10.0 probs[0.374157 0.339336 0.286507]可以看到T0.1时概率几乎变成 one-hotT10.0时三个选项的概率趋于均匀。生产中需要注意temperature 0在很多实现里不是“把标准 Softmax 的分母变成 0”而是直接走贪心解码。不同框架对 temperature 为 0 的处理不完全一致。如果自己实现采样器最好在入口处把temperature 0单独处理不要让它进入除法。4.3 结合 Top-k 与 Top-p 的现代采样方式只用 Temperature 生成容易遇到两个问题。温度太低模型会陷入重复温度太高模型会胡言乱语。因此现代推理引擎一般还会组合 Top-k 和 Top-p。Top-k 的做法是只保留概率最大的 k 个 token把其余 token 的概率清零然后重新归一化。Top-p 的做法是按概率从大到小累加直到累计概率超过 p只保留这些 token再重新归一化。在实现 Top-p 时一定要在过滤后重新做一次归一化。不能只把低概率 token 置零否则保留 token 的概率之和小于 1。下面给出一个同时包含 temperature、Top-k 和 Top-p 的最小采样函数。import numpy as np def sample_with_nucleus(logits, temperature1.0, top_p0.9, top_k50): logits np.asarray(logits, dtypenp.float32) if temperature 0: return int(np.argmax(logits)) probs softmax_temperature(logits, temperature) if top_k is not None and top_k 0: threshold np.sort(probs)[-top_k] probs np.where(probs threshold, probs, 0.0) if top_p 1.0: sorted_probs np.sort(probs)[::-1] cumulative_probs np.cumsum(sorted_probs) cutoff_idx int(np.searchsorted(cumulative_probs, top_p) 1) cutoff_prob sorted_probs[cutoff_idx - 1] probs np.where(probs cutoff_prob, 0.0, probs) if probs.sum() 0: probs np.ones_like(probs) / len(probs) probs probs / probs.sum() return int(np.random.choice(len(probs), pprobs))这个示例代码用于理解流程不直接等于 vLLM、Ollama 等引擎内部的 C 实现。它说明了几个容易被忽略的要点过滤之后必须重新归一化温度很低时直接走 argmax所有分支都要避免概率全为 0 的异常情况。4.4 调用大模型 API 或本地部署时这些参数如何设置如果你正在用本地推理工具部署大模型常见的 OpenAI-compatible 接口或 Ollama 接口通常支持类似字段。以 OpenAI-compatible 风格接口为例请求体可能长这样{ model: your-model-name, prompt: 用三句话解释什么是大模型, temperature: 0.2, top_p: 0.9, max_tokens: 200 }Ollama 的原生/api/generate接口会把采样参数放在options中curl http://localhost:11434/api/generate -d { model: your-model, prompt: 用三句话解释什么是大模型, stream: false, options: { temperature: 0.2, top_p: 0.9 } }不同框架对参数的名称、范围和默认值并不完全相同。落地前不要照抄网上配置要看当前使用版本的 API 文档。特别是temperature范围有的实现接受 0 到 1 或 0 到 2有的实现只接受大于 0 的浮点数。把读取到的参数打印到日志里比盲目传参更容易定位问题。5. 排查 Softmax 相关异常从现象到根因5.1 一张表快速定位实际项目里很多看起来像“模型问题”“显存问题”“上下文长度问题”的故障根因落在 Softmax 附近。下面表格适合作为第一轮筛查。问题现象常见原因检查方式处理建议概率输出出现 NaNlogits 包含 inf/NaN或 float16 溢出查看 logits 的 max、min、isnan 统计稳定化 Softmax必要时用更高精度累积注意力权重每行不为 1在错误维度做了 softmax检查dim是否指向 key 所在序列维度输出权重并沿关键维度求和确认每个位置和为 1训练 loss 异常低但生成效果差Decoder 因果 Mask 没有屏蔽未来 token打印注意力权重检查未来位置是否为 0用-inf填充不能填 0输出全是同一个 tokentemperature 过低或 logits 数值崩溃检查采样参数和 logits 分布适当提高 temperature或开启 Top-p 采样API 报 temperature 相关错误传入 0、负数或超出范围的值查看服务端参数校验逻辑参数入口统一校验0 时走贪心分支自定义 kernel 结果与 PyTorch 不一致Kernel 没有使用减去 max 的稳定策略用随机输入做差分测试参考 Online Softmax 算法处理运行最大值和指数和5.2 排错案例一logits 出现 inf 或 NaN现象是模型训练几步后 loss 变成 NaN或者推理时概率输出直接是 NaN。第一步先看 logits而不要先看概率。概率由 logits 经过指数运算而来只要 logits 异常概率必然异常。import torch def check_logits(logits, namelogits): print(f{name}.shape{tuple(logits.shape)}) print(f{name}.dtype{logits.dtype}) print(fmax{logits.max().item():.6f}) print(fmin{logits.min().item():.6f}) print(fmean{logits.mean().item():.6f}) print(fnan_count{torch.isnan(logits).sum().item()}) print(finf_count{torch.isinf(logits).sum().item()})如果 logits 里出现大量inf或极小数优先怀疑损失函数之前是否为数值不稳定提供了土壤。检查是否直接对很大 logits 调用了exp检查是否没有除sqrt(d_k)检查mask是否出现错误检查混合精度参数是否合理。在纯 Python 实现中解决方式是在做exp前先减去 max。在 PyTorch 中优先使用内置torch.softmax、F.log_softmax、F.cross_entropy不要自己重写不必要的指数计算。5.3 排错案例二因果 Mask 失效导致训练泄漏曾有人在一个小型 GPT 实验中遇到这种现象训练 loss 降得非常快看起来模型“学得不错”但生成文本却完全没有语法顺序。后来检查发现注意力 Mask 在送入 Softmax 前被填成了0而不是-inf。排查方法很直接。写一个 3 乘 3 的随机因果 Mask 单测打印注意力权重矩阵查看上三角位置是否全部为 0。import torch import math scores torch.randn(1, 1, 3, 3) seq_len scores.shape[-1] causal_mask torch.tril(torch.ones(seq_len, seq_len)).bool() scores_masked scores.masked_fill(~causal_mask, float(-inf)) weights torch.softmax(scores_masked, dim-1) print(注意力权重:) print(weights[0, 0].numpy().round(4))正确的上三角位置应该全为 0。例如第一行只能看到第一个位置所以行向量形如[1.0, 0.0, 0.0]第三行可以看到前三个位置行内三个权重都不为 0且相加为 1。如果 Mask 是二维的(seq, seq)要 reshape 到(1, 1, seq, seq)广播到每个 batch 和每个 head。有些框架接受布尔 Mask有些接受浮点 Mask且布尔值 True 代表“允许”别把语义弄反。5.4 排错案例三temperature 参数越界导致服务异常容器化部署大模型时API 层会接收用户输入的 temperature。有的用户为了稳定输出传temperature0有的代码传了负数还有的端把 top_p 传成 1.5。这些参数最终都会进入采样逻辑如果服务端没有做参数校验轻则采样异常重则触发除零和 NaN。建议在 API 入口统一做参数校验。def validate_sampling_params(temperature, top_p): if not isinstance(temperature, (int, float)): raise ValueError(temperature 必须是数值类型) if temperature 0: raise ValueError(temperature 必须大于 0需要贪心输出时请单独指定 greedy 参数) if not (0.0 top_p 1.0): raise ValueError(top_p 必须位于 (0, 1] 区间)这个判断是为了避免真实除法中出现非正温度。如果你使用某个已经封装好的推理框架且框架把temperature0解释为贪心那可以在框架层单独处理但自定义代码不要依赖这种隐含行为。6. 从公式到工程部署大模型时的最佳实践6.1 学习阶段和生产阶段不要使用同一套实现学习阶段最重要的是能看见每一步。用 NumPy 写朴素 Softmax、推导 Mask、观察温度变化这些做法都有价值。生产阶段完全不一样。生产上更关注吞吐、显存占用、数值鲁棒性和可观测性。自己写 Python 版 Softmax 放到服务链路里性能通常远低于 Torch 原生算子或推理引擎的融合内核。推荐做法是服务层使用torch.softmax、F.log_softmax、F.cross_entropy等成熟算子。长序列推理使用引擎内置的 FlashAttention 或高效注意力内核。采样层不要手写完整的贪心逻辑除非你正在验证算法。自定义算子上线前必须和参考实现做差分测试不能只看一两个用例。表格对比能看出两阶段差异。阶段典型需求计算公式风险点学习看中间过程、验证推导朴素 Softmax、减去 max 版本exp 溢出、维度错误微调计算稳定的交叉熵损失LogSoftmax NLLLoss手动 log 导致 -inf训练加速融合算子、混合精度FlashAttention 内嵌 Online SoftmaxKernel 数值结果不一致服务部署低延迟高吞吐引擎内置算子参数越界、采样异常6.2 动手实验清单如果想把 Softmax 真正变成自己的调试工具可以按顺序做下面几个实验。分别实现朴素版本和减去 max 的稳定版本输入[1000, 1001, 999]观察结果差异。用torch.softmax处理形状为(2, 4, 8)的随机张量分别在dim-1、dim-2上测试打印每行、每列的和理解维度含义。实现一个带因果 Mask 的最小注意力打印 4 乘 4 的注意力矩阵确认上三角为 0。写一个温度采样函数用同一组 logits 在温度 0.2 和 1.5 下各采样 200 次统计输出类别频率。使用F.cross_entropy(logits, labels)和手动softmax log各计算一次 loss比较在极端 logits 下的差异。如果是全职做推理优化再用随机输入比较自己写的 NumPy 版 attention 和 PyTorch 版 attention 的误差上限。每个实验都建议有断言。比如“所有概率都在 0 到 1 之间”“每行概率之和接近 1”“attention 矩阵未来位置全为 0”。断言写出来排查时才不会靠肉眼观察。6.3 代码审查与发布前检查在把代码或配置投入生产之前建议检查下面几点。是否在损失函数里直接调用了torch.softmax后再torch.log如果是改成F.cross_entropy或F.log_softmax。是否对 logits 做了exp而没有先减 max如果是改成稳定版本或内置算子。注意力 Mask 使用的是不是-inf而不是 0如果用了 0未来 token 仍然参与注意力。采样函数的 temperature 是否为 0 或负数做了单独分支如果没做会出现除零或参数异常。是否对输入的 logits 做过 dtype 和 NaN 检查float16 下尤其要关注极大值。是否在自定义 Kernel 里保存了完整注意力矩阵如果内存有限优先选择分块式 Online Softmax。是否对自定义算子做了差分测试同一组随机输入下自定义输出和参考输出的误差上限应该被记录。这些检查不只是理论洁癖。在大模型训练和推理中一个小数点后很多位的误差经过多层 attention 和反向传播会被不断放大。公式名字可能不重要但公式在哪儿被调用、调用时有没有避开数值陷阱直接影响项目能否上线。6.4 大模型学习路线里值得先吃透的几个公式如果把“大模型学习路线”展开第一个值得吃透的公式就是 Softmax。它是理解注意力、理解交叉熵、理解采样入口的基础。之后可以按顺序继续学习其他高频公式。LayerNorm 或 RMSNorm稳定层输出很多大模型使用 RMSNorm 减少计算量。RoPE 旋转位置编码给 token 位置信息让注意力分数感知相对位置。SwiGLU 或 GELU 类激活函数改变 FFN 层的信息变换方式。Adam/AdamW 更新公式训练优化器背后的动量与权重衰减机制。理解这些公式时不要停留在“调用过”层面至少要能用手写代码复现输入输出能解释一个数量级变化对结果的影响能说清某处eps加在除法中的原因。这样在看大模型推理引擎日志、阅读模型源码、复现训练实验时才不会觉得每个算子都是一个黑盒。Softmax 是否被人记住名字并不重要。关键是这样一类数学基础组件会被每一代 Transformer、每一位研究者和每一个推理引擎反复使用。它能帮你快速定位“概率为什么不对”也能帮你理解“温度为什么这样影响输出”。建议现在打开 Python 环境把上面几个最小实验跑一遍。当你亲手写出能处理[1000, 1001, 999]的稳定版本再回头看大模型 API 返回的结果会明显感觉到公式没有停在纸张上而是真实运行在每一次 token 生成里。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价