Transformer Attention 数学:Q/K/V、Softmax 权重、Mask 与 KV Cache
Transformer Attention 数学:Q/K/V、Softmax 权重、Mask 与 KV Cache
站内搜索
直接问 AI

Transformer Attention 数学:Q/K/V、Softmax 权重、Mask 与 KV Cache

Transformer 的注意力机制(Self-Attention)可以被通俗地理解成:序列中的每个词(token)用自己的查询向量(Query)去评估所有词的键向量(Key),从而决定对整个句子的“注意力分配”,最后再用这个分配权重去加权每个词的信息向量(Value)。它的数学形式仅仅短短一行代码,但在工程落地和模型训练中却暗藏无数细节。

公式本身好背,难的是知道每一步在算什么。所以这篇用 3 个 token 从头手算一遍 Scaled Dot-Product Attention——每个中间矩阵都写出来,可以自己对着验算。手算完之后再看 Q/K/V 投影为什么要分开、除以根号 d 是在防什么、因果掩码具体掩掉哪些位置、多头拆开之后每个头看到的是什么,以及推理时 KV Cache 存的到底是哪一部分。

一、核心数学公式解析

这是整个大语言模型时代的基石公式:

Attention(Q, K, V) = softmax((Q @ K^T) / sqrt(d_k)) @ V

其中:

  • Q @ K^T 产生的是一个 `[seq_len, seq_len]` 大小的注意力分数矩阵。由于是点积,它衡量了每对 token 之间在多维空间中的“相似度”或“关联度”。
  • 为什么必须除以 sqrt(d_k)?假设 Q 和 K 的维度 `d_k = 4096`,且元素服从均值为 0,方差为 1 的独立分布,那么点积的方差会高达 `4096`。方差过大会导致极端的分数值(如 100 和 -100),在经过 Softmax 时就会导致梯度几乎为零(梯度消失),即“Softmax 饱和”。

二、架构图解:数据流与维度变化


graph TD
    Input[Input Sequence: B, L, d_model] --> WQ(W_q Linear)
    Input --> WK(W_k Linear)
    Input --> WV(W_v Linear)
    
    WQ --> Q[Q: B, h, L, d_k]
    WK --> K[K: B, h, L, d_k]
    WV --> V[V: B, h, L, d_v]
    
    Q --> Dot[Dot Product: Q @ K^T]
    K --> Dot
    Dot --> Scale[Scale by 1/sqrt(d_k)]
    Scale --> Mask[Apply Causal Mask]
    Mask --> Softmax[Softmax along dim L]
    Softmax --> AttentionWeights[Attention Weights: B, h, L, L]
    
    AttentionWeights --> MatMulV[MatMul with V]
    V --> MatMulV
    
    MatMulV --> Context[Context Output: B, h, L, d_v]
    Context --> Concat[Concat Heads: B, L, d_model]
    Concat --> Out[W_o Linear]

三、实战演示:用 Numpy 手写自注意力

光看公式太抽象,我们来跑一段可执行的 Numpy 纯手写代码。假设输入是一个只有 3 个 token(例如 “AI”, “needs”, “math”)的序列,维度为 4:

import numpy as np

# 1. 模拟 Q, K, V 矩阵 (Seq_len=3, d_k=4)
# 代表 "AI", "needs", "math" 三个词
Q = np.array([
    [ 1.0,  0.5, -0.2,  0.1],  # AI
    [-0.5,  1.2,  0.8, -0.4],  # needs
    [ 0.2, -0.1,  1.5,  0.9]   # math
])
K = np.array([
    [ 0.8,  0.4, -0.3,  0.0],
    [-0.2,  1.0,  0.5, -0.1],
    [ 0.1, -0.2,  1.1,  0.7]
])
V = np.array([
    [ 1.0,  0.0],
    [ 0.0,  1.0],
    [-1.0, -1.0]
])

d_k = Q.shape[1]

# 2. 计算打分 (Scores) 并进行缩放 (Scaling)
scores = (Q @ K.T) / np.sqrt(d_k)
print("Scaled Scores:\\n", scores)

# 3. 因果掩码 (Causal Mask)
# 屏蔽未来位置,防止模型作弊
mask = np.triu(np.ones((3, 3)), k=1)
scores[mask == 1] = -np.inf

# 4. Softmax 归一化
def softmax(x):
    e_x = np.exp(x - np.max(x, axis=-1, keepdims=True))
    return e_x / e_x.sum(axis=-1, keepdims=True)

weights = softmax(scores)
print("Attention Weights:\\n", np.round(weights, 3))

# 5. 值加权 (Context)
context = weights @ V
print("Context Output:\\n", context)

跑完这段代码你会发现:第一行对应 “AI” 这个词,它的注意力权重只会分配给自己;而第三行的 “math” 会将注意力分配给前两个词。这就体现了自回归模型的本质:只能用历史信息生成未来信息。

四、因果掩码(Mask)到底改变了什么?

正如上面代码所示,在自回归(Autoregressive)生成任务中,如果当前在预测第 3 个词,它绝不能“看到”第 4、5 个词的信息。我们在 Softmax 之前,强行把上三角矩阵的注意力分数赋值为负无穷大(-inf)。经过 Softmax 后,这些位置的权重会被精确地压为 0。所以,Mask 不是删除 token,而是在概率层面做切断,让非法的注意力分配变成绝对不可能发生的事。

五、工程师的填坑经验:显存杀手与 KV Cache

实战视角:在书本上你看到的是优雅的矩阵公式,但在工业界部署 LLM 时,你看到的往往是一次次无情的 OOM (Out of Memory) 报错。

在推理阶段,大模型是以逐字生成(Token-by-token)的方式运行的。生成第 $t+1$ 个词时,前面的 $t$ 个词的 K 和 V 都是完全不变的!如果我们每次都用全尺寸的 L x d_model 矩阵去重新乘一遍,那就是巨大的算力浪费。

KV Cache 的本质,就是用空间换时间。

  • 我们会在 GPU 显存里开辟一块连续区域,把历史生成的 K 和 V 保存下来。
  • 每生成一个新词,只需要计算当前这 1 个 token 的 $Q_{new}, K_{new}, V_{new}$,然后把 $K_{new}$ 拼接到显存里。
  • 代价极其高昂:一个稍微长一点的上下文,哪怕只有 10K tokens,单 batch 消耗的 KV Cache 可能就会超过模型权重本身的显存占用!这就是为什么现在工业界会发明 PagedAttention(vLLM 的核心)、MQA (Multi-Query Attention) 和 GQA (Grouped-Query Attention),全都是为了削减 KV Cache 的显存体积。

六、实现时最容易错的三个 shape

第一处是 batch 维度。教学代码经常写成 Q @ K.T,这只适合单个序列;真实模型通常是 batch x heads x tokens x dim。这时应当转置最后两个维度,而不是把 batch 或 head 维度也混进去。shape 写错时,程序有时不会报错,只会通过广播得到完全错误的注意力矩阵。

第二处是 mask 维度。自回归 mask 应该覆盖 query-token 到 key-token 的二维关系,并且能广播到 batch 和 head。padding mask 则表示哪些 token 是填充位。两类 mask 的语义不同,不能简单相加了事。第三处是 softmax 的轴,必须沿 key 维度归一化;如果沿 query 维度归一化,每一列会变成概率分布,注意力含义就反了。

七、怎么检查 attention 实验结果

最基础的检查是每一行 attention weight 的和是否接近 1。然后检查 mask 后的未来位置是否接近 0。再检查 context 的 shape 是否和 Value 的最后一维一致。对于这篇文章里的三个 token toy example,你还可以手算第一行 softmax,确认权重变化不是因为代码排序错误或 mask 方向写反了。

注意力热力图适合调试,但不等价于完整解释。一个 token 权重高,只表示这一步加权读取更多地使用了某个 Value;它不直接证明模型“因为什么原因”做出最终预测。把 heatmap 当成排查工具,而不是因果证据,能避免很多误读。

八、Attention 验证矩阵

自注意力实现最容易出现“shape 能跑但语义错”的问题。下面的矩阵把检查点固定下来,读者可以用它复核本文的 NumPy toy example,也可以迁移到批量、多头或推理缓存实现中。

检查点 正确证据 常见错误
score 形状 Q @ K.T 得到 query-token 到 key-token 的二维矩阵。 转置错维度,把 batch/head 维度混进注意力矩阵。
缩放与 softmax 除以 sqrt(d_k) 后沿 key 维度归一化,每行和约等于 1。 沿 query 维度 softmax,或不缩放导致权重过早饱和。
causal mask 未来位置在 softmax 后接近 0,历史位置仍可分配权重。 mask 方向反了,让当前 token 只能看未来而不能看历史。
KV Cache 新 token 只追加 K_newV_new,历史缓存不重复计算。 每步重算全部 K/V,或 cache 长度与位置编码不同步。

九、图示与数据流总结

三个 token 的 scaled dot-product attention 权重热力图
每一行代表一个 Query token 对所有历史 Key token 的注意力分布。这就是所谓的 Attention Heatmap,模型之所以“懂”语言,就藏在这张图里的每一丝权重变化中。

这套机制看似只是矩阵乘法,但却支撑起了当今很多前沿 AI 系统。下次再遇到 Transformer 报错时,第一反应应该是:打印所有张量的 shape,然后在纸上画一遍矩阵乘法的过程

发表回复

向下探索