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

注意力代码最难排查的错误,往往不是 shape 报错,而是返回了形状正确、含义却错误的结果:读取了未来 token、沿错误的轴归一化,或者缓存里明明有历史信息却没有读取。本篇围绕同一组 3 个 token 的输入,逐步核对打分、掩码、权重、输出和增量解码。

实验修订日期:2026 年 9 月 8 日。请用本文的 Attention 与 KV Cache 实验包复算。页面其他位置关联的旧版 Deep Learning Math Lab 使用另一组无掩码输入,不能核对下面的因果注意力图表。本实验不需要训练好的模型或 GPU。

一、先把公式中的形状说清楚

scores = (Q @ transpose_last_two_axes(K)) / sqrt(d_k)
weights = softmax(masked_scores, axis=keys)
context = weights @ V

单头计算中,Q 的形状是 (Lq, dk),K 是 (Lk, dk),V 是 (Lk, dv)。分数矩阵为 (Lq, Lk),输出为 (Lq, dv)。Lq 与 Lk 相等只是其中一种情况;后面的缓存反例只有一个 Query,却有三个 Key。

缩放的动机来自一个明确假设:若 Q/K 各分量相互独立、均值为零、方差为一,点积的方差为 dk,除以其平方根后方差为一。这不是对训练后激活值的保证。dk 是每个头的维度,不能直接当成模型总宽度。公式出处见 Attention Is All You Need。

分开的 Q/K/V 投影允许查询、匹配键和待读取内容采用不同坐标。本例直接给出投影之后的数值,没有训练投影矩阵,也没有证明这些向量学到了“AI needs math”的词义。

二、多头数据流与维度

graph TD
    Input[Input: B, L, d_model] --> WQ[W_q projection and split]
    Input --> WK[W_k projection and split]
    Input --> WV[W_v projection and split]
    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[Q times transposed K: B, h, L, L]
    K --> Dot
    Dot --> Scale[Divide by sqrt d_k]
    Scale --> Mask[Mask future keys]
    Mask --> Softmax[Softmax over keys]
    Softmax --> Read[Weights times V]
    V --> Read
    Read --> Context[Context: B, h, L, d_v]
    Context --> Concat[Concatenate: B, L, h times d_v]
    Concat --> Out[W_o projection: B, L, d_model]

批量计算时只转置 K 的最后两个轴。NumPy 的 K.T 会反转四维张量的全部轴,应该使用 K.swapaxes(-1, -2)。拼接多个头后宽度为 h * dv,再由输出投影映射到 dmodel;两者只有在配置满足条件时才相等。

三、同一组输入的全部中间结果

位置编号为 0、1、2,每个 Query 可以读取自己及之前的位置。Q/K 维度为 4,V 维度为 2。下一 token 训练通常用输入位置 i 的表示预测 i+1 的目标,所以“当前输入”和“接下来要预测的目标”不能混为一谈。

查看可直接运行的 NumPy 示例
import numpy as np

# Fixed illustrative tensors, not embeddings from a trained model.
Q = np.array([[1.0, 0.5, -0.2, 0.1],
              [-0.5, 1.2, 0.8, -0.4],
              [0.2, -0.1, 1.5, 0.9]], dtype=np.float64)
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]], dtype=np.float64)
V = np.array([[1.0, 0.0], [0.0, 1.0], [-1.0, -1.0]])

scores = (Q @ K.T) / np.sqrt(Q.shape[-1])
allow = np.arange(3)[None, :] <= np.arange(3)[:, None]
masked = np.where(allow, scores, -np.inf)
shifted = masked - masked.max(axis=-1, keepdims=True)
weights = np.exp(shifted)
weights /= weights.sum(axis=-1, keepdims=True)
context = weights @ V

if __name__ == "__main__":
    np.set_printoptions(precision=6, suppress=True)
    print("Scaled scores:\n", scores)
    print("Masked scores:\n", masked)
    print("Attention weights:\n", weights)
    print("Context:\n", context)

第一个分数可以直接手算:(1*0.8 + 0.5*0.4 + (-0.2)*(-0.3) + 0.1*0) / 2 = 0.53。全部点积、掩码和输出如下:

Scaled scores                 After causal masking
[[ 0.530  0.095 -0.075]        [[ 0.530   -inf   -inf]
 [-0.080  0.870  0.155]         [-0.080  0.870   -inf]
 [-0.165  0.260  1.160]]        [-0.165  0.260  1.160]]

Attention weights             Context = weights @ V
[[1.000000 0.000000 0.000000]  [[ 1.000000  0.000000]
 [0.278885 0.721115 0.000000]   [ 0.278885  0.721115]
 [0.158938 0.243109 0.597953]]  [-0.439015 -0.354843]]

以 Query 1 为例,减去本行最大可见分数后得到 [-0.95, 0, -inf],归一化分母为 exp(-0.95) + 1,两个可见权重约为 0.278885 和 0.721115。Query 2 也会读取自己;由于 V2 为 [-1,-1],它的输出就是 [w20 - w22, w21 - w22]。这解释了两个负分量的来源,不需要借助语言模型语义。

与正文同一组 dk4 输入的三 token 因果注意力热力图,三个未来位置权重均为零
plot_attention.py 直接读取上面的数组生成此图。灰格代表禁止读取,并非数据缺失。这些是注意力权重,不是梯度,也不是模型理解语言的证据。

四、只有一个 Query 时,缓存掩码为什么容易写错

从位置零开始计算完整序列时,3×3 的下三角掩码是正确的。但最后一个缓存 Query 位于绝对位置 2,不能重新把它当成位置零,直接套用 1×3 的普通下三角掩码。

掩码 输出 (x, y)
偏移 2:读取 Key 0 至 2 (-0.439, -0.355)
错误偏移 0:只读 Key 0 (1.000, 0.000)

错误输出的最大分量误差为 1.439015。它的权重行和仍然为一,输出形状也正确,因此只检查行和无法发现这个问题。应该按绝对位置建立可见关系:key_position <= query_start + local_query_position。本文完整前缀缓存的 query_start,就是已经缓存的 token 数。

迁移到库函数时也要核对约定。PyTorch 2.14 的 SDPA 文档说明非方形因果掩码采用左上对齐,其布尔 attention mask 中 True 表示允许参与。不要默认其他接口的 padding mask 与它同义。这里核对的是文档约定,并未运行 PyTorch 数值一致性测试。

实验实现会拒绝整行都被遮挡的 Query。否则减去行最大值时可能出现 -inf - -inf,继而得到 NaN。这是本文实验包主动规定的行为,不代表所有框架的统一行为。因果可见性与 padding 可见性应按逻辑条件组合,并在 softmax 前明确完全填充的 Query 如何处理。

五、实际运行过的缓存与数值检查

参考环境为 Python 3.13.9、NumPy 2.3.5、float64。下载包保留源文件哈希、完整精度 CSV 和 审计 JSON。另外用独立循环与 Decimal.exp 做了 50 位十进制精度计算,输入从原始 binary64 数值精确转换,避免把十进制字面量与二进制浮点数混为一谈。

检查 结果
Decimal 输出 5.6e-17
逐 token 缓存 [1,1,1] 5.6e-17
分块缓存 [2,1] 0
批量缓存 [2,3,2] 2.2e-16
修改未来 K/V 0
非法输入 拒绝 12 类

前五项结果为输出的最大绝对误差。Decimal 对照覆盖全部 9 个分数、9 个权重和 6 个输出,最大分数误差为 1.1e-16。批量测试采用 B=2、H=3、L=7、dk=4、dv=5 和种子 42;缓存同时追加 K/V 并使用正确的前缀偏移。修改未来 K/V 时,前两个输出保持不变。拒绝的输入包括非有限值、全遮挡行、错误形状和非法偏移。

另设两个故意写错的反例:反转因果掩码后,禁止读取的未来位置权重总量为 0.872569;沿列归一化后,最大行和误差为 0.596413。测试均检测到了错误。另外,缓存追加失败后,原有 K/V 必须保持不变。

缓存测试确实追加了数组,并检查历史 K/V 未被改变,但它不是完整 Transformer:Q/K/V 已经预先给定。重复使用 NumPy concatenate 会复制历史数组,不适合作为高效推理服务实现。本实验没有测量加速比、GPU 内核一致性、模型精度或文本生成质量。

六、先算清 KV 字节数,再讨论显存不足

当 K/V 每头维度和数据类型一致时,原始缓存负载为:

bytes = 2 * layers * batch * tokens * kv_heads * head_dim * bytes_per_element

取 32 层、batch 为 1、10,000 tokens、每头维度 128、每元素 2 字节,得到下表。这里是公式估算,不是 GPU 显存实测:

KV 头数 原始 GiB 对应情形
32 4.8828 32 个 Query 头的 MHA
8 1.2207 32 个 Query 头的 GQA
1 0.1526 32 个 Query 头的 MQA

脚本还用一对实际 NumPy fp16 小数组核对公式,总计 480 字节,没有申请表中的大块内存。作为比较,恰好 70 亿个双字节参数的原始权重约为 13.04 GiB;因此本例 10K-token MHA 缓存并没有超过这些权重。要判断是否 OOM,还必须计入激活、内存分配、元数据、batch 和具体架构。

PagedAttention处理缓存分配浪费与共享,分页本身不会减少上式中的原始元素数;GQA/MQA则减少 KV 头数。两者不是同一种节省机制,也不能据此声称缓存必须占用一整块连续显存。本实验包没有实现这些推理系统。

七、如何复现并主动制造一个错误

解压实验 ZIP,在解压目录运行:

python3 -m venv .venv
source .venv/bin/activate
python -m pip install -r requirements.txt
python attention_example.py
python audit_attention.py --out results

成功时最后输出 ATTENTION_AUDIT_OK,默认 results 目录不会覆盖包内 reference 参考结果。重新画图需安装 requirements-plot.txt,再运行 plot_attention.py。CSV 保留 17 位有效数字;跨平台比较应使用容差,不要求最后几个浮点位完全相同。

可以做两个有明确预期的练习:先把最后一次查询的偏移从 2 改成 0,复现错误的 [1,0];恢复后只修改最后一个 V,前两个因果输出应完全不变。这是在检验可见性,而不是只看热力图是否“像是正确的”。

八、结论与适用边界

因果推理中的缓存复用依赖前缀、模型参数和位置约定保持不变;修改前缀或改变位置后可能需要失效重算,训练和 dropout 也要单独处理。Hugging Face 的缓存说明介绍了每层历史 K/V 与新 token 数据的组合方式。

本文给出的证据是可复算的算术与掩码测试,不是因果可解释性或生产部署保证。下一步可以对照反向传播计算图实验:那里用有限差分检查导数,这里检查注意力的读取边界,两类证据不能互相替代。

发表回复

向下探索