第 13 章把 KV cache
压成固定大小的记忆板,换来的是精确检索变弱。这一关走另一条路:K/V 照样全存,
但每个 query 只挑一小部分去算,这叫
稀疏注意力。
挑法取自 DeepSeek 的
NSA:
压缩、挑选、滑窗三条支路,再用门控加权。配套代码 phase4-efficiency/12_sparse_attn.py,
含四组语言建模训练(Mac MPS)与三组"远处找钥匙"训练(RTX 4090)。
第 2 章的因果注意力里,位置 t 的 query 要和 t+1 个 key 逐个打分(见
因果掩码)。
上下文越长,每一步读得越多。拖滑块选一个 query 位置,切换三种读法,看这一行掩码里哪些格子亮着:
tokens_read(读数公式)· 掩码 CAUSAL / WINDOW / BLOCK_DONE
论文默认:压缩块 l=32、步长 d=16;挑选块 l′=64、挑 n=16 块;滑窗 w=512。
解码时每一步最多读 s/d 个压缩 token、n·l′ 个挑选 token、w 个邻近 token。拖上下文长度:
s/16(论文配置)或 t/8(本关配置),
和上下文长度成正比,只是除掉了一个常数。挑选和滑窗两条支路读的量是固定的(1024 与 512)。
上下文到 1M 时,压缩摘要有 65,536 个,占 NSA 读取量的 98% 左右;这也是后来 DeepSeek-V3.2、V4 和 Qwen3.8 换用更轻的打分器、或把压缩率拉大的原因之一(STEP 5)。
tokens_read 的注释写的是"去重之前的上界":当前块同时在滑窗和挑选里,最近的几块也可能被挑中。
上面的格子图在 t 有真实挑选记录时(t = 15, 31, 47 … 255)会额外给出去重后的原始 token 数。
三条支路各有自己的 K/V 投影,所以即使位置重复,读到的向量也不同,按支路分开计数也说得通。
下面是验证集里一段真实的 256 个字符(NSA 那组训练脚本存下的 sample_text),按
8block_len 切成 32 块。
拖 query 位置:已经写满的块才有摘要可看(橙框),query 所在的块还没写完(紫色虚框),后面的块看不见。
点任意一块,看 8 个 key 怎么变成 1 个摘要 key:
SparseAttention.forward「支路 1:压缩」· blk_pos / cmp_k / cmp_v / sink_kvBLOCK_DONE 做掩码:块的最后一个位置 ≤ t,这块的摘要才对 t 可见。
query 所在的那块没有摘要可看,它的内容交给挑选支路(当前块必选)和滑窗支路。
sink_kv(可学习参数,初值为 0),
拼在所有摘要的最前面。于是 t = 0…6 的 query 至少能看见这一格,读数公式里压缩支路那一项也因此多了 +1。
l=32、步长 d=16,
相邻块有一半重叠,论文的解释是减轻信息被切碎。本关用一个 Linear 加可学习的 blk_pos 代替 MLP,
步长等于块长(不重叠),这样挑选块和压缩块能共用同一套切分,打分可以直接复用(下一步)。
另外,摘要 key 没有加 RoPE,块内先后顺序由 blk_pos 和拼接的顺序表达。
压缩支路算注意力时已经有了 query 对每个摘要的分数。挑选支路直接复用这组分数,取最高的 top_n = 4 块(query 所在的当前块强制入选),把这几块里的原始 token 逐个拿来做注意力。 下面是训练结束后最后一层真实挑中的块。点一个 query 位置:
SparseAttention.forward「支路 2:挑选」· score.topk · 记录在 self.last["chosen"]softmax(q·摘要key / √head_size),和 NSA 论文 §3.3.2 式 (8) 一致。
同一组 K/V 的两个 query 头把分数相加后一起挑(论文式 (10)),这样一组只需要加载一份被选中的 K/V。
当前块的分数被直接设成 1e4,保证入选。
torch.no_grad() 里,挑哪几块这件事本身不产生梯度。
摘要 key 通过压缩支路的输出拿到梯度,越训越能反映块的内容,挑选复用它的分数,也就跟着变准。
挑中之后,块里 token 的注意力计算是正常可导的。
chosen[0, t])。
每层有 2 组 K/V 头,第 1 组挑了什么没有存;前三层挑了什么也没有存。
答案要从值 token 那里取,取的动作可以发生在更早的层,之后再由残差流一路带到最后一层,
所以最后一层这一组不一定还需要回头看钥匙。
门控数据也和这一点对得上:找钥匙模型第 2、3 层挑选门的均值是 0.609 与 0.565,最后一层是 0.458(STEP 4 可切换查看)。
这张图能说明的是:挑选确实会越过 64 个噪声位置跳回前面的钥匙区,但不能说明检索具体发生在哪一层。
第三条支路是
滑窗注意力:
最近 32 个 token 必看。三条支路各算出一个输出,再各乘一个门值相加。门是
sigmoid(Linear(x)),每个头、每个位置、每条支路各一个,三个门互不约束。
开关下面三条支路,看读数怎么变;对应组合真训练过的,会画出每层的平均门值:
SparseAttention.forward「支路 3:滑窗」与「门控加权」· self.gate = nn.Linear(n_embd, 3 * n_head)g ∈ [0,1],由输入特征经 MLP 和 sigmoid 得到。
softmax 会强迫三个门加起来等于 1,一条支路变重,另外两条必须变轻;sigmoid 允许三条同时开大或同时关小,
等于顺带学了一个整体缩放。图里语言建模 NSA 每层三个门之和从 0.65 到 1.03,找钥匙 NSA 在 1.08 到 1.26 之间,用的就是这个自由度。
gate_mean_last_layer 只喂了 1 条验证样本(语言建模 NSA:cmp 0.145、sel 0.232、win 0.652),
图里的 gate_mean_per_layer 喂的是整批(语言建模 32 条,找钥匙 128 条),最后一层是 cmp 0.14、sel 0.255、win 0.637。
样本不同,平均值略有差别,两者给出的排序一致。
Linear(128 → 2×2×32),NSA 那组总参数 0.947M,全注意力 0.743M。STEP 5 比 loss 时要算上这一点。
同一份代码、同一个种子,只换 --branches。两套实验在不同硬件上跑:
语言建模 4 组在 Mac MPS(各 4000 步),远处找钥匙 3 组在 RTX 4090 · CUDA · PyTorch 2.4.1(各 6000 步,batch 128)。
两套之间的训练秒数不能互相比较。切换指标,点行看细节:
F.scaled_dot_product_attention 照样把整张矩阵算完,
被掩掉的位置只是不参与 softmax。三条支路就是三次完整的注意力,再加压缩和打分,所以 MPS 上 NSA 那组训练 2756 秒,全注意力 801 秒。
算出来的数值和真稀疏实现一致,速度不一致。
12 × (3 × Gated DeltaNet + 1 × QSA):四层里三层用线性注意力省存储,剩下一层用稀疏注意力省读取。
DeepSeek-V4 在稀疏之外还做了 KV 压缩(CSA 每 4 个 token 压成一条,HCA 每 128 个压成一条),
论文报告 V4-Pro 的 KV cache 是 V3.2 的 10%。
l=32 步长 d=16;挑 n=16 块 × l′=64(含开头 1 块、本地 2 块);滑窗 w=512I = Σⱼ wⱼ · ReLU(qⱼ · k)n_win=128;CSA、HCA 逐层交错12 × (3 GDN + 1 QSA);在 256K 长度的持续预训练阶段把全注意力层换成 QSApython 12_sparse_attn.py --branches full | win | cmp,sel | cmp,sel,win · 加 --task recall 换成找钥匙