对应论文 §2.3 "Compressed Sparse Attention 2 (CSA2)" 的开头和 §2.3.1 "Cross-Layer KV and Index Reuse",层的具体排布在 §4.2.1。CSA2 是在 V4 的 CSA 上改的。压缩算子怎么推出来、lightning indexer 怎么打分,在 V4 连载的第 1 篇 和第 2 篇 里,这里只摆出结论,篇幅留给改动。
先用七步看 CSA2 是怎么从 V4 的 CSA 变过来的,每一步只改一处。看完再往下读,每一节对应其中一步。
encoder 的一组:第 2 – 7 层(共 3 组) 第 2 层 Full main KV indexer K 打分,选 512 条 注意力 第 3 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 4 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 5 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 6 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 7 层 Reuse main KV indexer K 打分,选 512 条 注意力 decoder:第 20 – 39 层(画前 8 层) 第 20 层 Full main KV indexer K 打分,选 512 条 注意力 第 21 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 22 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 23 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 24 层 Reindex main KV indexer K 打分,选 512 条 注意力 第 25 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 26 层 Reuse main KV indexer K 打分,选 512 条 注意力 第 27 层 Reuse main KV indexer K 打分,选 512 条 注意力 … 第 28、32、36 层是 Reindex,其余是 Reuse,到第 39 层 候选池 → 第 24 层 全局 KV:每层各一份 V4-Flash:5.41 条 / token V4 的做法:每层各做一遍 一组 6 层,每层自己生成 main KV 和 indexer K,自己打分选 512 条,再做注意力。V4-Flash 带全局注意力的 41 层各存一份。
▶ 播放 ‹ 0 V4 的做法:每层各做一遍 1 共用 main KV 2 indexer K 跟着共用 3 Reuse:选出的 512 条也共用 4 decoder:一份 KV,每 4 层重选一次 5 候选池:Reindex 只在 16384 个位置里打分 6 合起来:890 B / token › CSA2 的演进。左边是 encoder 的一组 6 层,右边是 decoder 的前 8 层。可以点步骤按钮,也可以按键盘左右键。
图像 文本 token 视觉特征写到图像 token 的位置上, 和文本 embedding 排成同一条序列 复制成 4 份,进入 4 条残差流 残差流 ×4 每条 5120 维 重复 3 组 第 2 – 19 层,6 层一组 每组 1 层 Full + 5 层 Reuse 重复 4 组 第 24 – 39 层,4 层一组 每组 1 层 Reindex + 3 层 Reuse encoder:第 0 – 19 层 decoder:第 20 – 39 层 prefill 时,绝大部分 prompt token 只算到这里 decoder 的全局 KV 全部由 投影得到 只有滑窗分支,不读全局 KV 40 层的注意力后面都接一个 MoE, 下面各行省略不画 写 读 写 读 重选索引 读 encoder 的三组各有一份全局 KV,由该组的 Full 层写入,组内 6 层共用 encoder 共享池(每组一份) main KV:每 2 个 token 一条,512 维 FP4 indexer K:每条 128 维 FP4 top-512 索引:每个 query 一份,不进缓存 decoder 只有一份全局 KV,由第 20 层从 encoder 末态投影出来,20 层共用 decoder 共享池(只有一份) main KV:每个 token 一条,512 维 FP4 indexer K:每条 128 维 FP4 候选池:第 20 层选出的 2048 块 × 8 = 16384 个位置, Reindex 层只在池内打分 top-512 索引:第 20 层先写, 每个 Reindex 层覆盖一次 全局 KV 合计 890 B / token Single-Pass mHC:每个子层一次读、一次写 + 混 读用的 是上一个子层算好的 查表结果经门控后加进残差流 logits 最后一次只读:4 条流压成 1 条 读主干第 37 – 39 层入口处 4 条流的平均 一次前向出 5 个草稿 token 和各自的置信度 encoder 末态:第 19 层的输出,也就是第 20 层的输入。decoder 所有层的全局 KV 都只从它投影(论文式 1) encoder 末态 mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好 读 A mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B mHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合 写 C · 混 B 从头训练的视觉编码器:2D-RoPE、RMSNorm、SwiGLU,hidden 1024 DeepSeek-ViT 32 层 · patch 14 3×3 pixel-unshuffle 把 9 个相邻 patch 拼到通道维,再过两层 MLP 投影到主干宽度 3×3 重排 + MLP token ÷ 9 → 5120 维 图像位置上的 embedding 被视觉特征覆盖 Token Embedding 129280 × 5120 前两层只有滑窗分支,没有全局 KV 滑窗注意力 第 0、1 层 · 窗口 128 sqrt(softplus) 打分;文本和图像 token 各用一套负载均衡偏置;routed expert 权重 FP4 DeepSeekMoE 每层的 FFN · 384 选 6 + 1 shared Full 模式:自己压缩 main KV(每 2 个 token 一条),投影出 indexer K,打分选 top-512,全部写进共享池 CSA2 · Full 第 2 / 8 / 14 层 · m = 2 Reuse 模式:只有自己的 query、滑窗 KV 和输出投影;全局 KV 与 top-512 索引都从共享池读 CSA2 · Reuse ×5 每组其余 5 层 decoder 唯一的 Full 层:输入是 encoder 末态,每个 token 一条 main KV;另选出最多 16384 个位置作为候选池 CSA2 · Full 第 20 层 · m = 1 · 建候选池 Reuse 模式 CSA2 · Reuse ×3 第 21 – 23 层 Reindex 模式:全局 KV 和 indexer K 从共享池读,用自己的 indexer query 在候选池里重新打分,选出新的 top-512 CSA2 · Reindex 第 24 / 28 / 32 / 36 层 Reuse 模式:用本组 Reindex 层选出的 top-512 CSA2 · Reuse ×3 每组其余 3 层 最终归一化 RMSNorm 不与 embedding 共享权重 LM Head → 129280 投机解码的草稿模块:三个 block 一次前向给出 5 个草稿 token 的 logits,Markov head 补上草稿 token 之间的依赖,confidence head 估计每个位置被接受的概率 DSpark 草稿层 ×3 滑窗注意力 + MoE(128 选 3) 按 2、3、4-gram 的哈希查表,取出的向量经门控后加进 4 条残差流;两张表共 196B 参数 Engram 第 1、14 层入口各一个 序列 · CSA2 序列 · 共享池 序列 · CED 序列 · 滑窗 深度 · Single-Pass mHC 宽度 · DeepSeekMoE 记忆 · Engram 解码 · DSpark 输入与输出
你现在在这里:三种模式的注意力层,和它们读写的两个共享池。
KV cache 的大小是三个量相乘
上一篇 解决了 prefill 的计算量。这一篇解决存储:全局 KV 随上下文长度线性增长,1M 上下文下它决定一台机器能放多少个会话。
每个 token 要存多少字节的全局 KV,可以写成对层求和:
字节 / token = ∑ 存 KV 的层 l b l ⏟ 每条多少字节 × 1 m l ⏟ 每个 token 几条 . \text{字节 / token} = \sum_{\text{存 KV 的层 } l}\ \underbrace{b_l}_{\text{每条多少字节}} \times \underbrace{\frac{1}{m_l}}_{\text{每个 token 几条}} . 字节 / token = 存 KV 的层 l ∑ 每条多少字节 b l × 每个 token 几条 m l 1 .
三个量可以分别压,效果相乘。论文 §2.3 就是按这三个维度回顾前人工作的:
每条多大。 GQA 减少 KV 的头数。MLA 让所有头共用一条低维的 latent。V4 的条目是一条 512 维向量,K 和 V 是同一个东西。
每个 token 几条。 V4 的 CSA 把每 m = 4 m = 4 m = 4 个 token 压成一条,HCA 把每 128 个压成一条。
多少层各存一份。 V4 没有动这个维度:V4-Flash 有 41 层带全局注意力,每层自己存一份。
第三个维度上已有的工作,论文点了四篇。把它们各自共用了什么列出来:
工作 跨层共用的是什么 不足 Cross-Layer Attention(2024) KV。相邻几层用同一份 论文没有评价,只把它当作「有些层复用别的层的缓存」的出处。我补一句:它是稠密注意力上的纯 KV 共享,没有稀疏选择 IndexCache(2026) 稀疏选择的结果。只有一部分层保留 indexer,其余层用最近一次选出的条目 每层的 KV 仍然各存各的,存储没有减少 YOIO(2026) KV 和选择结果。上半层共用一份 KV,全体只选一次 整个网络共用一次选择,限制了效果 HySparse(2026) KV 和块级的选择结果。来源是一个稠密注意力层 仍然保留了稠密注意力的层
后三行的不足是论文 §2.3 的原话。论文的结论是:这些方法没有一个同时覆盖三个维度。CSA2 要做的是在 V4 已经压过前两个维度的基础上,把第三个也压掉。
V4-Flash 平均每个 token 存 21 × 1 4 + 20 × 1 128 ≈ 5.41 21 \times \frac{1}{4} + 20 \times \frac{1}{128} \approx 5.41 21 × 4 1 + 20 × 128 1 ≈ 5.41 条。V4.1-Flash 是 2.5 条。这一篇讲 2.5 是怎么来的。每条多少字节留到下一篇。
一层全局注意力算四样东西,其中三样可以跨层共用
要跨层共用,先看一层里有什么。V4 的一个 CSA 层,全局分支按顺序算四样东西:
main KV。 把本层的隐藏态投影成 512 维,每 m m m 个 token 压成一条。这是注意力最后读的内容,要存进 KV cache。
indexer K。 每条 main KV 配一个 128 维的索引用的 key,也要存进 KV cache。
top-k 索引。 indexer 用当前 token 的 indexer query 和所有 indexer K 打分,选出分数最高的 k k k 条。这是一组位置编号,每个 query 现算,不进缓存。
注意力本身。 用本层的 query 读选中的 k k k 条 main KV,加上本层滑窗里的 128 条。
前两样是存储,第三样是计算,第四样是这一层真正的工作。
CSA2 的规定是:第四样每层自己做,前三样可以用前面某一层的。而且「共用缓存」和「共用选择结果」分开决定。把前三样的各种组合列出来,有意义的只有三种:
模式 main KV indexer K top-k 索引 query、滑窗 KV、输出投影 Full 自己生成 自己生成 自己选 自己的 Reindex 用前面的 用前面的 自己选 自己的 Reuse 用前面的 不需要 用前面的 自己的
「自己生成 KV、但用别人的索引」这种组合没有意义:索引是针对某一份 KV 的位置编号,换了一份 KV,编号指的就不是原来的条目了。
每一层属于哪种模式是静态指定的,写在 config 里,不随输入变。
Full 等于 V4 的一个完整的 CSA 层。
Reindex 不占新的存储,但要付一次 indexer 的计算。换来的是它能选出和前面不同的 512 条。
Reuse 既不占存储,也不跑 indexer。它只做第四样。
Full:层 2/8/14(m=2)、层 20(m=1,无 wgate) Reindex:层 24/28/32/36 Reuse:30 层 跨层共享池 SharedAttentionRuntime(356 B/条 × 2.5 条/token = 890 B/token) [64,512] RoPE 前 latent index_k top-512 写 compress_kv 写 index_k 写 topk_idxs 层 20:写 candidates 读 compress_kv 、 新 top-512 读 index_k 读 candidates 写回 topk_idxs(覆盖) 读 compress_kv 读 compress_kv 读 topk_idxs(最新) [B, N, 5120] mHC + norm 后 Q 低秩(每层自有) wq_a 5120→1280 + q_norm → wq_b 1280→64×512 RoPE 末 64 维 [64, 512] indexer Q(每层自有) wq_b 1280→32×128 weights_proj 5120→32 FP4 SWA 路径(每层自有) wkv 5120→512 + kv_norm RoPE 末 64 → FP8 ring cache 128 槽 sparse_attn(每层自有) SWA 128 + 全局 top-512 + sink → 反旋转 RoPE Compressor wkv、wgate 5120→512 softmax(gate)⊙kv 按 m 求和(不重叠) m=1 无 wgate(层 20) latent [N/m, 512] RoPE 末 64 维 → FP4 E2M1 + E4M3 scale/16 indexer K wk 512→128 + k_norm RoPE θ=160000 → FP4 E8M0/32,无 Hadamard 打分 因果掩码 → top-512 按位置 sort 层 20 另产 candidates (2048 块 × 8 = 16384) 输出投影(每层自有) wo_a 8 组 ×(4096→1024) wo_b 8192→5120 Q 低秩(每层自有) wq_a 5120→1280 + q_norm → wq_b 1280→64×512 RoPE 末 64 维 [64, 512] 重打分 读共享 index_k candidates 掩码 → top-512,sort 写回 topk_idxs(覆盖) sparse_attn(每层自有) SWA 128 + 全局 top-512 + sink → 反旋转 RoPE SWA 路径(每层自有) wkv 5120→512 + kv_norm RoPE 末 64 → FP8 ring cache 128 槽 indexer Q(每层自有) wq_b 1280→32×128 weights_proj 5120→32 FP4;K 从共享池读 本层不产 KV / index_k 无 Compressor 无 wk / k_norm 输出投影(每层自有) wo_a 8 组 ×(4096→1024) wo_b 8192→5120 Q 低秩(每层自有) wq_a 5120→1280 + q_norm → wq_b 1280→64×512 RoPE 末 64 维 [64, 512] 无 indexer、无 Compressor 注意力参数清单 与纯 SWA 层相同 SWA 路径(每层自有) wkv 5120→512 + kv_norm RoPE 末 64 → FP8 ring cache 128 槽 sparse_attn(每层自有) 读共享 compress_kv + 共享 topk_idxs + SWA 128 + sink → 反旋转 输出投影(每层自有) wo_a 8 组 ×(4096→1024) wo_b 8192→5120 compress_kv 主 KV latent,FP4(288 B/条) index_k indexer K,FP4(68 B/条) topk_idxs 最新 top-512 索引 candidates(仅 decoder) ≤ 2048 块 × 8 = 16384 层 20 产出 三种模式各自算什么、复用什么:Full 自己压 main KV(softmax(gate)⊙kv 按 m 个不重叠 token 求和、无 APE)、从 latent 投出 indexer K(无 Hadamard)、跑 indexer 产 top-512(层 20 另产 candidates 候选池),统统写进跨层共享池;Reindex 只带自己的 indexer Q,读共享 index_k 重打分(先经 candidates 掩码),把新 topk_idxs 写回;Reuse 只有自己的 Q / SWA / O,直接读共享 compress_kv + topk_idxs 进 sparse_attn。实线是层内数据流与写池,虚线是从池里读。
上面这张图在注意力结构对比 那篇里出现过,它把三种模式各自的算子画全了。读的时候注意每一列都有的那几个方块:query 的低秩投影、滑窗 KV、注意力和输出投影。它们标着「每层自有」,三种模式下都一样。
40 层只有 4 个 Full
encoder(第 0 – 19 层) decoder(第 20 – 39 层) 第 0 层:纯滑窗注意力,窗口 128 第 1 层:纯滑窗注意力,窗口 128;入口处有一个 Engram 第 2 层:Full,m = 2:自己生成 main KV 和 indexer K,选 top-512 第 3 层:Reuse:读共享的 main KV 和 top-512 索引 第 4 层:Reuse:读共享的 main KV 和 top-512 索引 第 5 层:Reuse:读共享的 main KV 和 top-512 索引 第 6 层:Reuse:读共享的 main KV 和 top-512 索引 第 7 层:Reuse:读共享的 main KV 和 top-512 索引 第 8 层:Full,m = 2:自己生成 main KV 和 indexer K,选 top-512 第 9 层:Reuse:读共享的 main KV 和 top-512 索引 第 10 层:Reuse:读共享的 main KV 和 top-512 索引 第 11 层:Reuse:读共享的 main KV 和 top-512 索引 第 12 层:Reuse:读共享的 main KV 和 top-512 索引 第 13 层:Reuse:读共享的 main KV 和 top-512 索引 第 14 层:Full,m = 2:自己生成 main KV 和 indexer K,选 top-512;入口处有一个 Engram 第 15 层:Reuse:读共享的 main KV 和 top-512 索引 第 16 层:Reuse:读共享的 main KV 和 top-512 索引 第 17 层:Reuse:读共享的 main KV 和 top-512 索引 第 18 层:Reuse:读共享的 main KV 和 top-512 索引 第 19 层:Reuse:读共享的 main KV 和 top-512 索引 第 20 层:Full,m = 1:自己生成 main KV 和 indexer K,选 top-512,并建候选池 第 21 层:Reuse:读共享的 main KV 和 top-512 索引 第 22 层:Reuse:读共享的 main KV 和 top-512 索引 第 23 层:Reuse:读共享的 main KV 和 top-512 索引 第 24 层:Reindex:读共享的 indexer K,在候选池里重新选 top-512 第 25 层:Reuse:读共享的 main KV 和 top-512 索引 第 26 层:Reuse:读共享的 main KV 和 top-512 索引 第 27 层:Reuse:读共享的 main KV 和 top-512 索引 第 28 层:Reindex:读共享的 indexer K,在候选池里重新选 top-512 第 29 层:Reuse:读共享的 main KV 和 top-512 索引 第 30 层:Reuse:读共享的 main KV 和 top-512 索引 第 31 层:Reuse:读共享的 main KV 和 top-512 索引 第 32 层:Reindex:读共享的 indexer K,在候选池里重新选 top-512 第 33 层:Reuse:读共享的 main KV 和 top-512 索引 第 34 层:Reuse:读共享的 main KV 和 top-512 索引 第 35 层:Reuse:读共享的 main KV 和 top-512 索引 第 36 层:Reindex:读共享的 indexer K,在候选池里重新选 top-512 第 37 层:Reuse:读共享的 main KV 和 top-512 索引 第 38 层:Reuse:读共享的 main KV 和 top-512 索引 第 39 层:Reuse:读共享的 main KV 和 top-512 索引 第 40 层:DSpark 草稿层:纯滑窗注意力 + MoE(128 选 3) 第 41 层:DSpark 草稿层:纯滑窗注意力 + MoE(128 选 3) 第 42 层:DSpark 草稿层:纯滑窗注意力 + MoE(128 选 3) 0 2 8 14 20 24 28 32 36 39 40 – 42 3 组,每组共用一份全局 KV(m = 2) 20 层共用一份全局 KV(m = 1) DSpark Full ×4:自己生成全局 KV(第 2、8、14、20 层) Reindex ×4:只重新选 top-512(第 24、28、32、36 层) Reuse ×30:KV 和索引都用现成的 纯滑窗 ×2(第 0、1 层) DSpark 草稿层 ×3 Engram:第 1、14 层入口 DSpark 读第 37 – 39 层入口的残差流 格子的颜色是模式。下方的括号标出共用同一份全局 KV 的范围。
排布来自论文 §4.2.1,和 config 的 kv_source_layer_ids、index_source_layer_ids 一致:
范围 分组 每组的排法 m m m 第 0、1 层 — 纯滑窗,没有全局分支 — encoder 第 2 – 19 层 3 组,每组 6 层 1 层 Full + 5 层 Reuse 2 decoder 第 20 – 23 层 1 组 1 层 Full + 3 层 Reuse 1 decoder 第 24 – 39 层 4 组,每组 4 层 1 层 Reindex + 3 层 Reuse 1
带全局分支的 38 层里,Full 4 层,Reindex 4 层,Reuse 30 层。
每个 token 存几条全局 KV,只看 Full 层:
3 × 1 2 ⏟ encoder 的 3 个 Full + 1 × 1 1 ⏟ decoder 的 1 个 Full = 2.5 条 / token . \underbrace{3 \times \frac{1}{2}}_{\text{encoder 的 3 个 Full}} + \underbrace{1 \times \frac{1}{1}}_{\text{decoder 的 1 个 Full}} = 2.5\ \text{条 / token}. encoder 的 3 个 Full 3 × 2 1 + decoder 的 1 个 Full 1 × 1 1 = 2.5 条 / token .
V4-Flash 是 5.41 条。条数减少到不足一半,同时每一条覆盖的范围更细:V4 最细是 4 个 token 一条,V4.1 的 decoder 是 1 个 token 一条。层这个维度省出来的空间,一部分还给了序列这个维度。
decoder 只有一个 Full 层,这和上一篇的 CED 对上了:decoder 的全局 KV 必须由 H 20 H_{20} H 20 投影,第 20 层的输入正好就是 H 20 H_{20} H 20 。
论文没有解释为什么 encoder 是 6 层一组、decoder 是 4 层一组,也没有解释为什么 encoder 用 3 个 Full 而不用 Reindex。能确定的只有存储上的代价:encoder 每多一个 Full 层,每 token 多 0.5 条;decoder 每多一个,多 1 条。
压缩算子去掉了重叠和位置偏置
Full 层要自己生成 main KV。先把 V4 的压缩算子摆出来。记 C t C_t C t 是 token t t t 的 KV 向量,Z t Z_t Z t 是它的压缩权重,两者都是 512 维,由隐藏态各经一个线性层得到。V4 的 CSA 让相邻两个条目共用一半的 token:
C i Comp = ∑ j = m ( i − 1 ) m i − 1 S j a ⊙ C j a + ∑ j = m i m ( i + 1 ) − 1 S j b ⊙ C j b , C^{\text{Comp}}_i = \sum_{j=m(i-1)}^{mi-1} S^a_j \odot C^a_j + \sum_{j=mi}^{m(i+1)-1} S^b_j \odot C^b_j , C i Comp = j = m ( i − 1 ) ∑ mi − 1 S j a ⊙ C j a + j = mi ∑ m ( i + 1 ) − 1 S j b ⊙ C j b ,
其中 S S S 是对 2 m 2m 2 m 个槽的 Z Z Z 加上一个可学习的块内位置偏置 B B B 之后做的 softmax。
token 0 1 2 3 4 5 6 7 8 9 10 11 块 0( 个 token) 块 1( 个 token) 块 2( 个 token) a, 位 0 b, 位 0 a, 位 1 b, 位 1 a, 位 2 b, 位 2 a, 位 3 b, 位 3 a, 位 0 b, 位 0 a, 位 1 b, 位 1 a, 位 2 b, 位 2 a, 位 3 b, 位 3 a, 位 0 b, 位 0 a, 位 1 b, 位 1 a, 位 2 b, 位 2 a, 位 3 b, 位 3 → 给条目 3 只有块 0 的 b(a 槽填 ) 块 1 的 b + 块 0 的 a 块 2 的 b + 块 1 的 a 压缩条目 个 每个条目:把 个槽的 做一次 softmax(逐通道, 个通道各一套权重),用它给 加权求和。每个 token 参与两个条目。 CSA 的重叠压缩(公式 11、12)。橙色是 a 系列,供下一个 条目用;蓝色是 b 系列,供当前 条目用。条目数仍是 n/m,但每个条目看到 2m 个 token。
这个设计有三样成本:每个 token 要算两套 C C C 和 Z Z Z ,共 4 个投影矩阵;有一个位置偏置;decode 时要多缓冲一个块。
CSA2 把它退回最简单的形式。论文只用文字描述,下面的式子是我按官方代码写出来的:
[ S m i ; … ; S m ( i + 1 ) − 1 ] = Softmax ( [ Z m i ; … ; Z m ( i + 1 ) − 1 ] ) , C i Comp = ∑ j = m i m ( i + 1 ) − 1 S j ⊙ C j . [S_{mi};\ \dots;\ S_{m(i+1)-1}] = \operatorname{Softmax}\big([Z_{mi};\ \dots;\ Z_{m(i+1)-1}]\big), \qquad
C^{\text{Comp}}_i = \sum_{j=mi}^{m(i+1)-1} S_j \odot C_j . [ S mi ; … ; S m ( i + 1 ) − 1 ] = Softmax ( [ Z mi ; … ; Z m ( i + 1 ) − 1 ] ) , C i Comp = j = mi ∑ m ( i + 1 ) − 1 S j ⊙ C j .
每个条目只由自己这一组的 m m m 个 token 得到,相邻条目不共用 token。softmax 仍然是逐通道的:512 个通道各有一套「组里哪个 token 更重要」的权重。没有位置偏置 B B B 。
encoder:m = 2,每 2 个 token 一条 decoder:m = 1,每个 token 一条 条目 0 条目 1 条目 2 ,512 个通道各做一次 条目 0 条目 1 条目 2 条目 3 只有一项时 softmax 的权重恒为 1,不需要 Z CSA2 的压缩算子。左边是 encoder 的 Full 层,右边是 decoder 的第 20 层。
decoder 的 m = 1 m = 1 m = 1 是它的特例。一组只有一个 token,softmax 只有一项,权重恒为 1:
C t Comp = C t . C^{\text{Comp}}_t = C_t . C t Comp = C t .
压缩权重 Z Z Z 不起任何作用。官方代码对这个情况单独处理,连 Z Z Z 的投影矩阵都不建:
def forward ( self , x , start_pos ):
ratio = self . compress_ratio
if ratio == 1 : # one token per group: nothing to pool, so no gate and no fp32
return self . norm ( self . wkv ( x ))
x = x . float ()
kv , score = self . wkv ( x ) , self . wgate ( x )
...
kv = kv . unflatten ( 1 , ( - 1 , ratio ) )
score = score . unflatten ( 1 , ( - 1 , ratio ) )
kv = ( kv * score . softmax ( dim = 2 ) ). sum ( dim = 2 )
...
return self . norm ( kv . to ( dtype ))
wkv 是 C C C 的投影,wgate 是 Z Z Z 的投影。权重文件里第 2、8、14 层有 compressor.wgate,第 20 层没有。
论文给的理由是这两处简化让实现更简单、训练更快,没有给消融。V4 第 1 篇说过重叠是为了缓解块边界的硬切。这是我的解读:m m m 从 4 降到 2 和 1 之后,一个条目只覆盖一两个 token,边界本身的影响已经小了,重叠的必要性也跟着下降。
压缩之后的处理和 V4 一样:过 RMSNorm,最后 64 维加 RoPE,位置取这一组第一个 token 的位置。
indexer K 改由 main KV 投影
Full 层的第二样东西是 indexer K。V4 的做法是再做一遍压缩:indexer 有自己的一套压缩算子,从隐藏态出发,产生 128 维的 key。它也是 4 个投影矩阵加一个位置偏置。
CSA2 不再从隐藏态另走一路。indexer K 直接由刚算好的 main KV 条目投影:
k i I = RMSNorm ( C i Comp W I K ) , W I K ∈ R 512 × 128 . k^I_i = \operatorname{RMSNorm}\big(C^{\text{Comp}}_i\, W^{IK}\big), \qquad W^{IK} \in \mathbb{R}^{512 \times 128}. k i I = RMSNorm ( C i Comp W I K ) , W I K ∈ R 512 × 128 .
这个式子也是按代码写的。几个细节只在代码里能看到:
投影用的是加 RoPE 之前 的 main KV。投影完,indexer K 的最后 64 维再加自己的 RoPE。
V4 在量化前对 indexer K 做一次 Hadamard 旋转,V4.1 没有这一步。
indexer K 按 FP4 存,每 32 维一个 scale。
W I K W^{IK} W I K 只有 512 × 128 = 65536 512 \times 128 = 65536 512 × 128 = 65536 个参数。V4-Flash 的 indexer 压缩算子是 4 个 4096 × 128 4096 \times 128 4096 × 128 的矩阵,约 210 万。
打分公式没有变,仍然是 V4 第 2 篇的那个:
I t , s = ∑ j = 1 H I w t , j I ⋅ ReLU ( q t , j I ⋅ k s I ) . I_{t,s} = \sum_{j=1}^{H^I} w^I_{t,j} \cdot \operatorname{ReLU}\big(\mathbf{q}^I_{t,j} \cdot \mathbf{k}^I_s\big). I t , s = j = 1 ∑ H I w t , j I ⋅ ReLU ( q t , j I ⋅ k s I ) .
q t , j I \mathbf{q}^I_{t,j} q t , j I 是 token t t t 的第 j j j 个 indexer query,k s I \mathbf{k}^I_s k s I 是第 s s s 条的 indexer K,w t , j I w^I_{t,j} w t , j I 是由 token t t t 决定的每个头的权重。头数 H I H^I H I 从 V4-Flash 的 64 降到 32,每头仍是 128 维。每个 query 选 512 条。
共享池里有四样东西,两样进缓存,两样现算
三种模式靠一个共享的对象传递数据。官方代码里它只有四个槽:
class SharedAttentionRuntime :
def __init__ ( self ):
self . compress_kv = None # main KV
self . index_k = None # indexer K
self . topk_idxs = None # 每个 query 的 top-512 索引
self . candidates = None # 候选池,下一篇讲
每个槽只有一个位置,没有按层编号。这样就够了,因为层是按顺序执行的:一个 Full 层写入之后,直到下一个 Full 层覆盖它之前,中间的层读到的都是这一份。
注意力层的代码里,三种模式的分支只有两个判断:
def _compress_topk_idxs ( self , x , qr , latent , start_pos , offset , compress_len ):
if not self . is_index_source :
return shared_attn . topk_idxs # Reuse:用现成的索引
idxs = self . indexer ( x , qr , latent , start_pos , offset )
shared_attn . topk_idxs = idxs # Full / Reindex:自己选,写回去
return idxs
def _compress_kv ( self , x , qr , start_pos , offset ):
latent = None
if self . is_kv_source : # Full:自己生成 main KV
latent = self . compressor ( x , start_pos )
shared_attn . compress_kv = self . compress_kv_cache
idxs = self . _compress_topk_idxs ( x , qr , latent , start_pos , offset , compress_len )
...
return shared_attn . compress_kv [ : bsz , : compress_len ] , idxs
is_kv_source 为真是 Full。is_kv_source 为假、is_index_source 为真是 Reindex。两个都为假是 Reuse。
0 2 8 14 20 24 28 32 36 39 第 0 层 第 1 层 第 2 层 第 3 层 第 4 层 第 5 层 第 6 层 第 7 层 第 8 层 第 9 层 第 10 层 第 11 层 第 12 层 第 13 层 第 14 层 第 15 层 第 16 层 第 17 层 第 18 层 第 19 层 第 20 层 第 21 层 第 22 层 第 23 层 第 24 层 第 25 层 第 26 层 第 27 层 第 28 层 第 29 层 第 30 层 第 31 层 第 32 层 第 33 层 第 34 层 第 35 层 第 36 层 第 37 层 第 38 层 第 39 层 层的模式 encoder decoder main KV compress_kv 写 读 读 读 读 读 写 读 读 读 读 读 写 读 读 读 读 读 写 读 读 读 读 读 读 读 读 读 读 读 读 读 读 读 读 读 读 读 进 KV cache indexer K index_k 写 写 写 写 读 读 读 读 进 KV cache 候选池 candidates 写 读 读 读 读 每个 query 现算 top-512 索引 topk_idxs 写 读 读 读 读 读 写 读 读 读 读 读 写 读 读 读 读 读 写 读 读 读 写 读 读 读 写 读 读 读 写 读 读 读 写 读 读 读 每个 query 现算 Full Reindex Reuse 纯滑窗 这一层生成并写入 这一层读取 共享池的四样东西在 40 层里各由谁写、谁读。深色是写,浅色是读。
从这张表能读出几件事。
main KV 一共 4 份,各自一直留着。 第 8 层写入时,共享对象里的槽指向了新的一份,但第 2 层那一份没有被删掉。它存在第 2 层自己的缓存里,下一个 token 进来时,第 3 – 7 层还要读它。所以 KV cache 里常驻的是 4 份 main KV 和 4 份 indexer K。
indexer K 只有 Reindex 层去读。 Full 层用自己刚算的。Reuse 层不跑 indexer,用不到它。encoder 里没有 Reindex 层,所以 encoder 的 3 份 indexer K 只被生成它的那一层用。
top-512 索引不进缓存。 它是一个 query 对一份 KV 的选择,每来一个 token 重算一次。Full 或 Reindex 层算出来,后面几个 Reuse 层在处理同一个 token 时接着用,下一个 Full 或 Reindex 层再覆盖。
decoder 的 20 层共用一份 KV,但有 5 组不同的索引。 第 20 层选一次,第 24、28、32、36 层各重选一次。
Reuse 层只共用条目的集合,注意力权重仍由本层的 query 算
一个自然的疑问:第 2 层的 indexer 按第 2 层的 query 选了 512 条,第 3 – 7 层的 query 和它不一样,凭什么能用同一批?
先把共用的范围说清楚。Reuse 层共用的是「看哪 512 条」这个集合。在这 512 条上怎么分配注意力,是每层用自己的 query 算的:
输出 = ∑ s ∈ S t ∪ 滑窗 softmax s ( q t ( l ) ⋅ k s ) v s . \text{输出} = \sum_{s \in \mathcal{S}_t \,\cup\, \text{滑窗}} \operatorname{softmax}_s\big(\mathbf{q}^{(l)}_t \cdot \mathbf{k}_s\big)\, \mathbf{v}_s . 输出 = s ∈ S t ∪ 滑窗 ∑ softmax s ( q t ( l ) ⋅ k s ) v s .
S t \mathcal{S}_t S t 是共用的,q t ( l ) \mathbf{q}^{(l)}_t q t ( l ) 是本层的。两个 Reuse 层可以把注意力集中在这 512 条里完全不同的几条上。
所以真正的假设是:相邻几层想看的条目,大体落在同一个 512 条的集合里。 这个假设有实验依据。IndexCache 在 V3.2 的稀疏注意力上做过实验:一个 30B 的模型去掉 75% 的 indexer,让那些层直接用最近一次选出的条目,质量只有可忽略的下降。IndexCache 的说法是相邻层选出的 top-k 高度相似。这是我的解读:CSA2 的 Reuse 模式沿用的就是这个发现。
Reindex 让 decoder 的 20 层不必只选一次
如果 decoder 只有一个 Full 层,其余 19 层全是 Reuse,结构最简单。但那样整个 decoder 对每个 token 只选一次 512 条,20 层都被限制在这 512 条里。论文评价 YOIO 时说的正是这一点:全网络共用一次选择会限制效果。
多放几个 Full 层能解决,但每个 Full 层要多存一份 KV,每 token 多 356 字节。
Reindex 是第三条路。它不生成 KV,只重新选一次。代价和收益分开看:
存储。 不增加。Reindex 层读的 main KV 和 indexer K 都是第 20 层的。
参数。 只多一个 indexer query 的投影和一个头权重的投影,约 540 万。没有压缩算子,也没有 W I K W^{IK} W I K 。
计算。 要跑一遍 indexer,给所有看得见的条目打分。decoder 的 m = 1 m = 1 m = 1 ,1M 上下文下就是 100 万条。这是 Reindex 的主要开销,下一篇的 Hierarchical Sparse Indexer 专门处理它。
收益。 decoder 每 4 层换一批条目,20 层一共有 5 次选择。
「共用缓存」和「共用选择结果」分开决定,指的就是这个:Reindex 层共用了缓存,没有共用选择结果。
三种模式的参数清单
把每种模式多出来的参数列出来。「公共部分」是每一层都有的,共约 127M。它包括 query 的两个投影矩阵、滑窗 KV 的投影、输出投影的两个矩阵,还有每个头一个的 attention sink。
模式 层 公共部分之外的参数 约多少 纯滑窗 0、1 无 0 Reuse 30 层 无 0 Reindex 24、28、32、36 indexer 的 wq_b(1280 × 4096)和 weights_proj(5120 × 32) 5.4M Full,m = 1 m = 1 m = 1 20 上一行,加 compressor.wkv(5120 × 512)和 indexer 的 wk(512 × 128) 8.1M Full,m = 2 m = 2 m = 2 2、8、14 上一行,加 compressor.wgate(5120 × 512) 10.7M
Reuse 层的注意力参数和第 0 层的纯滑窗注意力一模一样,权重文件里的名字都相同。
全模型和压缩、索引有关的参数合计约 62M。按 V4-Flash 的 config 手算,它的同类参数约 0.5B。这部分参数在 552B 里微不足道,这里列出来只是为了说明:CSA2 的三种模式不是靠增加参数换来的。
训练时共用的东西要跨流水线传递
推理时层按顺序执行,共享池只是一个全局对象。训练时没这么简单。大模型训练用流水线并行 ,40 层被切成几段放在不同的机器上。一个 Reuse 层和它依赖的 Full 层可能不在同一段。
论文 §3.1.2 列了三项为此做的支持:
shadow indexer。 共用的 indexer 参数只有一个归属方,负责优化和存 checkpoint。其他用到它的流水线段各放一个可执行的副本,参数和梯度靠同步保持一致。
扩展流水线传递的内容。 段与段之间原本只传隐藏态。现在还要把下游要用的中间表示和选择结果一起传过去,并保证梯度能传回来。
按 micro-batch 管理共享状态的生命周期。 流水线里同时有多个 micro-batch 在前向、重算和反向。每份共享状态要留到最后一个用它的层算完,再立刻释放。
这三项的细节放在第 8 篇。
另有一处和 V4 不同:V4 先用稠密注意力训练 1T token,再换成稀疏注意力。V4.1 从第一步起就是稀疏的,序列长度 64K,没有稠密的热身阶段。V4 有一个短的 indexer 热身阶段,按 V3.2 的做法是让 indexer 拟合主注意力的分布。V4.1 没有这个阶段时 indexer 怎么训练,论文没有写。
容易混淆的几点
Reuse 不等于这一层没有注意力。 Reuse 层有自己的 query、自己的滑窗 KV、自己的输出投影和 attention sink。它共用的只是全局 KV 和选中的位置。
共享池不是一份新的缓存。 它只是几个指针。真正占空间的是 4 个 Full 层各自的 main KV 和 indexer K。「共享池」这个名字是本连载起的,论文没有给它命名。
encoder 的 3 组不共用 KV。 每组自己的 Full 层生成一份,组与组之间没有关系。
CSA2 不是 CSA 的小改版。 V4 的 CSA 与 HCA 交错在 V4.1 里没有了,论文的说法是 "pure CSA2"。压缩率、压缩算子、indexer K 的来源、层与层的关系都变了。
m = 1 m = 1 m = 1 也叫 CSA2。 论文明确说 CSA2 把不压缩的情况当作压缩率为 1 的特例。decoder 的全局 KV 没有压缩,但仍然是稀疏读的。
下一篇
Reindex 层要给 100 万条打分,这是 decode 时剩下的最大一块随长度增长的计算。下一篇讲 Hierarchical Sparse Indexer 怎么把它变成常数,再讲 FP4 怎么把每条 KV 的字节数减半,最后把 890 字节的分项和 V4-Flash 的 3514 字节逐项对比。
资料
评论