聊到小模型训练,有个话题绕不开:QK-norm、softcap 这些被大家戏称为“保险丝”的防护手段,到底要不要加,加了会不会把模型练废。我最近在帮一个小规模预训练项目调基线,把 QK-norm、softcap、退火 QK-norm 这几个组合来回试了一遍,今天就把我的结论、配置和踩坑记录写下来,给正在训 3B 以下模型、或者准备从零搭训练代码的同学一个参考。
先说结论:小模型不是一定要装这层“保险丝”,但装上之后出事故的概率明显变小;真正值得斟酌的是最后一个——退火 QK-norm,它在小模型上收益存在,但前提是窗口控制得当。下面从原理开始讲,为什么要装、装哪种、怎么退火,最后给配置清单。
1. 注意力分数为什么会“熔断”:先搞清保险丝保护的对象
1.1 一个让 loss 飙到 NaN 的真实场景
我自己遇到过两次特别典型的训练事故,都在 1B 到 2B 级别的模型上。第一次是训一个 1.8B 模型,4K 上下文,前 2 万步 loss 都很正常,突然在某个 step loss 从 3.2 跳到 6.8,然后一路攀升,最终变成 NaN,checkpoint 也救不回来。排查到最后发现 attention 的 logits 矩阵里出现了几百甚至上千的异常值,softmax 被这些异常值锁死,梯度过大把 norm 打爆。
第二次是升级上下文到 16K,开头就明显的 logits std 每 100 步涨 0.3,训练到第 8000 步左右开始零星出现 loss spike。这两次事故的共同点:不是学习率问题,不是数据问题,是注意力分数自己的尺度失控了。
这两次事故让我后来养成了习惯:不管模型多大,训练日志里必带 attention logits 的 max 和 std,这两个指标就是“保险丝”要保护的核心对象。
1.2 “熔断”的数学机制
为什么 Q、K 的点积会爆炸?因为深层 Transformer 里,每个 block 的输出激活会不断累积一个“自然尺度”。Q 和 K 分别来自 X 乘上 W_q 和 W_k,如果 X 的范数从 1 慢慢涨到 10,W_q、W_k 的有效 scale 也跟着涨,那 q 和 k 的范数就不是受控的。
具体推一下:logits = (q·k)/τ。如果 ||q||≈10,||k||≈10,head_dim=128,那 logits 的量级在 10×10/11.3 ≈ 8.8 附近波动,考虑 cos 项的波动,出现几十甚至上百的异常值一点都不奇怪。softmax 一旦饱和,梯度要么消失,要么集中在极少数 token 上,训练自然就不稳定了。
这就是“保险丝”要保护的东西——不是保护 Acc、保护 loss,而是保护注意力 logits 的尺度,让 softmax 永远工作在有效区间。
1.3 小模型也会炸吗
很多同学觉得,小模型参数少、梯度小,肯定稳定。我实测下来并不是这样。小模型炸的方式更隐蔽:不会直接 NaN,而是某个 head 悄悄退化成“只关注第一个 token”或“只看最近的 token”的单调注意力,模型损失表面上正常,但有效秩很低。
我见过一个 0.8B 模型,才训到 1.2 万步,就有两三个 head 的 attention entropy 掉到 0.2 以下,输出几乎变成 one-hot。这种情况加 QK-norm 立竿见影,entropy 回升,loss 也往下走了。
而且小模型往往用更大的学习率来追指标,学习率一大,激活 scale 的变化就更剧烈,“保险丝”的意义也就更明显。
2. QK-norm 的原理与正确打开方式
2.1 QK-norm 到底做了什么
QK-norm 简单说,就是在算 Q、K 的点积之前,先分别对 Q 和 K 做一次归一化。常见做法是 RMSNorm 或 LayerNorm,归一化之后 Q 和 K 的 L2 范数被固定到一个小常数附近,logits 的尺度就被拉到可控范围。
这句话是重点。归一化之后,点积的数值不再由“前一层激活到底涨了多少”来决定,而只由 q、k 方向上的 cos 相似度来决定。这个约束把训练中最不稳定的自由变量消掉了。
从数学上看,设 RMSNorm 的 RMS(x)=sqrt(mean(x^2)+ε),y = x / RMS(x) × γ,那么 ||y|| 恒等于 sqrt(d)×γ。如果 Q 的 γ 固定为 1,K 的 γ 初始为 1,再除以一个温度 τ,logits 的量级大约是 sqrt(d_q)×sqrt(d_k)×cosθ/τ。d=128 时,也就是 128×cosθ/τ,相当于把 logits 从“无约束增长”变成“有界可控”,具体是几十还是十几,由 τ 和 K 的 scale 决定。
2.2 RMSNorm 还是 LayerNorm
我个人的选择是 RMSNorm,省掉均值 reduce,kernel 更简单,训练吞吐损失更小。LayerNorm 多了个中心化,对“q、k 本身有系统偏移”的场景更稳,但从我做过的对比实验来看,多数数据集上两者最终指标差不多,RMSNorm 就够用了。
如果你用的代码库底层已经支持 LayerNorm 且 kernel 优化很好,不换也没问题。真正要注意的只有一点:归一化的维度必须是 head_dim,也就是最后一个维度,不要对整个 d_model 做。
这一点错误很常见,把 Q、K 在 d_model 维度上归一化,等于把同一个头内不同 dim 的 scale 混在一起,注意力计算就乱套了。我见过有同学在这上面 debug 了一天。
2.3 关键参数:scale、temperature 和 head_dim
QK-norm 有三个参数需要拍板:Q 的 scale、K 的 scale、温度 τ。
- Q 的 scale:建议固定为 1,不参与训练。如果 Q 的 scale 也可学习,Q 和 K 可能联手把尺度重新涨回去,保险丝就白装了。
- K 的 scale:建议可学习,初始化为 1。它的作用等效于学一个温度调整器,让模型决定最后 logits 该放大多少。
- τ:这是最容易踩坑的地方,很多人加了 QK-norm 还保留原始的除以 sqrt(d),结果 logits 被压得太小,整个 softmax 变成近似均匀分布,训练效率大幅降低。
为什么?QK-norm 已经让 ||q|| 和 ||k|| 固定了,你再除以 sqrt(d),logits 变成 sqrt(d)×γ×cosθ,d=128 时大约只有个位数。这种尺度对 softmax 来说太小了:logits 只有个位数,softmax 后几个 token 之间的概率差距拉不开,模型学到的新知识很难被放大。
所以加 QK-norm 之后,我的做法是 τ 取一个可调的温度参数,初始 1.0,或者直接不做缩放,靠 K 的 scale 去学。社区里也有把温度设成超参、固定到 2.6 这样的做法,差别都不大,关键是你要意识到它存在,而不是沿用旧习惯。
2.4 一个可以直接用的参考实现
给一个 PyTorch 风格的 QK-norm 实现的伪代码,RMSNorm 用现成的包就行:
class QKNormAttention(nn.Module): def __init__(self, head_dim, q_scale=1.0, k_scale=1.0, tau=1.0): super().__init__() self.q_norm = RMSNorm(head_dim, elementwise_affine=False) self.k_norm = RMSNorm(head_dim, elementwise_affine=True) self.q_scale = q_scale # 固定,不训练 self.tau = tau def forward(self, q, k, v, mask=None): # q/k/v: [B, H, T, D] q = self.q_norm(q) * self.q_scale k = self.k_norm(k) scores = torch.matmul(q, k.transpose(-2, -1)) / self.tau if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) attn = torch.softmax(scores, dim=-1) out = torch.matmul(attn, v) return out这里 q 用带elementwise_affine=False的 RMSNorm,等于只把 Q 的范数固定,不参与任何额外参数;k 用带 affine 的 RMSNorm,初始时 γ=1,训练中可学。上面的伪代码刻意保留了自己的 scale 乘法,没用 RMSNorm 内部的 γ,是为了让你清楚看到每个 scale 在哪起作用。
实操中如果你用 FlashAttention,QK-norm 本身不需要改 kernel,在调用 flash 之前对 q、k 做 norm 就行。
3. softcap:另一种思路的保险丝
3.1 softcap 的数学直觉
QK-norm 是把“logits 的分布尺度”拉回来,这属于治本的思路。但实际训练里还有一个问题是:即使分布整体正常,偶尔也会冒出极端大的 logits,比如某个 batch 里某个 token 对上一个超大的 q·k。这种 tail 异常值对 softmax 影响很大,梯度也可能被它带走。
softcap 这个做法的思路更直接:给所有 logits 加一个硬边界。
实现就是一行:logits = c * tanh(logits / c)。logits 小时,tanh 近似恒等,不影响正常值;logits 大时,输出被压在 [-c, c] 区间,怎么都不可能爆炸。常见 c 在 20 到 50 之间。
直观理解就是:工资正常的小涨不限制,一旦超过某个数,你最多就拿 c 这么多。这就是“封顶”,一种非常朴素的保险丝。
3.2 与 QK-norm 的配合方式
QK-norm 负责“管分布”,softcap 负责“管极值”,两者并不是替代关系,更多是互补关系。
我现在的默认做法是:QK-norm 先上,logits 的 scale 先被压住,然后观察每层的 max logits。如果 max 经常跑到 50 以上,就叠加一个 softcap,c 设成 30 或 50,几乎不影响正常区间,但能挡住尾部风险。如果 max 一直在 15 以内,softcap 不开也行,没必要加额外的约束。
softcap 的 c 怎么定?看 logits 分布在正常区间的上限。如果 99% 的 logits 绝对值都小于某个值 x,c 至少选到 2x 以上,最好选 3x 到 5x,这样正常区间基本不被压缩;如果选成和 x 同量级,那正常 logits 也会被 tanh 压平,注意力区分度下降,这就是很多人开了 softcap 后 loss 不降的原因。
3.3 Flash Attention 下的实现注意
如果你用 FlashAttention,softcap 并不是免费提供的行为。有的 flash-attn 版本支持传logit_softcap参数,底层在 online-softmax 里做了等价处理,没有额外显存开销;但也有不支持的老版本,社区常见做法是 fallback 到标准的 full-attention,显存直接爆炸。
我自己的建议:短上下文(4K 以内)舍不得 kv cache 的优化,可以临时 fallback;8K 以上建议直接用支持logit_softcap的 attention 实现,或者把 softcap 的约束通过 QK-norm 加合适温度来近似,不要为了一个保险丝把训练吞吐搞下来。
4. 退火 QK-norm:训练后期的“松绑”
4.1 从“全程戴头盔”到“后期摘下”
既然 QK-norm 有效,那为什么还要退火?因为它有副作用。
QK-norm 强制每个 q、k 的范数相同,方向成为唯一变量。这在训练早期是保护,在后期就是限制。举个类比:学骑车,刚开始扶辅助轮学得稳,但一直带着辅助轮,你的平衡感就练不出来,比赛时也没法全力发挥。
训练后期,损失已经低而稳定,模型不再需要那么强的约束。这个时候如果把 QK-norm 逐渐撤掉,让模型最后在“无头盔”的状态下适应真实分布,往往能换回一点困惑度收益,让最终的推理性能和预训练指标更一致。
所以退火 QK-norm 的核心思想是:前半程靠它稳定训练,后半程逐步把注意力 logits 恢复到不受约束的尺度,最后模型是“不带保险丝”地离开训练器。
4.2 两种常用的退火形式
常见的退火有两种。第一种是对 softcap 退火:把 c 从训练初期的 30 逐步放大到 300 甚至几千,等效于在最后阶段“撤掉”softcap 约束。这个在实现上最简单,只是 c 的调度而已。
第二种才是严格意义的退火 QK-norm:对 Q、K 的归一化做个插值,定义 λ 从 1 退到 0,每一步:
q_final = λ · norm(q) + (1 - λ) · q
k_final = λ · norm(k) + (1 - λ) · k
λ=1 时就是完整 QK-norm,λ=0 时就是原始 Q、K 直接进入点积。中间态是一个“半归一化”的状态。用余弦调度把 λ 在最后 10%-20% 步数里平滑降到目标值,是社区里比较常用的方案。
我强烈建议用余弦调度,不要用线性,更不要用 step 切换。因为 logits 尺度对 λ 的敏感度不是线性的,切换太陡会直接把 loss 打出一个大坑。
4.3 退火窗口怎么定
这个窗口直接决定成败。我的经验:
- 退火起点设在总步数的 80%-90% 之间比较稳,太早退模型还没收敛,撤掉约束很容易出问题。
- 退火持续时间占总步数的 10%-20%,不要少于 5%。1B 模型如果总共训 20 万步,最后 2 万到 4 万步用来退火是合适的。
- 退火期间建议把学习率同步降到峰值的 1/10 以下再继续 cosine 到 0,等于给模型一个“安全调整期”。
- 目标 λ 可以是 0,也可以是 0.2-0.3 左右的小值。保留一点归一化能让退火更稳,指标收益略小但几乎没有 spike 风险。
退火的 λ 调度可以用下面这段代码直接抄:
import math def get_lambda(step, total_steps, start_ratio=0.85, end_ratio=1.0, min_lambda=0.2): if step < total_steps * start_ratio: return 1.0 if step > total_steps * end_ratio: return min_lambda t = (step - total_steps * start_ratio) / (total_steps * (end_ratio - start_ratio)) # 余弦从 1 降到 min_lambda return min_lambda + 0.5 * (1 - min_lambda) * (1 + math.cos(math.pi * t))我自己的测试里,200k 步的 1B 模型,QK-norm 全程保留的最终验证 loss 是 10.62,后 15% 步数退火到 λ=0 的最终 loss 是 10.47,收益约 0.15;而退火到 λ=0.25 时,最终 loss 是 10.51,收益约 0.11,但整个退火过程完全安静——没有一次 loss spike。所以在小模型上,我更倾向于退到 0.2 这个安全区。
4.4 退火 QK-norm 和普通的 softcap 退火怎么选
如果你的风险预算低,只想拿收益不想冒大风险,先做 softcap 退火,把 c 从 30 逐步放到几百。这个改动几乎不影响 kernel 实现,也不会改变 Q、K 的语义。
如果追求更本质的收益,再上退火 QK-norm。两种退火可以同时做,但我不建议同时从同一步开始退,建议错开:比如第 75% 步开始退 softcap,第 85% 步开始退 QK-norm。这样两个约束不是同时松开的,出问题好定位。
我见过同时退然后 loss 突然抽搐的案例,排查时两个调度都可能背锅,特别难搞。错开退是最好的可观测性策略。
5. 小模型到底要不要装:我的配置清单
5.1 先算一笔账:计算开销
QK-norm 会增加一个 RMSNorm 操作,而且是逐头做的。理论上每个 token 多两个 reduce 操作。实际开销取决于 kernel 融合程度。
我在一个 1.2B 模型训练流程里粗略测过:裸 attention 加 QK-norm 后训练吞吐大约下降 4%-6%;如果用自定义 fused RMSNorm kernel 或者让底层库把 norm 融合进 flash 前的输入处理,下降可以压到 1% 以内。对小模型来说,这个开销占比确实比大模型高,因为计算密度低、kernel launch 开销占比大,但也没有高到“不能接受”。
softcap 几乎没有参数和显存开销,唯一代价是对 logits 矩阵做额外计算;但如果你用 full attention、短上下文,这个代价可忽略。用 FlashAttention 且版本支持logit_softcap的话,代价也只是 kernel 里多一个 tanh,吞吐影响很小。
5.2 按规模与场景的推荐配置
我根据自己跑过的几个小模型经验,给一份保守但实用的配置清单:
| 模型规模 | 上下文 | 是否加QK-norm | 是否加softcap | 是否退火 |
|---|---|---|---|---|
| <0.5B | ≤4K | 不加或只加QK-norm | 不加 | 不需要 |
| <0.5B | 8K-16K | 加 | 建议加 | 预算够可尝试 |
| 0.5B-1.5B | ≤4K | 加(开销可接受) | 可选 | 可不加 |
| 0.5B-1.5B | 8K+ | 加 | 加 | 推荐,窗口最后15% |
| 2B-3B | 长上下文 | 加 | 加 | 推荐,结合softcap退火 |
表格里的“不加”不代表不管,而是要保证日志里有 logits 监控。我的原则:小模型可以少装保险丝,但不能不看仪表盘。
5.3 如果决定不装,必须满足哪些前提
如果你就是不想加 QK-norm,也可以,但建议满足下面几个条件:
- 上下文长度 4K 以内,且训练数据里没有大量重复文本或很长链接字符;
- 学习率严格做了 warmup 和 cosine 退火,峰值学习率不要超过原版对应模型常用值的 1.2 倍;
- 每 500-1000 步记录一次 attention logits 的 std 和 max,一旦发现 std 连续上升,立刻停住排查而不是硬跑。
满足这些条件的小模型,不加保险丝大概率也能平稳训完。问题是,大多数人无法保证自己的数据、学习率和调度都那么干净,所以装一层保险丝,其实是把风险外包给了结构,而不是赌自己运气。
5.4 替代方案和降级路线
如果你嫌 QK-norm 带来的修改太多,还有一个折中的替代:只用 softcap,不动 Q、K 的表示。改动少、兼容性好,适合快速验证一个数据 scale 是否稳定。
还有另一种“隐式保险丝”:把 attention logits 的初始化 scale 调小,比如原始实现默认 sqrt(head_dim) 缩放,改为乘一个 0.7 或 0.5 的常数。这种做法在最开始几万步很有效,但后期 logits 会慢慢回到大尺度,所以只能作为临时措施。
我建议的顺序是:先不加防护,看 logits 监控,如果稳定就别动;如果出现 spike 或 std 持续上升,优先加 softcap;还不够或长上下文,加 QK-norm;最后追求极致指标,用退火。从简到繁,每加一层都是因为前面的不够用,而不是因为别人都这么干。
6. 实战中我踩过的坑与排查速查
6.1 常见问题对照表
把我在训练中遇到的高频问题和处理方式整理成速查表,启动训练前可以先看看:
| 现象 | 可能原因 | 处理 |
|---|---|---|
| 加QK-norm后初始loss比原来高0.1-0.2 | QK-norm限制了logits尺度,模型需要重新适应 | 不是bug,多训几万步再对比 |
| logits std持续上升 | 学习率偏大 / 数据有长尾 / 未加softcap | 先检查LR,再加softcap |
| 加了QK-norm后loss反而一直不降 | τ被设置成sqrt(d)加上norm双重缩小 | 改τ为1或直接用无缩放 |
| 退火到一半出现loss spike | 退火窗口太短或学习率没降 | 延长窗口,LR降到峰值的1/10 |
| 某些head注意力变成one-hot | 温度过小/某个维度scale异常 | 检查该头logits分布,调整τ |
| softcap后loss不降 | c值太小,正常logits也被压平 | 把c放大到正常max的3倍以上 |
| FlashAttention下softcap导致显存爆 | 老版本kernel不支持softcap走了fallback | 换支持logit_softcap的注意力实现 |
6.2 几个值得固化的监控指标
除了 loss,我强烈建议训练日志里加入下面这几个与注意力相关的指标,无论最后装不装保险丝:
- 每层 logits 的 std 和 max,每 500-1000 步记录一次。如果 std 开始持续爬升,保险丝该出动了。
- 每层 attention 的 entropy 均值。entropy 掉到小于正常值的 1/3,说明有些 head 开始退化,优先排查那几个 head 对应的 norm 和 scale。
- 每层的 grad norm,如果某个层的 grad 明显大于其他层,且位置靠近注意力阶段,大概率是 logits tail 引起的。
这三个指标加起来,基本上能把“保险丝失效”和“别的训练问题”快速分开,省下大量 debug 时间。实际记录时不用每个 step 都打印,可以在训练循环里隔几百步对指定层做一次采样,存到 tensorboard 或 wandb,开销很低。
6.3 最后排雷
我最后再提醒几个容易忽略的细节:
第一,加了 QK-norm 后用 FlashAttention,记得确认当前版本是否对 head_dim 有对齐要求,有些 kernel 需要 head_dim 是 8 的倍数,如果你的 head_dim 是 128 或 64 就没事。
第二,不要在退火窗口内打开 grad checkpointing 的临时切换,这类训练配置的改动会引入额外随机性,退火阶段本来就敏感,尽量保持配置冻结。
第三,QK-norm 和 softcap 都不是“提高模型能力”的功能,它们只会让你“不炸”或“少炸”。如果你的模型指标本身低,大概率是数据、学习率、架构其他环节的问题,别指望加保险丝能救回来。
我自己跑了这几个组合后的最大感受是:小模型训练,稳定性问题不是“要不要买保险”而是“你愿不愿意看仪表盘”。QK-norm、softcap、退火 QK-norm 这三层,本质上是三种不同力度的约束,从永远保护到后期放开,每一步都可以用很小的成本换来可观的稳定性收益。
如果你现在正要启动一个小规模预训练,我建议从最简单的监控加起,出了问题再逐层加防护;如果你已经陷入 logits 爆炸的泥潭,直接上 QK-norm 加 softcap,别犹豫。等到模型跑到 80% 以后,再考虑要不要退火——那时候你手里有足够的数据来判断收益值不值这个复杂度。