Attention Cache and Sequence Parallelism
导言
KV cache 最容易造成的误解,是把所有名为 K/V 或 state 的张量都看成同一种“缓存”。事实上,训练时为反向传播保留的 K/V activation、推理时跨 decode step 存活的 KV cache、GDN 跨 token 改写的固定状态,生命周期和分布式处理都不同。
本文从一次 token 生成开始,逐对象解释 MHA、DSA 与 GDN 保存什么,再讨论 CP、USP 切开长序列后,临时 activation 和持久 cache 分别需要怎样的通信。
结论¶
先记住六个判断:
- KV cache 主要服务自回归推理。Prefill 一次构建 prompt 的逐层历史,decode 每步追加新 K/V;旧 Q 不缓存,因为未来不会再用它去发起查询。
- 训练通常没有跨 forward 持久增长的 KV cache。Teacher forcing 已知整段目标,一次并行形成所有 Q/K/V;K/V 只是当前 batch 的 activation,为 backward 保存或重算。
- MHA 保存逐 token、逐 KV head 的 K/V,cache 随上下文长度线性增长。
- DSA 仍有两套随长度增长的历史:Lightning Indexer 的低维 K/scale,以及 MLA 的
cKV/kR;Indexer 扫描全历史,主 MLA 只读取 global top-k。 - GDN 没有逐 token KV cache。它保存固定大小矩阵状态和有界短卷积状态;代价是历史被压进有限容量,而不是无损保留每个 token。
- CP/USP 的训练通信不能直接等同于 decode cache 分片。训练时 Q/K/V 都有长序列维;decode 的 Query 长度通常只有 1,若把历史 cache 按 token 分片,还要额外做 distributed softmax 或 global top-k。
先分清三种“保存”¶
为什么只缓存 K/V¶
在第 l 层、第 t 个位置,MHA 先把当前层输入 h_t^l∈R^D 投影成:
q_t^l = reshape(h_t^l W_Q^l) [B,Hq,1,d]
k_t^l = reshape(h_t^l W_K^l) [B,Hkv,1,d]
v_t^l = reshape(h_t^l W_V^l) [B,Hkv,1,d]
对于 causal decoder,历史 token j≤t 的 k_j^l,v_j^l 一旦算出,以后不会因为生成了新 token 而改变。第 t+1 步的新 Query 要再次读取它们,所以 cache 执行:
K_cache^l ← concat(K_cache^l, k_t^l, axis=time)
V_cache^l ← concat(V_cache^l, v_t^l, axis=time)
a_t^l = (q_t^l (K_cache^l)^T) / sqrt(d) + causal_mask
p_t^l = softmax(a_t^l, axis=history)
o_t^l = p_t^l V_cache^l
旧 q_j^l 的工作在第 j 步已经结束;未来发起查询的是新 q_t^l,所以通常没有 Q_cache。Transformers 的通用动态 cache 也按层保存 shape 为 [B,Hkv,T,d] 的 K 和 V,并沿倒数第二维 append。1
Prefill 与 decode¶
- Prefill:输入 prompt 的
S个 token 已知,可以用矩阵乘法一次形成所有层的 Q/K/V,并把各层 K/V 写入初始 cache。它不是逐 token Python 循环,但 causal mask 仍保证位置i看不到未来位置。 - Decode:每步通常只有一个新 token。模型只为它计算新 Q/K/V,把新 K/V 追加到各层 cache;新 Q 读取从 prompt 到当前步的全部历史。
- 每层独立保存:第
l层的 K/V 来自该层输入,不能拿第l-1层 cache 代替。
以 BF16、L=32 层、B=1、Hkv=32、d=128、T=8192 为例,忽略对齐和管理开销:
这里第一个 2 是 K 与 V 两份,最后一个 2 bytes 是 BF16 元素大小。换成 GQA/MQA 会减少 Hkv,但 attention score 仍要覆盖全部历史 token。
训练为什么不同¶
训练使用 teacher forcing:目标序列在一次 forward 前已经全部给定。所有位置的 Q/K/V 可以并行计算,causal mask 只限制可见性。K/V 可能因 autograd 被保存到 backward,也可能由 activation checkpointing 重算,但它们通常在该次迭代结束后释放,不会在下一批数据上继续 append。
同一个 use_cache 不代表训练也该开
Hugging Face 官方文档明确把 cache 限定为 inference;训练时启用可能引发非预期错误。1 训练图需要梯度、dropout 和完整序列语义,生成 cache 的原地更新与跨 forward 生命周期并不匹配这些要求。
MHA:逐 token 保存 K/V¶
MHA 的对象账本如下。为了突出 cache,省略 MLP 和残差分支。
| 对象 | 公式 / 代码名 | Shape | 生产者 → 消费者 | 生命周期 |
|---|---|---|---|---|
| 当前层输入 | h_t / hidden_states |
[B,1,D] |
上一子层 → Q/K/V 投影 | 临时 activation |
| 当前 Query | q_t / query_states |
[B,Hq,1,d] |
WQ、reshape、RoPE → QK MatMul |
临时 activation |
| 当前 Key/Value | k_t,v_t |
[B,Hkv,1,d] |
WK/WV、reshape、RoPE(K) → cache update |
临时 activation |
| 历史 Key/Value | K_cache,V_cache |
[B,Hkv,T,d] |
旧 cache + 当前 K/V → score/value MatMul | 跨 decode step cache |
| 分数与概率 | a_t,p_t |
[B,Hq,1,T] |
QK MatMul、scale、mask、Softmax → PV MatMul | 临时 activation |
| 输出 | o_t,y_t |
[B,Hq,1,d] → [B,1,D] |
PV MatMul、concat、WO → 下一子层 |
临时 activation |
下面是完整的单层 decode 伪代码。repeat_kv 表示 GQA/MQA 为计算建立的逻辑 head view;MHA 中 Hkv=Hq,它是恒等操作。实际 kernel 可以不物化复制。
def mha_decode(hidden_t, layer_cache, wq, wk, wv, wo, rope, scale):
batch, one, model_dim = hidden_t.shape
q = (hidden_t @ wq).view(batch, one, Hq, head_dim).transpose(1, 2)
k = (hidden_t @ wk).view(batch, one, Hkv, head_dim).transpose(1, 2)
v = (hidden_t @ wv).view(batch, one, Hkv, head_dim).transpose(1, 2)
q = rope(q, position=layer_cache.length)
k = rope(k, position=layer_cache.length)
k_all = concat([layer_cache.k, k], dim=2)
v_all = concat([layer_cache.v, v], dim=2)
layer_cache.k = k_all
layer_cache.v = v_all
k_for_q = repeat_kv(k_all, groups=Hq // Hkv)
v_for_q = repeat_kv(v_all, groups=Hq // Hkv)
scores = matmul(q, transpose(k_for_q, -1, -2)) * scale
probs = softmax(scores, dim=-1)
heads = matmul(probs, v_for_q)
merged = heads.transpose(1, 2).reshape(batch, one, model_dim)
return merged @ wo, layer_cache
DSA:两套随序列增长的压缩缓存¶
DSA 不是把 MHA cache 直接截成 top-k。DeepSeek-V3.2-Exp 在 MLA 前放置 Lightning Indexer:先用低维 Key 为每个 Query 选择历史位置,再让主 MLA 只读取这些位置。5
| 对象 | 源码 / 公式名 | Shape | 生产者 → 消费者 | 生命周期 |
|---|---|---|---|---|
| Indexer Query | q_i |
[B,Hidx,Didx] |
当前 hidden → Indexer Q 投影、RoPE、量化 | 临时 activation |
| Indexer Key | k_j |
[B,T,Didx] FP8 + scale |
每个历史 hidden → k_cache/k_scale_cache |
随 T 增长的 cache |
| Indexer head 权重 | w_i / weights |
[B,Hidx] |
当前 hidden → 加权跨 head score | 临时 activation |
| Index score | I_j |
[B,T] |
FP8 dot、ReLU、跨 head 求和 → top-k | 临时 activation |
| 选择位置 | J=TopK(I,k) |
[B,k] |
Indexer → 主 MLA gather / mask | 临时控制对象 |
| MLA latent cache | cKV_j,kR_j |
[B,T,Rkv]、[B,T,DR] |
KV down projection、RMSNorm、RoPE → 主 MLA | 随 T 增长的 cache |
| 主 MLA 输出 | y_t |
[B,D] |
只聚合 J 中 latent → WUV/WO |
临时 activation |
固定 revision 的 Indexer 把低维 K 和量化 scale 物理写入两块 cache,再对全历史求分数:7
q_fp8, q_scale = act_quant(q, block_size, self.scale_fmt)
k_fp8, k_scale = act_quant(k, block_size, self.scale_fmt)
self.k_cache[:bsz, start_pos:end_pos] = k_fp8
self.k_scale_cache[:bsz, start_pos:end_pos] = k_scale
weights = self.weights_proj(x.float()) * self.n_heads ** -0.5
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
index_score = fp8_index(
q_fp8,
weights,
self.k_cache[:bsz, :end_pos],
self.k_scale_cache[:bsz, :end_pos],
)
topk_indices = index_score.topk(min(self.index_topk, end_pos), dim=-1)[1]
算法目标的 decode 语义可以写成:
I_j = Σ_i w_i ReLU(q_i k_j^T) [B,T]
J = TopK(I,k) [B,k]
a_j = (qC' cKV_j^T + qR kR_j^T) / sqrt(d), j∈J [B,H,k]
p = Softmax(a, axis=selected_history) [B,H,k]
u = Σ_(j∈J) p_j cKV_j [B,H,Rkv]
y = WO(Concat_h(u WUV_h)) [B,D]
其中 qC'=qC WUK 是 memory-optimal decode 中的代数吸收对象,不是另一份 cache;WUK/WUV 都来自联合 KV up-projection 的 K/V 分区。固定 Python reference 为验证语义,实际先形成 dense 主 score 再用 topk_indices scatter mask;论文目标的 production sparse kernel 则应先 gather 选中 latent,再做主 MLA。6
训练时也会形成整段 Indexer K 与 MLA latent,但它们是当前 batch 的 activation,而不是跨请求持久 cache。若沿序列切分,top-k 不能只在本 rank 决定:每个 rank 的 local top-k 只是候选,必须再合并成 global top-k,主 MLA 才能读取正确位置。
GDN:固定状态替代逐 token cache¶
GDN 把历史写入每个 head 的矩阵状态 S_t∈R^(K×V),不保存每个 token 的 K/V。它先衰减旧状态,再计算旧状态对当前 key 的 value 预测,最后只写入预测误差。8
| 对象 | 公式 / 源码名 | Shape | 生产者 → 消费者 | 生命周期 |
|---|---|---|---|---|
| Q/K/V | q_t,k_t,v_t |
[B,H,K]、[B,H,K]、[B,H,V] |
投影、ShortConv、SiLU、Q/K L2Norm → recurrence | 临时 activation |
| 衰减与写入率 | lambda_t,beta_t |
[B,H] |
a_proj/b_proj 与参数化 → state update |
临时 activation |
| 衰减后状态 | Sbar_t |
[B,H,K,V] |
lambda_t S_(t-1) → 预测、写入 |
语义中间量;kernel 可融合 |
| 预测与误差 | vhat_t,e_t |
[B,H,V] |
Sbar_t^T k_t、beta(v-vhat) → delta 写入 |
临时 activation |
| 递归状态 | S_t / h |
[B,H,K,V];cache 可为 [B,H,V,K] |
上一步 → 当前读写 → 下一步 | 固定大小跨 token state |
| 短卷积状态 | conv_state_q/k/v |
三个有界窗口 | ShortConv 上一步 → 下一步 | 固定大小 cache |
| 读出 | r_t,y_t |
[B,H,V] → [B,D] |
q_t^T S_t、Gated RMSNorm、WO |
临时 activation |
完整单 token 递推为:
Sbar_t = lambda_t S_(t-1) [B,H,K,V]
vhat_t = Sbar_t^T k_t [B,H,V]
e_t = beta_t (v_t - vhat_t) [B,H,V]
S_t = Sbar_t + k_t e_t^T [B,H,K,V]
r_t = q_t^T S_t [B,H,V]
y_t = WO(GatedRMSNorm(r_t,z_t)) [B,D]
FLA naive reference 中 h 是 S,b_v 先是 v_t,随后被复用为误差 e_t:10
h = h.clone() * g[:, :, i].exp()[..., None, None]
b_v = v[:, :, i].clone()
b_k = k[:, :, i]
b_v = b_v - (h.clone() * b_k[..., None]).sum(-2)
b_v = b_v * beta[:, :, i][..., None]
h = h.clone() + b_k.unsqueeze(-1) * b_v.unsqueeze(-2)
o[:, :, i] = torch.einsum("bhd,bhdm->bhm", q[:, :, i], h)
训练为了并行,不会真的用 Python 按 token 慢循环。FLA 使用 chunk recurrence 或 fused kernel,并在需要时输出 final state;短序列 decode 则走 recurrent kernel。两条路径实现同一递推,但物化的中间量和 workspace 不同。9
固定大小不等于无损
MHA/DSA 仍为每个历史 token 保存独立表示;GDN 让多个 token 竞争同一状态矩阵容量。它用固定内存换取了历史压缩归纳偏置,不应只凭 cache 大小判断模型质量。
CP:训练时交换临时 K/V 分块¶
Megatron Context Parallelism 沿 sequence 维切分网络输入和所有 activation。Linear、RMSNorm 等不跨 token,可直接在本地 S/P 个位置上执行;Attention 的本地 Query 仍必须看到全局 K/V。官方语义是前向收集 K/V,反向对 K/V activation gradient 做 reduce-scatter,高性能路径以 P2P Ring 实现。2
设 Rank r 初始拥有:
Ring 路径固定 Q_r,让 (K_j,V_j) 依次经过所有 rank。每轮用 online softmax 合并局部分数,避免物化完整 [S/P,S] 分数矩阵:
def cp_ring_forward(q_local, k_local, v_local, cp_group):
running_max = full([B, H, S_local, 1], -inf)
running_sum = zeros([B, H, S_local, 1])
running_out = zeros([B, H, S_local, head_dim])
k_block, v_block = k_local, v_local
owner = cp_group.rank
for round_id in range(cp_group.size):
scores = matmul(q_local, transpose(k_block, -1, -2)) * scale
scores = apply_global_causal_mask(scores, q_owner=cp_group.rank, kv_owner=owner)
block_max = scores.max(dim=-1, keepdim=True)
new_max = maximum(running_max, block_max)
old_factor = exp(running_max - new_max)
block_exp = exp(scores - new_max)
running_out = running_out * old_factor + matmul(block_exp, v_block)
running_sum = running_sum * old_factor + block_exp.sum(dim=-1, keepdim=True)
running_max = new_max
k_block, v_block, owner = ring_send_recv(k_block, v_block, owner, cp_group)
return running_out / running_sum
这里 K_j,V_j 是当前训练 forward 的临时 activation block。反向还需把各轮产生的 dK_j,dV_j 送回 owner 并累加;它们不会像推理 cache 那样在下一次 forward 继续存在。
对于 MQA/GQA,Hkv 更小,Ring 发送的 K/V 体积也更小,NVIDIA 文档明确把它列为降低 CP 通信量的手段。2 但这不改变“本地 Q 要覆盖全局 token”的语义。
USP:先换轴,再沿 Ring 扩展上下文¶
USP 是 Ulysses 与 Ring 的二维组合。令总 sequence-parallel degree 为:
初始每个 rank 拥有 [B,S/(PuPr),H,D]。Ulysses group 中的 All-to-All 把一部分 sequence 分片换成 head 分片:
随后 Ring group 在 Pr 个 rank 之间轮转 K/V,最后 inverse A2A 把输出恢复到初始 layout。YunChang 的公开实现正是按照 process group 初始化、local shard 提取、LongContextAttention 的顺序组织,并要求 world_size=Pu×Pr。4
def usp_attention(q_local, k_local, v_local, ulysses_group, ring_group):
q_head = all_to_all(q_local, scatter_axis="head", gather_axis="sequence",
group=ulysses_group)
k_head = all_to_all(k_local, scatter_axis="head", gather_axis="sequence",
group=ulysses_group)
v_head = all_to_all(v_local, scatter_axis="head", gather_axis="sequence",
group=ulysses_group)
out_head = ring_online_attention(q_head, k_head, v_head, group=ring_group)
out_local = all_to_all(out_head, scatter_axis="sequence", gather_axis="head",
group=ulysses_group)
return out_local
纯 Ulysses 的并行度受可切分 head 数限制;Ring 不切 head,但增加 P2P 轮次。USP 只让 Pu 消耗 head 维,Pr 可继续扩展,因此适合把节点内高带宽交给 A2A、节点间交给 Ring。论文在特定 2×8 A800、LLaMA3-8B 配置下报告 208K 序列的 47% MFU;这只是该硬件和配置的证据,不是普遍性能保证。3
训练切分与推理 cache 分片¶
MHA decode:需要 distributed softmax¶
训练 CP/USP 的输入有长 sequence 维,可以把 Q/K/V activation 一起切开。decode 通常只有 q_t 一个 Query;若把持久 K/V cache 按历史 token 分给 P 个 rank,每个 rank 只能算局部 score,必须按全局 Softmax 合并。
Rank r 对本地 cache 得到局部最大值 m_r、局部分母 l_r 与未归一化输出 u_r:
全局结果是:
所以“每卡只放 T/P cache”并不免费:还要复制或切分新 Query、归并 max/sum/output,并处理 PagedAttention 的 block table、请求迁移和负载均衡。只有 cache 容量或带宽收益超过通信成本时才值得。
DSA:先全局选,再稀疏读¶
DSA 的 cache 若按历史 token 分片,每个 rank 可先为本地 Indexer K 算 local_topk,但随后必须:
- 合并所有 rank 的候选分数与全局位置,得到真正的
global_topk。 - 按 owner gather 选中的
cKV/kR,或把当前 Query 发到 owner 上完成局部主 MLA。 - 对选中位置的主 attention 做全局 Softmax 合并。
若 USP 的 Ulysses 维还切分了 Indexer heads,I_j=Σ_i w_i ReLU(q_i k_j^T) 的跨 head 求和也要先 reduce,才能做 top-k。标准 MHA 的 A2A/Ring wrapper 因此不能不加修改地保证 DSA 语义。
GDN:传状态,不传 KV¶
GDN 没有历史 K/V block 可供 Ring 轮转。若训练序列按连续 chunk 分在多个 rank,后一个 chunk 的第一个 token 需要前一个 chunk 的结束状态:
最直观实现是按 causal 顺序传 boundary state;更并行的实现可以让每个 chunk 形成“输入状态 → 输出状态”的 transition summary,再做 prefix scan。反向传播依赖方向相反。若用 Ulysses 只按 head 切分,GDN 的 head-wise 状态相互独立,通常更自然;若再增加 Ring 序列维,就要把 KV Ring 改成状态传递或 scan。
这三段是从计算依赖推出的实现条件,不表示 Megatron CP 或 YunChang 已经为任意 DSA/GDN 模型提供开箱即用的训练与 serving 支持。
联合比较¶
| 场景 | MHA | DSA | GDN |
|---|---|---|---|
| 训练历史对象 | 全序列 Q/K/V activation | Indexer K + MLA latent activation | chunk activation + boundary state |
| 推理持久对象 | K_cache,V_cache |
Indexer K/scale + cKV/kR |
S_t + 短卷积 state |
随上下文 T 增长 |
是,O(T Hkv d) |
是,两套压缩历史 | 否,相对 T 为固定大小 |
| Decode 主读取 | 全历史 K/V | Indexer 全历史;主 MLA top-k | 当前固定状态 |
| CP 序列切分 | KV block 全局交互 | global top-k + selected latent | boundary state / prefix scan |
| USP 的 Ulysses 维 | 切 heads 后本地 attention | Indexer 跨 head score 可能需 reduce | head-wise state 较自然 |
| USP 的 Ring 维 | KV block 轮转 | Indexer/latent 分片与全局 top-k | 不轮转 KV,改传 state |
显存峰值应分项记录:
M_peak =
M_model_state
+ M_saved_activation
+ M_persistent_cache_or_state
+ max_t(M_operator_workspace(t) + M_comm_buffer(t))
+ M_allocator_margin
不能只把某个对象除以 P 就声称总显存严格缩小 P 倍。All-to-All 接收 buffer、Ring 双缓冲、FlashAttention workspace、GDN chunk 中间量和 allocator 碎片都可能形成新的峰值。
常见误区¶
- “训练也有 K/V,所以就是 KV cache”:名字相同,生命周期不同。训练 K/V 属于本轮计算图;推理 cache 跨 decode step 存活。
- “用了 KV cache,decode 就是 O(1)”:省掉的是历史 K/V 重算;dense MHA 的新 Query 仍读取
T个历史位置,单步 attention 仍随T线性。 - “DSA 只保存 top-k”:错误。为了下一个 Query 能重新选择,Indexer 与 MLA 历史表示仍要保存;top-k 是每个 Query 的读取集合。
- “GDN 是另一种压缩 KV cache”:更准确地说它是递归状态。旧 token 不再拥有可逐项寻址的独立 K/V。
- “CP 就是每卡只看局部上下文”:每卡只保存局部 activation,但本地 Query 的数学结果仍须包含全局允许的 K/V。
- “训练 USP 配置可直接照搬到 decode”:decode 的 Query sequence 长度通常为 1,持久 cache 的放置、请求调度和归并通信需要独立设计。
- “USP 总并行度只要不超过 Q heads 就行”:真正消耗 head 维的是
Pu,而 GQA/MQA、Indexer 和 TP 共存时还要核对 KV/head 映射与跨 head reduce。
总结¶
理解 cache 的最好方法不是背缩写,而是对每个对象问四遍:
- 谁产生它?
- 谁在什么时候再次读取它?
- 它活到当前 forward、当前 chunk,还是下一次 decode step?
- 沿 sequence 或 head 切分后,哪个归约才能保持与单卡相同的数学结果?
沿这四个问题看,MHA 是逐 token 历史表,DSA 是“全历史轻索引 + top-k 主读取”,GDN 是固定状态机;CP/USP 则是训练时重排全局交互的办法。它们可以组合,但绝不能仅凭都出现了 K、V、state 或 sequence shard,就认为保存与通信语义相同。
参考资料¶
-
Hugging Face, Transformers Caching, accessed 2026-07-23. ↩↩
-
NVIDIA, Megatron Core Context Parallel Package, accessed 2026-07-23. ↩↩
-
Fang et al., USP: A Unified Sequence Parallelism Approach for Long Context Generative AI, first submitted 2024-05-13. ↩
-
YunChang, hybrid Attention implementation, commit
56118e0d. ↩ -
DeepSeek-AI, DeepSeek-V3.2-Exp, released 2025-09-29. ↩
-
DeepSeek-AI,
inference/model.py, commit87e509a2. ↩ -
DeepSeek-AI,
inference/kernel.py, commit87e509a2. ↩ -
Yang et al., Gated Delta Networks: Improving Mamba2 with Delta Rule, first submitted 2024-12-09. ↩
-
FLA,
fla/layers/gated_deltanet.py, commitc70f11c5. ↩ -
FLA,
fla/ops/gated_delta_rule/naive.py, commitc70f11c5. ↩







