跳转到主要内容

Causal Attention:为什么 KV hit 后 Attention 按 1 - h² 缩放

用 causal attention 的下三角区域推导 prefix cache hit 后 attention 近似按 1-h² 缩放

· 约 3 分钟阅读

Prefix cache 命中前 hNhN 个 token 后,新 suffix 仍要访问 cached prefix。Attention 省掉的是 prefix-prefix 的二维三角区域,而不是简单省掉 hh 比例的所有工作。

30 秒复习
  • 一句话:Dense causal prefill 的工作量像下三角;命中 prefix 后去掉左上角小三角,所以剩余比例近似为 1-h²
  • 三个判断:suffix-prefix 矩形仍需计算;逐 token projection/FFN 更接近 1-h;理论交互比例不等于 kernel 延迟比例。
  • 核心模型W_hit = M·H + M(M+1)/2,其中 H=hNM=N-H,因此 W_hit/W_0 ≈ 1-h²
  • 边界:只适用于复用连续 prefix 的 dense causal 可见区域;滑窗、稀疏、CSA/HCA、递推结构及拆分后的 projection 都需重建自己的缩放模型。

1. 0-hit 是一个下三角

对长度为 NN 的 causal prompt,第 ii 个 Query 只能看位置 1i1\ldots i

       key/value token
       1 2 3 4 5 6 7 8
query
1      x
2      x x
3      x x x
4      x x x x
5      x x x x x
6      x x x x x x
7      x x x x x x x
8      x x x x x x x x

若把每个可见 Query-Key 交互记作一份工作:

W0=i=1Ni=N(N+1)2N22W_0=\sum_{i=1}^{N}i =\frac{N(N+1)}{2} \approx\frac{N^2}{2}

这里讨论的是 attention 交互区域,不是整个 Transformer block 的 FLOPs 或实测时延。

2. Prefix hit 省掉哪一块

令命中长度为 H=hNH=hN,未命中 suffix 长度为 M=NHM=N-H。以 8 个 token、前 4 个命中为例:

       1 2 3 4 | 5 6 7 8
1      .
2      . .
3      . . .
4      . . . .
-------------------------
5      x x x x | x
6      x x x x | x x
7      x x x x | x x x
8      x x x x | x x x x

. 是已复用的 prefix-prefix 区域。suffix 每一行仍需读 prefix K/V,因此剩余工作包含:

  1. suffix 对 prefix 的矩形:M×HM\times H
  2. suffix 内部的下三角:M(M+1)/2M(M+1)/2

3. 1-h² 的推导

剩余交互量为:

Whit=MH+M(M+1)2W_{\text{hit}} =MH+\frac{M(M+1)}{2}

也可以直接从完整三角减去命中三角:

Whit=N(N+1)H(H+1)2W_{\text{hit}} =\frac{N(N+1)-H(H+1)}{2}

所以精确离散比例为:

WhitW0=1H(H+1)N(N+1)\frac{W_{\text{hit}}}{W_0} =1-\frac{H(H+1)}{N(N+1)}

NN 足够大且 H=hNH=hN

WhitW01h2\frac{W_{\text{hit}}}{W_0} \approx1-h^2

因此常用一阶近似是:

Tattn(h)Tattn,0(1h2)T_{\text{attn}}(h) \approx T_{\text{attn},0}(1-h^2)

它表达的是几何区域比例;是否能用于时间缩放,还要看 bucket 里包含哪些算子。

4. 为什么 FFN 更接近 1-h

FFN/MLP 以及 QKV/O projection 是逐新 token 执行的线性层。复用 HH 个 prefix token 后,只需处理 MM 个 suffix token,因此计算量更接近:

Ttoken-linear(h)Ttoken-linear,0(1h)T_{\text{token-linear}}(h) \approx T_{\text{token-linear},0}(1-h)

h=0.5h=0.5

逐 token 线性部分:剩 50%
dense causal 交互:剩 75%

这解释了为什么不能对整个 prefill breakdown 统一乘 1-h,也不能把包含 projection 的总 “attention” bucket 全部无条件乘 1-h²

5. 如何交给端到端模型

本页的输出只有一个局部控制量:

dense causal interaction scale ≈ 1 - h²

端到端 TTFT 还要加入 FFN/projection、固定开销、KV load、传输与计算 overlap,并与 decode chip-time 合并。这些内容由 KV Cache Hit Ratio 修正模型 负责,本页不再展开第二套模型。

6. 适用边界

使用 1-h² 前检查:

  1. 命中的是从位置 0 开始的连续 prefix,而不是离散块。
  2. 可见区域是标准 dense causal 下三角。
  3. 目标 bucket 是 Query-Key/Value 历史交互,而非全部 attention projections。
  4. 工作量近似与目标 kernel 的时延趋势经过校准。
  5. cache 已可用;若需从 Host/远端加载,传输时间单独进入 TTFT。

Sliding Window、Block Sparse、CSA/HCA 的可见区域不同;MLA 虽压缩 payload,但若仍访问完整历史,其 token 交互区域可以相同、每次交互的字节和 kernel 路径却不同。

相关页面