QK-Norm(Query-Key Normalization,查询-键归一化)指的是:在注意力(attention)里算分数之前,先把查询向量 Q 和键向量 K 各做一次归一化,再让它们做点积。目的很直接——别让注意力分数(logits)越训越大,最后把训练带跑偏。
为什么需要它
自注意力的核心公式是:注意力权重 = softmax(QKᵀ / √d)。每个 token 会生成三个向量:查询(Query,Q)、键(Key,K)、值(Value,V)。Q 和 K 都是模型自己学出来的,训练过程中它们的"长度"(范数)可能一路往上涨。两个长向量做点积,结果自然就大;维度越高,这种放大越明显。
麻烦在于 softmax 对输入尺度极其敏感。logits 一小,输出是温和的分布,注意力均匀分给多个 token;logits 一大,输出就趋近 one-hot,某个 token 独吞几乎全部权重,其余接近零。这时候梯度要么消失要么爆炸,loss 曲线上冒出尖峰,训练直接跑飞。这是大模型训练里最典型的翻车方式之一。
打个比方:Q 和 K 像 KTV 调音台上的两个增益旋钮,而模型自己负责拧。它一兴奋就把两个都拧到底,混音台立刻爆音。softmax 相当于后面那个削波压缩器,能救一点,但声音已经失真了。更靠谱的办法是给每路话筒先装一个自动增益:不管说话人多大声,输出电平都被拉回合理区间。QK-Norm 干的就是这件事。
具体做法通常是对 Q 和 K 分别做层归一化(AI 词典:LayerNorm">LayerNorm)或均方根归一化(RMSNorm),把向量长度归到统一尺度,保留方向信息,再算点积。这样点积的幅度不再随训练无限增长,softmax 分布保持在"可软可硬"的健康区间,梯度稳定,注意力也不容易塌缩成只盯一个位置。
它和几个邻居的区别
| 做法 | 作用位置 | 机制 | 特点 |
|---|---|---|---|
| 1/√d 缩放 | 点积之后 | 固定常数缩放 | 便宜,但管不住不同 token 之间的范数差异 |
| QK-Norm | 点积之前 | 对 Q、K 逐 token 归一化 | 自适应,随训练动态变化 |
| logits 截断 / 软上限 | logits 上 | 硬性截断或压顶 | 简单粗暴,可能损失区分度 |
| Pre-LN 里的 LayerNorm | 子层入口 | 归一化整个子层输入 | 稳的是残差流,不专门管注意力分数 |
一句话:常规的 1/√d 是"一刀切的固定缩放",QK-Norm 是"每个 token 各调各的自动增益"。
对实际工作的意义
对做训练的人来说,QK-Norm 是一个性价比很高的稳定性开关:多了两次归一化,算力开销几乎可以忽略,却常常能消掉早期的 loss 尖峰,让大模型训得更稳、更少需要回滚重跑。近几年不少超大规模视觉和语言模型都用它来压住注意力 logits。具体支持情况以各框架官方页面为准。
对普通职场人来说,它属于"看不见但受益"的那类改动——你不会直接感受到它,但模型训练更稳,意味着更少的返工、更可复现的结果,以及更少因为一次训练炸掉而浪费的时间和电费。
