专栏DeepSeek-V4.1 模型结构·序列3 / 9
15 min学习

DeepSeek-V4.1 模型结构(2):序列维度(中),CSA2 的三种模式

KV cache 的大小是三个量相乘:每条多大、每个 token 几条、多少层各存一份。前两个量 V4 已经压过,CSA2 压第三个。这一篇把一层全局注意力拆成四样东西,看哪些能跨层共用,由此得到 Full、Reindex、Reuse 三种模式;再讲压缩算子和 indexer 相对 V4 的两处简化,共享池里谁写谁读,以及 38 层为什么只需要 4 份全局 KV。

目录13 节

对应论文 §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 层Fullmain KVindexer K打分,选 512 条注意力第 3 层Reusemain KVindexer K打分,选 512 条注意力第 4 层Reusemain KVindexer K打分,选 512 条注意力第 5 层Reusemain KVindexer K打分,选 512 条注意力第 6 层Reusemain KVindexer K打分,选 512 条注意力第 7 层Reusemain KVindexer K打分,选 512 条注意力
decoder:第 20 – 39 层(画前 8 层)第 20 层Fullmain KVindexer K打分,选 512 条注意力第 21 层Reusemain KVindexer K打分,选 512 条注意力第 22 层Reusemain KVindexer K打分,选 512 条注意力第 23 层Reusemain KVindexer K打分,选 512 条注意力第 24 层Reindexmain KVindexer K打分,选 512 条注意力第 25 层Reusemain KVindexer K打分,选 512 条注意力第 26 层Reusemain KVindexer K打分,选 512 条注意力第 27 层Reusemain KVindexer 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 层各存一份。

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 层 Reuseencoder:第 0 – 19 层decoder:第 20 – 39 层prefill 时,绝大部分 prompt token 只算到这里decoder 的全局 KV 全部由投影得到只有滑窗分支,不读全局 KV40 层的注意力后面都接一个 MoE,下面各行省略不画写读写读重选索引读encoder 的三组各有一份全局 KV,由该组的 Full 层写入,组内 6 层共用encoder 共享池(每组一份)main KV:每 2 个 token 一条,512 维 FP4indexer K:每条 128 维 FP4top-512 索引:每个 query 一份,不进缓存decoder 只有一份全局 KV,由第 20 层从 encoder 末态投影出来,20 层共用decoder 共享池(只有一份)main KV:每个 token 一条,512 维 FP4indexer K:每条 128 维 FP4候选池:第 20 层选出的2048 块 × 8 = 16384 个位置,Reindex 层只在池内打分top-512 索引:第 20 层先写,每个 Reindex 层覆盖一次全局 KV 合计 890 B / tokenSingle-Pass mHC:每个子层一次读、一次写 + 混读用的是上一个子层算好的查表结果经门控后加进残差流logits最后一次只读:4 条流压成 1 条读主干第 37 – 39 层入口处 4 条流的平均一次前向出 5 个草稿 token 和各自的置信度encoder 末态:第 19 层的输出,也就是第 20 层的输入。decoder 所有层的全局 KV 都只从它投影(论文式 1)encoder 末态mHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 读算子:对 4 条流加权求和,得到子层的 5120 维输入。权重由上一个子层算好读 AmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 BmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 BmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 BmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 BmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 BmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 BmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 BmHC 写算子 C 与混合矩阵 B:子层输出按 C 分到 4 条流,同时 4 条流按双随机矩阵 B 互相混合写 C · 混 B从头训练的视觉编码器:2D-RoPE、RMSNorm、SwiGLU,hidden 1024DeepSeek-ViT32 层 · patch 143×3 pixel-unshuffle 把 9 个相邻 patch 拼到通道维,再过两层 MLP 投影到主干宽度3×3 重排 + MLPtoken ÷ 9 → 5120 维图像位置上的 embedding 被视觉特征覆盖Token Embedding129280 × 5120前两层只有滑窗分支,没有全局 KV滑窗注意力第 0、1 层 · 窗口 128sqrt(softplus) 打分;文本和图像 token 各用一套负载均衡偏置;routed expert 权重 FP4DeepSeekMoE每层的 FFN · 384 选 6 + 1 sharedFull 模式:自己压缩 main KV(每 2 个 token 一条),投影出 indexer K,打分选 top-512,全部写进共享池CSA2 · Full第 2 / 8 / 14 层 · m = 2Reuse 模式:只有自己的 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-512CSA2 · Reindex第 24 / 28 / 32 / 36 层Reuse 模式:用本组 Reindex 层选出的 top-512CSA2 · 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 bl⏟每条多少字节×1ml⏟每个 token 几条.\text{字节 / token} = \sum_{\text{存 KV 的层 } l}\ \underbrace{b_l}_{\text{每条多少字节}} \times \underbrace{\frac{1}{m_l}}_{\text{每个 token 几条}} .

三个量可以分别压,效果相乘。论文 §2.3 就是按这三个维度回顾前人工作的:

  • 每条多大。 GQA 减少 KV 的头数。MLA 让所有头共用一条低维的 latent。V4 的条目是一条 512 维向量,K 和 V 是同一个东西。
  • 每个 token 几条。 V4 的 CSA 把每 m=4m = 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×14+20×1128≈5.4121 \times \frac{1}{4} + 20 \times \frac{1}{128} \approx 5.41 条。V4.1-Flash 是 2.5 条。这一篇讲 2.5 是怎么来的。每条多少字节留到下一篇。

一层全局注意力算四样东西,其中三样可以跨层共用

要跨层共用,先看一层里有什么。V4 的一个 CSA 层,全局分支按顺序算四样东西:

  1. main KV。 把本层的隐藏态投影成 512 维,每 mm 个 token 压成一条。这是注意力最后读的内容,要存进 KV cache。
  2. indexer K。 每条 main KV 配一个 128 维的索引用的 key,也要存进 KV cache。
  3. top-k 索引。 indexer 用当前 token 的 indexer query 和所有 indexer K 打分,选出分数最高的 kk 条。这是一组位置编号,每个 query 现算,不进缓存。
  4. 注意力本身。 用本层的 query 读选中的 kk 条 main KV,加上本层滑窗里的 128 条。

前两样是存储,第三样是计算,第四样是这一层真正的工作。

CSA2 的规定是:第四样每层自己做,前三样可以用前面某一层的。而且「共用缓存」和「共用选择结果」分开决定。把前三样的各种组合列出来,有意义的只有三种:

模式main KVindexer Ktop-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/36Reuse:30 层跨层共享池 SharedAttentionRuntime(356 B/条 × 2.5 条/token = 890 B/token)[64,512]RoPE 前 latentindex_ktop-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×512RoPE 末 64 维[64, 512]indexer Q(每层自有)wq_b 1280→32×128weights_proj 5120→32FP4SWA 路径(每层自有)wkv 5120→512 + kv_normRoPE 末 64 → FP8ring cache 128 槽sparse_attn(每层自有)SWA 128 + 全局 top-512+ sink → 反旋转 RoPECompressorwkv、wgate 5120→512softmax(gate)⊙kv按 m 求和(不重叠)m=1 无 wgate(层 20)latent [N/m, 512]RoPE 末 64 维 → FP4E2M1 + E4M3 scale/16indexer Kwk 512→128 + k_normRoPE θ=160000 → FP4E8M0/32,无 Hadamard打分因果掩码 → top-512按位置 sort层 20 另产 candidates(2048 块 × 8 = 16384)输出投影(每层自有)wo_a 8 组 ×(4096→1024)wo_b 8192→5120Q 低秩(每层自有)wq_a 5120→1280+ q_norm →wq_b 1280→64×512RoPE 末 64 维[64, 512]重打分读共享 index_kcandidates 掩码→ top-512,sort写回 topk_idxs(覆盖)sparse_attn(每层自有)SWA 128 + 全局 top-512+ sink → 反旋转 RoPESWA 路径(每层自有)wkv 5120→512 + kv_normRoPE 末 64 → FP8ring cache 128 槽indexer Q(每层自有)wq_b 1280→32×128weights_proj 5120→32FP4;K 从共享池读本层不产 KV / index_k无 Compressor无 wk / k_norm输出投影(每层自有)wo_a 8 组 ×(4096→1024)wo_b 8192→5120Q 低秩(每层自有)wq_a 5120→1280+ q_norm →wq_b 1280→64×512RoPE 末 64 维[64, 512]无 indexer、无 Compressor注意力参数清单与纯 SWA 层相同SWA 路径(每层自有)wkv 5120→512 + kv_normRoPE 末 64 → FP8ring cache 128 槽sparse_attn(每层自有)读共享 compress_kv+ 共享 topk_idxs+ SWA 128 + sink → 反旋转输出投影(每层自有)wo_a 8 组 ×(4096→1024)wo_b 8192→5120compress_kv主 KV latent,FP4(288 B/条)index_kindexer 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)0281420242832363940 – 423 组,每组共用一份全局 KV(m = 2)20 层共用一份全局 KV(m = 1)DSparkFull ×4:自己生成全局 KV(第 2、8、14、20 层)Reindex ×4:只重新选 top-512(第 24、28、32、36 层)Reuse ×30:KV 和索引都用现成的纯滑窗 ×2(第 0、1 层)DSpark 草稿层 ×3Engram:第 1、14 层入口DSpark 读第 37 – 39 层入口的残差流
格子的颜色是模式。下方的括号标出共用同一份全局 KV 的范围。

排布来自论文 §4.2.1,和 config 的 kv_source_layer_ids、index_source_layer_ids 一致:

范围分组每组的排法mm
第 0、1 层—纯滑窗,没有全局分支—
encoder 第 2 – 19 层3 组,每组 6 层1 层 Full + 5 层 Reuse2
decoder 第 20 – 23 层1 组1 层 Full + 3 层 Reuse1
decoder 第 24 – 39 层4 组,每组 4 层1 层 Reindex + 3 层 Reuse1

带全局分支的 38 层里,Full 4 层,Reindex 4 层,Reuse 30 层。

每个 token 存几条全局 KV,只看 Full 层:

3×12⏟encoder 的 3 个 Full+1×11⏟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}.

V4-Flash 是 5.41 条。条数减少到不足一半,同时每一条覆盖的范围更细:V4 最细是 4 个 token 一条,V4.1 的 decoder 是 1 个 token 一条。层这个维度省出来的空间,一部分还给了序列这个维度。

decoder 只有一个 Full 层,这和上一篇的 CED 对上了:decoder 的全局 KV 必须由 H20H_{20} 投影,第 20 层的输入正好就是 H20H_{20}。

论文没有解释为什么 encoder 是 6 层一组、decoder 是 4 层一组,也没有解释为什么 encoder 用 3 个 Full 而不用 Reindex。能确定的只有存储上的代价:encoder 每多一个 Full 层,每 token 多 0.5 条;decoder 每多一个,多 1 条。

压缩算子去掉了重叠和位置偏置

Full 层要自己生成 main KV。先把 V4 的压缩算子摆出来。记 CtC_t 是 token tt 的 KV 向量,ZtZ_t 是它的压缩权重,两者都是 512 维,由隐藏态各经一个线性层得到。V4 的 CSA 让相邻两个条目共用一半的 token:

CiComp=∑j=m(i−1)mi−1Sja⊙Cja+∑j=mim(i+1)−1Sjb⊙Cjb,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 ,

其中 SS 是对 2m2m 个槽的 ZZ 加上一个可学习的块内位置偏置 BB 之后做的 softmax。

token01234567891011块 0(个 token)块 1(个 token)块 2(个 token)a, 位 0b, 位 0a, 位 1b, 位 1a, 位 2b, 位 2a, 位 3b, 位 3a, 位 0b, 位 0a, 位 1b, 位 1a, 位 2b, 位 2a, 位 3b, 位 3a, 位 0b, 位 0a, 位 1b, 位 1a, 位 2b, 位 2a, 位 3b, 位 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 要算两套 CC 和 ZZ,共 4 个投影矩阵;有一个位置偏置;decode 时要多缓冲一个块。

CSA2 把它退回最简单的形式。论文只用文字描述,下面的式子是我按官方代码写出来的:

[Smi; … ; Sm(i+1)−1]=Softmax⁡([Zmi; … ; Zm(i+1)−1]),CiComp=∑j=mim(i+1)−1Sj⊙Cj.[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 .

每个条目只由自己这一组的 mm 个 token 得到,相邻条目不共用 token。softmax 仍然是逐通道的:512 个通道各有一套「组里哪个 token 更重要」的权重。没有位置偏置 BB。

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=1m = 1 是它的特例。一组只有一个 token,softmax 只有一项,权重恒为 1:

CtComp=Ct.C^{\text{Comp}}_t = C_t .

压缩权重 ZZ 不起任何作用。官方代码对这个情况单独处理,连 ZZ 的投影矩阵都不建:

python
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 是 CC 的投影,wgate 是 ZZ 的投影。权重文件里第 2、8、14 层有 compressor.wgate,第 20 层没有。

论文给的理由是这两处简化让实现更简单、训练更快,没有给消融。V4 第 1 篇说过重叠是为了缓解块边界的硬切。这是我的解读:mm 从 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 条目投影:

kiI=RMSNorm⁡(CiComp WIK),WIK∈R512×128.k^I_i = \operatorname{RMSNorm}\big(C^{\text{Comp}}_i\, W^{IK}\big), \qquad W^{IK} \in \mathbb{R}^{512 \times 128}.

这个式子也是按代码写的。几个细节只在代码里能看到:

  • 投影用的是加 RoPE 之前的 main KV。投影完,indexer K 的最后 64 维再加自己的 RoPE。
  • V4 在量化前对 indexer K 做一次 Hadamard 旋转,V4.1 没有这一步。
  • indexer K 按 FP4 存,每 32 维一个 scale。

WIKW^{IK} 只有 512×128=65536512 \times 128 = 65536 个参数。V4-Flash 的 indexer 压缩算子是 4 个 4096×1284096 \times 128 的矩阵,约 210 万。

打分公式没有变,仍然是 V4 第 2 篇的那个:

It,s=∑j=1HIwt,jI⋅ReLU⁡(qt,jI⋅ksI).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).

qt,jI\mathbf{q}^I_{t,j} 是 token tt 的第 jj 个 indexer query,ksI\mathbf{k}^I_s 是第 ss 条的 indexer K,wt,jIw^I_{t,j} 是由 token tt 决定的每个头的权重。头数 HIH^I 从 V4-Flash 的 64 降到 32,每头仍是 128 维。每个 query 选 512 条。

共享池里有四样东西,两样进缓存,两样现算

三种模式靠一个共享的对象传递数据。官方代码里它只有四个槽:

python
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 层覆盖它之前,中间的层读到的都是这一份。

注意力层的代码里,三种模式的分支只有两个判断:

python
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。

02814202428323639第 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 层层的模式encoderdecodermain KVcompress_kv写读读读读读写读读读读读写读读读读读写读读读读读读读读读读读读读读读读读读读进 KV cacheindexer Kindex_k写写写写读读读读进 KV cache候选池candidates写读读读读每个 query 现算top-512 索引topk_idxs写读读读读读写读读读读读写读读读读读写读读读写读读读写读读读写读读读写读读读每个 query 现算FullReindexReuse纯滑窗这一层生成并写入这一层读取
共享池的四样东西在 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∈St ∪ 滑窗softmax⁡s(qt(l)⋅ks) vs.\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 .

St\mathcal{S}_t 是共用的,qt(l)\mathbf{q}^{(l)}_t 是本层的。两个 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 万。没有压缩算子,也没有 WIKW^{IK}。
  • 计算。 要跑一遍 indexer,给所有看得见的条目打分。decoder 的 m=1m = 1,1M 上下文下就是 100 万条。这是 Reindex 的主要开销,下一篇的 Hierarchical Sparse Indexer 专门处理它。
  • 收益。 decoder 每 4 层换一批条目,20 层一共有 5 次选择。

「共用缓存」和「共用选择结果」分开决定,指的就是这个:Reindex 层共用了缓存,没有共用选择结果。

三种模式的参数清单

把每种模式多出来的参数列出来。「公共部分」是每一层都有的,共约 127M。它包括 query 的两个投影矩阵、滑窗 KV 的投影、输出投影的两个矩阵,还有每个头一个的 attention sink。

模式层公共部分之外的参数约多少
纯滑窗0、1无0
Reuse30 层无0
Reindex24、28、32、36indexer 的 wq_b(1280 × 4096)和 weights_proj(5120 × 32)5.4M
Full,m=1m = 120上一行,加 compressor.wkv(5120 × 512)和 indexer 的 wk(512 × 128)8.1M
Full,m=2m = 22、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=1m = 1 也叫 CSA2。 论文明确说 CSA2 把不压缩的情况当作压缩率为 1 的特例。decoder 的全局 KV 没有压缩,但仍然是稀疏读的。

下一篇

Reindex 层要给 100 万条打分,这是 decode 时剩下的最大一块随长度增长的计算。下一篇讲 Hierarchical Sparse Indexer 怎么把它变成常数,再讲 FP4 怎么把每条 KV 的字节数减半,最后把 890 字节的分项和 V4-Flash 的 3514 字节逐项对比。

资料

评论