前两天一个朋友半夜发消息问我:“我的小模型训练到一半,loss 曲线突然开始出现尖峰,看到一个一个的注意力图全是‘那种只盯着自己的位置’的尖刺,这正常吗?”我回了一句:“你大概率是 QK 点积 logits 失控了,给 Q 和 K 加个保险丝吧。”他一脸懵:“就 300M 参数的小模型,也要搞 QK-norm 这种大模型技术?”
这个问题其实问得特别好,也是很多从单机训练、从几百万到十亿参数规模往上走的人都会遇到的分岔路口:小模型到底要不要装 QK 保险丝?为什么不装会出问题?装了之后用 QK-norm、softcap、还是退火 QK-norm,到底怎么选?这篇文章就专门聊这个,顺便把我在实际训练上百个中小规模模型过程中踩过的坑、用过的数据、看过的曲线和改过的代码一起拿出来说。
1. QK 点积失控时,训练里到底发生了什么
在很多人的直觉里,“注意力崩溃”是大模型专属病,小模型参数少、层数浅,应该挺健康。这个直觉错在最关键的一个环节上:注意力健康的判断维度不是模型参数量,而是每一层激活值的数值分布。
1.1 一个容易被忽略的方差放大现象
先回到最基本的 self-attention 公式。给定当前层的输入X,每个头会算出一组 query 和 key:
Q = 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, eps=1e-6): super().__init__() self.eps = eps self.scale = nn.Parameter(torch.ones(dim)) def forward(self, x): rstd = x.pow(2).mean(-1, keepdim=True).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_init=1.0, eps=1e-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, keepdim=True).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, c=30.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-norm | softcap |
|---|---|---|
| 作用位置 | 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 的原因是推理端在某个低算力部署环境,能少一个算子就少一个算子。于是先加了一个c=30的 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,训练早期 alpha=0,完全使用 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_ratio=0.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,推理多两个 RMSNorm |
| C | 仅 softcap(c=30) | 极低 | 一个 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(c=50),并在训练最后 15% 步数做退火 QK-norm 的 alpha 插值。
这里多提一句,QK-norm 和 softcap 并不冲突。很多听到“QK-norm”和“softcap”会觉得是二选一,实际可以同时用。QK-norm 负责源头上的幅度控制,softcap 负责兜住那些偶发的、分布外的极端 logits。两者配合起来,在线长尾序列、异常长文本、以及学习率调整产生短暂波动时,注意力层会更冷静。
最后说一个我在实践中会反复提醒自己的小细节:给注意力加保险丝时,要么所有层统一加,要么干脆别加“几层加、几层不加”这种方案。你可能会想只在前几层加,因为直觉觉得浅层更容易出问题。但我实际观察下来的情况是,不同模型的数据分布不一样,有的模型反而是深层注意力头先崩。做非对称配置当然可行,但会增加调试成本,而且你很难判断当前层崩是因为自身不稳定还是上一层 logits 异常传导过来的。统一配置,先保证整体稳定,再根据具体层的 logits 监控去寻找局部优化机会,是更省时间的路线。
小模型的优势就在于便宜、迭代快,可以反复做这种对照实验。保险丝装不装,与其在网上找“标准答案”,不如花半天时间把对照组跑完,用自己模型上的真实曲线说话。