注意力代码最难排查的错误,往往不是 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]。这解释了两个负分量的来源,不需要借助语言模型语义。

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