前两天一个朋友半夜发消息问我“我的小模型训练到一半loss 曲线突然开始出现尖峰看到一个一个的注意力图全是‘那种只盯着自己的位置’的尖刺这正常吗”我回了一句“你大概率是 QK 点积 logits 失控了给 Q 和 K 加个保险丝吧。”他一脸懵“就 300M 参数的小模型也要搞 QK-norm 这种大模型技术”这个问题其实问得特别好也是很多从单机训练、从几百万到十亿参数规模往上走的人都会遇到的分岔路口小模型到底要不要装 QK 保险丝为什么不装会出问题装了之后用 QK-norm、softcap、还是退火 QK-norm到底怎么选这篇文章就专门聊这个顺便把我在实际训练上百个中小规模模型过程中踩过的坑、用过的数据、看过的曲线和改过的代码一起拿出来说。1. QK 点积失控时训练里到底发生了什么在很多人的直觉里“注意力崩溃”是大模型专属病小模型参数少、层数浅应该挺健康。这个直觉错在最关键的一个环节上注意力健康的判断维度不是模型参数量而是每一层激活值的数值分布。1.1 一个容易被忽略的方差放大现象先回到最基本的 self-attention 公式。给定当前层的输入X每个头会算出一组 query 和 keyQ X W_Q K X W_K logits Q K^T / sqrt(d_k)如果只看X W_Q这一项它本质上是输入向量与权重矩阵的线性组合。输入向量经过足够多层的残差相加、LayerNorm、FFN 变换之后内部元素的方差早就不会保持初始的“单位方差”状态了。随着网络加深、学习率增大、并且W_Q和W_K处于初始化阶段或者震荡阶段Q和K各自的方差可能偏离 1 很远。接着Q K^T相当于把d_k个不相关随机项相加如果每一项的方差都是σ²点积结果就是方差约为d_k * σ²的量级。d_k在小模型里通常是 64 或 128。也就是说哪怕每个元素都很正常这个相加过程也会把 logits 的标准差放大约 8 到 11 倍。再配合训练初期权重不稳定logits 的值冲到 30、40、甚至 100 我并不奇怪。一个 softmax 喂进去 100 量级的输入输出基本就是 one-hot梯度传回去会变成一串接近 0 的小数注意力头就“死”了不再更新。这个问题不是大模型独有的小模型只是层数少但没有逃离“点积求和会放大方差”这个数学事实。100M 参数和 3B 参数在这个点上没有本质区别。1.2 失控的三个信号我自己的经验是训练时不会只从 loss 曲线看问题因为 loss 曲线往往要等到崩溃之后才给你颜色。更早的观察点是下面这几个注意力熵持续走低。过热 logits 会让每个 token 的注意力分布趋近 one-hot也就是说它只关心极小一部分 token。算一下每个注意力头在平均序列上的熵如果这个熵随着训练进行不是缓慢下降而是直接掉到正常值的 50% 以下就该停下来看看 logits 分布了。logits 绝对值上限异常放大。在训练脚本里打印每个注意力层在 softmax 之前的logits.abs().max()。小模型正常训练时这个值通常分布在 10 到 25 之间如果你看到超过 40并且不是训练初期那种偶发尖顶而是一路变大那就是 QK 动态出了问题。loss 曲线出现周期性的“心电图”。这本质上是因为 logits 偶尔突破某个阈值导致一小部分 token 的注意力完全集中梯度出现尖峰让局部参数在几步之内偏离轨迹然后又要花几百步拉回来。1.3 小模型真的不容易失控吗很多做小模型的人有个错觉小模型我随便训成本低跑几个 epoch 非常快崩了再调就行。这句话对一半。小模型虽然训练便宜但正因为便宜大家往往会加大学习率、减少 warmup、缩短验证间隔这些操作恰恰把训练推向更激进的区域让 QK logits 的失控变得更常见。另外小模型的容量有限一旦少数注意力头死掉其他头没有足够的冗余去补上它们的职责模型的有效表达力会真实地下降。我见过一个 300M 参数模型死掉 3 个注意力头后下游分类指标直接掉了 4 个点。所以“小模型”从来不是“不需要关心注意力稳定性”的借口反而因为容量紧张更承受不起内部组件失效。2. QK-norm给 Q 和 K 各装一根保险丝如果你还没接触过 QK-norm可以先把它理解为在 Q 和 K 进入点积之前分别做一次归一化把它们的尺度拉到可控范围然后乘上一个可学习的缩放因子。2.1 QK-norm 的数学直觉核心做法非常简单对每一个注意力头q rms_norm(q) * qk_scale k rms_norm(k) * qk_scale其中rms_norm会按最后一个维度计算均方根rms(x) sqrt(mean(x^2)) x_norm x / rms(x)为什么用 RMSNorm 而不是带均值偏移的 LayerNorm因为在 QK 这条路径上我们最关心的是方差、是幅度均值部分说实话不太影响点积的相对分布。RMSNorm 不需要计算均值实现更轻、算子更快而且在小模型上效果完全够用。关键点在于除以均方根之后每个头里的 q 向量和 k 向量都被拉到了“单位尺度”附近再通过可学习的qk_scale来恢复模型真正想要的 logits 幅度。模型不再被动接受权重初始化带来的尺度灾难而是把幅度控制权交到了可学习参数手里。2.2 一段可以直接用的实现这里给一个我平时经常直接复制进项目的 PyTorch 实现import torch import torch.nn as nn import torch.nn.functional as F class RMSNorm(nn.Module): def __init__(self, dim, eps1e-6): super().__init__() self.eps eps self.scale nn.Parameter(torch.ones(dim)) def forward(self, x): rstd x.pow(2).mean(-1, keepdimTrue).add(self.eps).rsqrt() return x * rstd * self.scale class QKNormHead(nn.Module): 对单个注意力头的 q 或 k 做归一化。 head_dim: 每个头的维度 scale_init: 初始缩放因子我习惯用 1.0 def __init__(self, head_dim, scale_init1.0, eps1e-6): super().__init__() self.eps eps self.scale nn.Parameter(torch.ones(head_dim) * scale_init) def forward(self, x): rstd x.pow(2).mean(-1, keepdimTrue).add(self.eps).rsqrt() return x * rstd * self.scale在注意力模块里这样接入class Attention(nn.Module): def __init__(self, hidden_dim, num_heads): super().__init__() self.num_heads num_heads self.head_dim hidden_dim // num_heads self.wq nn.Linear(hidden_dim, hidden_dim) self.wk nn.Linear(hidden_dim, hidden_dim) self.wv nn.Linear(hidden_dim, hidden_dim) self.q_norm QKNormHead(self.head_dim) self.k_norm QKNormHead(self.head_dim) def forward(self, x): batch, seq_len, _ x.shape q self.wq(x).view(batch, seq_len, self.num_heads, self.head_dim) k self.wk(x).view(batch, seq_len, self.num_heads, self.head_dim) v self.wv(x).view(batch, seq_len, self.num_heads, self.head_dim) q self.q_norm(q) k self.k_norm(k) logits torch.einsum(b h t d, b h s d - b h t s, q, k) logits logits / math.sqrt(self.head_dim) probs torch.softmax(logits, dim-1) out torch.einsum(b h t s, b h s d - b h t d, probs, v) return out注意我仍然保留了1 / sqrt(head_dim)的缩放这是很多实现里容易被忽略的点。QK-norm 把每个向量的分布拉回了单位尺度但如果完全不除以sqrt(head_dim)logits 的量级会变成head_dim乘以一个平均接近 1 的数对 64 维 head 来说就是 64 量级仍然偏大。保留标准缩放让可学习的qk_scale在这个基础上微调训练会更稳。2.3 与除以 sqrt(d) 的区别有人会问“传统 attention 不是已经有1 / sqrt(d_k)缩放了吗为什么还不够”这么理解1 / sqrt(d_k)是一个静态缩放它在模型初始化阶段尽可能把 logits 方差拉回 1。但它不知道训练进行到某一步时当前 Q 和 K 的实际分布已经膨胀到了什么程度。它就像一个固定大小的保险丝电流一直在变大但保险丝不会跟着调整。QK-norm 是动态的。每个 token、每个头都会在点积前实时计算输入向量的实际幅度然后主动把它压回单位尺度。这不只是“保险丝”还是一个带反馈调节的电路保护器。它能兜住那些1 / sqrt(d_k)兜不住的意外情况比如学习率突然调过头、数据批次里出现超长序列、某个 head 被梯度推到了异常区。在实际训练中我发现加了 QK-norm 之后logits 的 max 稳定在 15 到 25 之间几乎不会超过 30。没加 QK-norm 时同样的学习率和数据顺序能冲到 60 以上。这就是动态归一化和静态缩放之间最直观的差异。3. softcap另一根保险丝但位置和形态不同如果说 QK-norm 是“事前约束”那么 softcap 就是“事后限制”。它不动 Q、K 的分布而是在点积计算完、softmax 之前把已经算好的 logits 做一个软裁剪。3.1 softcap 到底做了什么softcap 的典型公式只有一行logits c * tanh(logits / c)对应的代码def softcap_logits(logits, c30.0): return c * torch.tanh(logits / c)这个操作的几何意义很直接tanh的输出被限制在[-1, 1]之间再乘回c最终 logits 必然落在[-c, c]区间内。logits 在 0 附近时tanh(x/c) ≈ x/c也就是说正常范围内的 logits 几乎不受影响基本是线性通过只有超过c的极端值才被压缩。它不是硬裁剪。如果是像“如果 logit c就赋值成 c”这样的硬裁剪在边界处梯度会直接消失这会让训练变得僵硬。softcap 用 tanh 做了一个平滑过渡越大的值被压制得越狠同时梯度不会瞬间归零会逐渐衰减给了模型缓和的适应性。3.2 软保险丝和硬保险丝的差别这里有一个经常被搞混的点QK-norm 和 softcap 解决的问题有重叠但并不完全相同。维度QK-normsoftcap作用位置Q/K 编码后、点积前点积后、softmax 前约束对象Q 和 K 向量的尺度最终的 logits 数值对正常 logits 的影响会被缩放幅度由可学习参数决定基本线性通过几乎不影响极端值处理能力强能把 logits 源头压到稳定量级强但只是把极端值截回有限范围推理额外算子两次 RMSNorm一个 tanh可算子融合实现成本需要新增模块、可学习参数一行公式无新增参数注意一个区别QK-norm 会改变正常 logits 的整体尺度因为归一化之后所有 logit 都按统一步伐缩放而 softcap 对处于安全区间的 logits 几乎没有影响它只负责“把最冒尖的那几个值拉回来”。所以如果你只想针对“偶尔出现的尖峰 logits”做治理softcap 的侵入性更小、效果更局部。如果你希望从根本上稳定整个 QK 分布的方差QK-norm 更彻底。3.3 小模型的取舍思路对小模型我实践下来的感觉是softcap 的商业性价比极高因为它实现成本几乎为零却能在很多场景下把注意力训练从崩溃边缘拉回来。给一个具体例子。我训过一个 500M 参数的模型一开始不加任何保护logits 的 max 经常冲到 45 左右导致注意力熵崩到非常低。我当时不想动 QK-norm 的原因是推理端在某个低算力部署环境能少一个算子就少一个算子。于是先加了一个c30的 softcap训练再跑 20 万步logits max 被锁在 30 以内loss 曲线的尖峰明显减少最终验证困惑度略微优于失控版本。但也要提醒softcap 的c参数有讲究。c设太大比如 100等于没装保险丝c设太小比如 10正常 logits 也会被压缩注意力分布会变得平滑过头模型表达能力反而下降。我自己的经验是先观察不加保护时的 logitsmax分布取一个“略低于失控上限、又高于正常波动上限”的值。常见选择是 30也可以在 20 到 50 之间扫一轮。小模型训练便宜这也是为什么我一直觉得小模型试 softcap 比直接试 QK-norm 更划算。4. 退火 QK-norm把保险丝用完就拆QK-norm 和 softcap 各有优劣但如果你把视角放到“部署阶段”QK-norm 会多出一个麻烦模型推理时也必须带上这两个 RMSNorm 算子哪怕它们的参数已经训练得比较好但在某些推理框架里每个注意力层多两个算子就意味着额外的延迟和显存带宽占用。于是有了第三种思路训练时用 QK-norm 稳定衰减推理时直接把 QK-norm 摘掉让模型表现得好像它从来没存在过。这就是退火 QK-norm。4.1 为什么会有退火 QK-norm退火 QK-norm 的出发点其实是部署友好。我在某个线上服务场景里见过具体的教训离线训练时大家觉得 QK-norm 好用就直接用上了结果到了上线前做延迟优化发现每个 attention 层多出来的两次 RMSNorm 让推理速度掉了 5%。5% 听起来不多但对一个高并发服务这段时间换算成成本就是真金白银。怎么既保住训练稳定性又不让推理背这个负担办法是训练前期完全使用 QK-norm在训练快结束时逐步把 QK-norm 的输出和原始的 Q、K 做线性插值让插值权重从 0 慢慢变到 1最终在模型内部让 QK-norm 的影响趋近于零。训练结束后推理代码直接走不带 QK-norm 的分支模型行为几乎不会发生变化。4.2 退火插值实现细节我用的实现思路是这样先定义一个退火系数 alpha训练早期 alpha0完全使用 QK-norm训练后期 alpha 线性增长到 1完全退回到原始 Q、K。def interpolated_qk(q, k, alpha, q_norm, k_norm): if alpha 0: return q_norm(q), k_norm(k) if alpha 1: return q, k qn q_norm(q) kn k_norm(k) q alpha * q (1.0 - alpha) * qn k alpha * k (1.0 - alpha) * kn return q, kalpha 的调度放在训练循环里total_steps 150_000 anneal_begin_ratio 0.8 anneal_begin int(total_steps * anneal_begin_ratio) def get_alpha(step): if step anneal_begin: return 0.0 return min((step - anneal_begin) / (total_steps - anneal_begin), 1.0)这里anneal_begin_ratio0.8表示最后 20% 的训练步数完成退火。退火过程不是一上来就线性我尝试过几种曲线包括余弦退火和线性退火实际差异不大重要的是别把退火周期压得太短。我甚至建议如果条件允许可以在退火结束后再加一个几十步到几百步的“尾巴训练”把学习率降到很低让模型在完全没有 QK-norm 的状态下稍微稳定一下参数。这样推理时剥掉 QK-norm 分支验证集困惑度几乎不会有变化。4.3 陷阱与注意点退火 QK-norm 不是无脑灵丹。我有几个实际操作中的注意点分享退火周期不能太短。如果最后 5% 的步数里把 alpha 从 0 拉到 1模型来不及适应loss 会出现明显回升。我一般建议退火区间至少占训练总步数的 10% 到 20%。对 150K 步的训练最后 30K 步退火是我比较稳的经验。只对 Q 和 K 插值不要动 V。V 的路径不参与点积 logits 的尺度问题硬把 V 也插值一遍反而会引入额外的分布偏移。这个错误我犯过看到损失在退火阶段异常爬升排查了很久才意识到是自己把插值接错了分支。attention 层的数量会影响退火时间。层数越多每层同时从 QK-norm 状态切到无 norm 状态累积起来的分布移动越明显。如果你的模型有 24 层以上退火区间建议靠更长的anneal_begin_ratio比如0.85。5. 小模型的最终决策装保险丝、装哪种、怎么验证讲了这么多原理回到最核心的问题我的小模型到底要不要装这不是一道判断题更像一道“先做两个小实验再决定”的工程题。小模型的优势就在于可以便宜地做对照组别浪费这个优势。5.1 一套可复现的对照实验方案我建议在小模型上跑一个四组对照实验每组都训到相同的 token 量或步数然后在相同的验证集上对比配置编号保护方式实现成本推理额外负担A不加任何保护最低无B仅 QK-norm低保持 QK-norm推理多两个 RMSNormC仅 softcapc30极低一个 tanh可融合D退火 QK-norm中推理时完全摘除每种配置都固定同一套种子、同一份 token 序列、同一个学习率计划。注意如果你想比较“保险丝的价值”学习率可以稍微激进一点点比如比平时默认值高 20% 左右因为故意引入压力才能看出保护机制的作用。如果一切正常到毫无压力那结论就失去参考意义了。然后在训练日志里同时记录这几项每个注意力层 logits 的abs().max()尤其关注第 1 层、中间层、最后几层平均注意力熵按所有头平均loss 曲线的低频波动情况看有没有周期性尖峰最终验证集困惑度或下游任务指标。5.2 怎么读实验结果看完四组结果通常会出现下面几种情况如果 A 组在激进学习率下没有崩溃logits 都稳定在 25 以内那不用装保险丝A 就是最优解。极小模型在适当 learning rate 和 warmup 下确实可以不需要任何额外保护不要为了“别人都在用”就硬加。如果 A 组崩了B、C、D 都稳住了优先考虑 C。因为 softcap 改动最小、参数最少、推理负担最轻训练收益却已经很明显。小模型的部署环境往往对模型大小和延迟敏感能少一个算子就少一个。如果 C 组虽然稳住了但 logits 的 max 始终贴着 c 的上限走说明有大量注意力头长期处于被压缩状态模型的可塑性受到了压制。这时换 B 或 D 通常更好因为 QK-norm 是动态归一化不会把大量正常范围内的 logits 压到同一个边界。如果 D 组和 B 组的验证指标差不多但你的部署端对延迟敏感选 D。这本质上是“拿训练时多一点的鲁棒性换推理时少两个算子”的权衡。5.3 我目前的默认方案与小技巧从很多次实验结果看我给一个比较稳的默认组合300M 到 1B 参数、8 到 16 个注意力头、训练 token 量在 1B 到 10B 之间的小模型我一般先跑一次不带保护的短线试验看 logits 分布是否正常。如果正常就直接开训不加任何保险丝只要出现过一次尖峰或者注意力熵异常我接下来的默认配置是QK-norm softcap(c50)并在训练最后 15% 步数做退火 QK-norm 的 alpha 插值。这里多提一句QK-norm 和 softcap 并不冲突。很多听到“QK-norm”和“softcap”会觉得是二选一实际可以同时用。QK-norm 负责源头上的幅度控制softcap 负责兜住那些偶发的、分布外的极端 logits。两者配合起来在线长尾序列、异常长文本、以及学习率调整产生短暂波动时注意力层会更冷静。最后说一个我在实践中会反复提醒自己的小细节给注意力加保险丝时要么所有层统一加要么干脆别加“几层加、几层不加”这种方案。你可能会想只在前几层加因为直觉觉得浅层更容易出问题。但我实际观察下来的情况是不同模型的数据分布不一样有的模型反而是深层注意力头先崩。做非对称配置当然可行但会增加调试成本而且你很难判断当前层崩是因为自身不稳定还是上一层 logits 异常传导过来的。统一配置先保证整体稳定再根据具体层的 logits 监控去寻找局部优化机会是更省时间的路线。小模型的优势就在于便宜、迭代快可以反复做这种对照实验。保险丝装不装与其在网上找“标准答案”不如花半天时间把对照组跑完用自己模型上的真实曲线说话。