Scaled Dot Product Attention
Scaled Dot-Product Attention
Attention 把“当前 query 应该从哪些 key 读取信息”变成一个可微分的加权聚合:
$$\operatorname{Attention}(Q,K,V)=\operatorname{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}+M\right)V$$
Q 表示当前要寻找什么,K 表示每个位置能被怎样匹配,V 表示匹配到后真正取回的内容。把 Q/K 的相似度和 V 的内容分开,是 attention 能够动态路由信息的原因。
一轮计算的因果链
1 | Q 与 K 做点积 |
| 张量 | 形状 | 归约/广播轴 |
|---|---|---|
| Query | [B,Hq,Tq,Dh] |
每个 query 保留 |
| Key | [B,Hkv,Tk,Dh] |
与 Q 在 Dh 上点积 |
| Scores | [B,Hq,Tq,Tk] |
softmax 沿 Tk |
| Value | [B,Hkv,Tk,Dh] |
按 Tk 加权求和 |
| Output | [B,Hq,Tq,Dh] |
每个 query 一个结果 |
在 MHA 中 Hkv=Hq;GQA/MQA 需要把 KV head 映射到多个 Q head,但不能把共享理解成“删除 Q head”。
为什么除以 sqrt(dk)
如果 Q/K 各维度近似零均值、方差固定,点积的方差会随 d_k 增长。维度越大,未缩放 logits 越容易让 softmax 接近 one-hot,梯度集中在少数位置。除以 sqrt(d_k) 把尺度拉回相近范围,使训练早期的权重和梯度更可控;它不是为了减少 FLOPs。
Mask 与复杂度
M 可以是 additive mask(禁止位置加负无穷)或布尔 mask,具体真值语义由 API 定义。Causal Mask 沿 query-key 坐标阻止未来 token;padding mask 通常沿 key 轴广播。标准全注意力的 score 矩阵有 Tq×Tk 项,长上下文的计算和临时显存因此成为瓶颈。
prefill 中许多 query 一起计算,GPU 更容易利用矩阵乘;decode 中 Tq=1,瓶颈更常转为读取历史 K/V 的内存带宽。KV Cache 与 PagedAttention 优化的是历史状态布局,不改变 attention 的数学定义。
实现检查与边界
- 先标出 softmax 轴,再检查 mask 是否在
[B,H,Tq,Tk]上按预期广播。 True=禁止与True=保留在不同 API 中可能相反;不能凭经验复制。- 不要在
forward内创建默认位于 CPU 的 scale tensor;标量head_dim ** -0.5更直接。 - fused attention kernel 可能不显式 materialize score 矩阵,但仍实现同一逻辑;“没看到矩阵”不等于没有
Tq×Tk的依赖。
它向上服务 Multi-Head Attention 和 Transformer Block,横向连接 Causal Mask 与 Positional Encoding and RoPE。
参考资料
Softmax 前后各自可能出错
QK^T 只是相似度,softmax 才把它变成沿 Tk 归一化的读取分布。实现通常会先减去每行最大值以避免指数溢出,再在 mask 后归一化;如果先把被禁止位置填成一个不够小的常数,长序列或低精度下它仍可能获得非零概率。反过来,如果一整行都被 mask,softmax 可能产生 NaN,这在 packed batch 的空片段和错误长度元数据里很常见。
dropout 也有明确边界:训练 attention dropout 作用于权重,不等价于随机丢弃 token;推理应关闭它,否则同一请求会因随机 mask 改变。fused/FlashAttention kernel 不输出完整 score 矩阵,节省了中间显存,但仍要满足同样的 causal、padding、scale 和 dropout 语义。验证 kernel 替换时应在小矩阵上逐元素比较输出和梯度,再测长序列性能。
把性能问题归因到正确的轴
当 Tq×Tk 很大时,attention 受算力和临时存储影响;当 decode 的 Tq=1 时,主要工作变成读取历史 K/V 并计算少量点积。于是同一模型会在 prefill 和 decode 呈现相反的瓶颈,不能用一次端到端平均延迟解释。记录 Tq/Tk/Hkv/Dh、实际读写字节和 kernel occupancy,才能判断该换 fused attention、压缩 KV,还是调整 batching。