← 返回首页
14_sparse_attention_viz.html

稀疏注意力:每个 query 只挑一部分历史去读

第 13 章把 KV cache 压成固定大小的记忆板,换来的是精确检索变弱。这一关走另一条路:K/V 照样全存, 但每个 query 只挑一小部分去算,这叫 稀疏注意力。 挑法取自 DeepSeek 的 NSA: 压缩、挑选、滑窗三条支路,再用门控加权。配套代码 phase4-efficiency/12_sparse_attn.py, 含四组语言建模训练(Mac MPS)与三组"远处找钥匙"训练(RTX 4090)。

STEP 1
每个 query 读多少
STEP 2
压缩:一块一个摘要
STEP 3
挑选:只细看 top-n 块
STEP 4
滑窗与门控
STEP 5
真实对照 · 工业版
这一关 block_size 256 上下文长度 block_len 8 一块几个 token top_n 4 每个 query 细看几块(含当前块) window 32 滑窗长度 n_kv_head 2 K/V 头数(同组共用一次挑选) 另:n_embd 128 · n_head 4 · head_size 32 · 4 层。window=32 与 top_n×block_len=32 数值相同,含义不同。

① 每个 query 到底读多少个 key

第 2 章的因果注意力里,位置 t 的 query 要和 t+1 个 key 逐个打分(见 因果掩码)。 上下文越长,每一步读得越多。拖滑块选一个 query 位置,切换三种读法,看这一行掩码里哪些格子亮着:

143
↳ 代码:12_sparse_attn.py 的 tokens_read(读数公式)· 掩码 CAUSAL / WINDOW / BLOCK_DONE
放大换成 NSA 论文的默认配置

论文默认:压缩块 l=32、步长 d=16;挑选块 l′=64、挑 n=16 块;滑窗 w=512。 解码时每一步最多读 s/d 个压缩 token、n·l′ 个挑选 token、w 个邻近 token。拖上下文长度:

64K
压缩支路读的数为什么还随长度增长?
压缩支路给每块留一个摘要,块数 = 上下文长度 / 步长。所以它读的量是 s/16(论文配置)或 t/8(本关配置), 和上下文长度成正比,只是除掉了一个常数。挑选和滑窗两条支路读的量是固定的(1024 与 512)。 上下文到 1M 时,压缩摘要有 65,536 个,占 NSA 读取量的 98% 左右;这也是后来 DeepSeek-V3.2、V4 和 Qwen3.8 换用更轻的打分器、或把压缩率拉大的原因之一(STEP 5)。
三条支路的读数加起来,会不会重复算了同一个 token?
会。tokens_read 的注释写的是"去重之前的上界":当前块同时在滑窗和挑选里,最近的几块也可能被挑中。 上面的格子图在 t 有真实挑选记录时(t = 15, 31, 47 … 255)会额外给出去重后的原始 token 数。 三条支路各有自己的 K/V 投影,所以即使位置重复,读到的向量也不同,按支路分开计数也说得通。
↳ 下一步:只读几十个 key,前提是知道该读哪几块。先给每块做一个便宜的摘要。

② 压缩:把一块 token 压成一个摘要

下面是验证集里一段真实的 256 个字符(NSA 那组训练脚本存下的 sample_text),按 8block_len 切成 32 块。 拖 query 位置:已经写满的块才有摘要可看(橙框),query 所在的块还没写完(紫色虚框),后面的块看不见。 点任意一块,看 8 个 key 怎么变成 1 个摘要 key:

100
已写满 · 摘要可见 query 所在块 · 摘要不可见 query 自己 · 空格 ↵ 换行
↳ 代码:12_sparse_attn.py 的 SparseAttention.forward「支路 1:压缩」· blk_pos / cmp_k / cmp_v / sink_kv
为什么摘要要等块写满才能看?
摘要是把块里 8 个 key 拼起来再过线性层,它混进了这 8 个位置的全部信息。 如果 query 在块中间就去看这个摘要,等于偷看了自己后面的 token,训练时就泄漏了答案。 所以代码用 BLOCK_DONE 做掩码:块的最后一个位置 ≤ t,这块的摘要才对 t 可见。 query 所在的那块没有摘要可看,它的内容交给挑选支路(当前块必选)和滑窗支路。
开头 7 个位置一块完整的都没有,压缩支路怎么算?
全部被掩掉时 softmax 的分母是 0,会算出 NaN。代码加了一个永远可见的"空槽" sink_kv(可学习参数,初值为 0), 拼在所有摘要的最前面。于是 t = 0…6 的 query 至少能看见这一格,读数公式里压缩支路那一项也因此多了 +1。
论文里的压缩和这里有哪些不同?
NSA 论文(§3.3.1 式 (7))的压缩函数 φ 是一个带块内位置编码的 MLP,块长 l=32、步长 d=16, 相邻块有一半重叠,论文的解释是减轻信息被切碎。本关用一个 Linear 加可学习的 blk_pos 代替 MLP, 步长等于块长(不重叠),这样挑选块和压缩块能共用同一套切分,打分可以直接复用(下一步)。 另外,摘要 key 没有加 RoPE,块内先后顺序由 blk_pos 和拼接的顺序表达。
↳ 下一步:query 对 32 个摘要打完分,分数最高的几块就值得逐字细看。

③ 挑选:只细看分数最高的 top-n 块

压缩支路算注意力时已经有了 query 对每个摘要的分数。挑选支路直接复用这组分数,取最高的 top_n = 4 块(query 所在的当前块强制入选),把这几块里的原始 token 逐个拿来做注意力。 下面是训练结束后最后一层真实挑中的块。点一个 query 位置:

↳ 代码:12_sparse_attn.py 的 SparseAttention.forward「支路 2:挑选」· score.topk · 记录在 self.last["chosen"]
挑选的分数怎么算?topk 不可导,挑选能学吗?
分数 = softmax(q·摘要key / √head_size),和 NSA 论文 §3.3.2 式 (8) 一致。 同一组 K/V 的两个 query 头把分数相加后一起挑(论文式 (10)),这样一组只需要加载一份被选中的 K/V。 当前块的分数被直接设成 1e4,保证入选。

整个挑选过程在 torch.no_grad() 里,挑哪几块这件事本身不产生梯度。 摘要 key 通过压缩支路的输出拿到梯度,越训越能反映块的内容,挑选复用它的分数,也就跟着变准。 挑中之后,块里 token 的注意力计算是正常可导的。
找钥匙准确率 99.85%,为什么图里挑中的块只有一半左右盖住钥匙?
脚本只保存了最后一层、第 0 组 K/V 头的挑选结果(chosen[0, t])。 每层有 2 组 K/V 头,第 1 组挑了什么没有存;前三层挑了什么也没有存。 答案要从值 token 那里取,取的动作可以发生在更早的层,之后再由残差流一路带到最后一层, 所以最后一层这一组不一定还需要回头看钥匙。 门控数据也和这一点对得上:找钥匙模型第 2、3 层挑选门的均值是 0.609 与 0.565,最后一层是 0.458(STEP 4 可切换查看)。 这张图能说明的是:挑选确实会越过 64 个噪声位置跳回前面的钥匙区,但不能说明检索具体发生在哪一层。
为什么当前块一定要选?
当前块还没写满,压缩支路里没有它的摘要,而紧挨着 query 的几个 token 对预测下一个字最重要。 NSA 论文的做法类似:16 个挑选块里固定包含开头 1 块和本地 2 块(§4.1)。
↳ 下一步:挑中的块是跳着的,最近的几十个字不一定都在里面。所以还要一条滑窗支路,并决定三条支路各听几成。

④ 滑窗兜底,门控分配权重

第三条支路是 滑窗注意力: 最近 32 个 token 必看。三条支路各算出一个输出,再各乘一个门值相加。门是 sigmoid(Linear(x)),每个头、每个位置、每条支路各一个,三个门互不约束。 开关下面三条支路,看读数怎么变;对应组合真训练过的,会画出每层的平均门值:

↳ 代码:12_sparse_attn.py 的 SparseAttention.forward「支路 3:滑窗」与「门控加权」· self.gate = nn.Linear(n_embd, 3 * n_head)
为什么用三个独立的 sigmoid,不用一个 softmax?
NSA 论文 §3.2 写的就是 sigmoid:g ∈ [0,1],由输入特征经 MLP 和 sigmoid 得到。 softmax 会强迫三个门加起来等于 1,一条支路变重,另外两条必须变轻;sigmoid 允许三条同时开大或同时关小, 等于顺带学了一个整体缩放。图里语言建模 NSA 每层三个门之和从 0.65 到 1.03,找钥匙 NSA 在 1.08 到 1.26 之间,用的就是这个自由度。
图里最后一层的数和 json 里的 gate_mean_last_layer 为什么差一点?
两个数取自不同的前向: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。 样本不同,平均值略有差别,两者给出的排序一致。
三条支路为什么各有一套 K/V 投影?
NSA 论文的理由是防止一条支路的梯度通过共享投影去影响另一条:滑窗支路学局部模式最快, 共用投影时它会主导 K/V 的学习,压缩和挑选支路就学不起来。代价是参数变多: 本关每层每条支路一个 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)。 两套之间的训练秒数不能互相比较。切换指标,点行看细节:

NSA 的 val loss 最低,能说它比全注意力好吗?
不能直接这么说。三个原因: ① 参数不一样多:NSA 0.947M,全注意力 0.743M,多 27.5%,多出来的主要是三套 K/V 投影和压缩器; ② 单种子:每组只跑了一次; ③ 评测本身有抖动:全注意力最后三次评测是 1.4946 → 1.4998 → 1.5038,相邻两次差到 0.005,NSA 是 1.4820 → 1.4970 → 1.4825,差到 0.015。

按这个抖动量级看:只滑窗 1.5052 与全注意力 1.5038 差 0.0014,是噪声;压缩+挑选(不带滑窗)1.5321,高出 0.028,更可能是真的, 说明去掉滑窗之后局部信息不够。NSA 低 0.021,方向和论文一致(论文 Figure 4:27B 模型上 NSA 的预训练 loss 低于全注意力),但这组数据不足以给出幅度。
只滑窗在语言建模上几乎不掉点,为什么找钥匙会掉到 3.1%?
tiny shakespeare 是字符级文本,预测下一个字符主要靠最近几十个字符,32 的窗口基本够用。 找钥匙任务是专门反过来设计的:128 个位置里,钥匙和值放在 1–32,中间 33–96 是噪声,提问在 97–127。 滑窗在提问区往回最远看到位置 66,碰不到任何一把钥匙。值有 32 种,乱猜的准确率是 1/32 = 3.125%,只滑窗那组 3.11% 就是乱猜。 全注意力 100%,NSA 99.85%:挑选支路把前面的块找了回来。
教学版为什么没有变快,反而更慢?
教学版是在稠密的 256×256 注意力矩阵上加掩码来"模拟"稀疏:F.scaled_dot_product_attention 照样把整张矩阵算完, 被掩掉的位置只是不参与 softmax。三条支路就是三次完整的注意力,再加压缩和打分,所以 MPS 上 NSA 那组训练 2756 秒,全注意力 801 秒。 算出来的数值和真稀疏实现一致,速度不一致。

真正变快要靠只取被选中块的 kernel。NSA 论文用 Triton 实现,在 8×A100 上 64K 长度训练前向快 9.0 倍、反向快 6.0 倍; 解码每步的访存量从 65,536 个 token 降到 5,632 个,按访存量估计的加速是 11.6 倍(表 4)。
稀疏注意力和第 13 章的线性注意力,各省的是什么?
稀疏注意力省"读":K/V 全部保存,每个 query 只读其中一部分,所以精确检索能力保留(本关找钥匙 99.85%),但 KV cache 照样随上下文增长。 线性注意力省"存":历史压进固定大小的状态,KV cache 不增长,代价是很难从状态里精确取回很久以前的某个 token。

两者可以叠加。Qwen3.8-Flash-Next 是 12 × (3 × Gated DeltaNet + 1 × QSA):四层里三层用线性注意力省存储,剩下一层用稀疏注意力省读取。 DeepSeek-V4 在稀疏之外还做了 KV 压缩(CSA 每 4 个 token 压成一条,HCA 每 128 个压成一条), 论文报告 V4-Pro 的 KV cache 是 V3.2 的 10%。
工业版四种真实模型怎么挑

NSA

DeepSeek · arXiv 2502.11089
  • 三支路 + sigmoid 门,和本关同一结构
  • 压缩 l=32 步长 d=16;挑 n=16 块 × l′=64(含开头 1 块、本地 2 块);滑窗 w=512
  • 挑选分数复用压缩支路的注意力分数
  • 用 27B 总参 / 3B 激活的 MoE 模型从预训练开始按稀疏方式训练(260B token)

DSA

DeepSeek-V3.2 · arXiv 2512.02556 §2.1
  • 单独一个 lightning indexer 打分:I = Σⱼ wⱼ · ReLU(qⱼ · k)
  • 按 token 挑,不按块:每个 query 选 2048 个 token
  • indexer 头数少,可用 FP8 实现
  • 基于 MLA 的 MQA 模式实例化;先稠密预热 1000 步,再稀疏训练 15000 步

CSA + HCA

DeepSeek-V4 · arXiv 2606.19348(2026 preview)
  • CSA:每 4 个 token 的 KV 压成 1 条,再用 indexer 挑 top-k 条:Flash 512,Pro 1024
  • HCA:每 128 个 token 压成 1 条,对压缩条目做稠密注意力
  • 另有滑窗支路 n_win=128;CSA、HCA 逐层交错
  • indexer 的注意力计算用 FP4

QSA

Qwen3.8-Flash-Next 技术报告 §2.1.2
  • indexer 是 MQA:4 个 query 头 + 1 个共享 key 头
  • key 每 4 个 token 平均池化成一块,只给已写满的块打分
  • 预算 2048 token,即最多 512 块,再加上最后未写满块的尾部 token
  • 48 层 = 12 × (3 GDN + 1 QSA);在 256K 长度的持续预训练阶段把全注意力层换成 QSA
四家共同点:K/V 全存,每个 query 先用便宜的方式给历史打分,再只对高分部分做完整注意力。差别在打分器(复用压缩分数 / 独立 indexer)和挑的单位(块 / token / 压缩条目)。
达标本关实测
↳ 跑法:python 12_sparse_attn.py --branches full | win | cmp,sel | cmp,sel,win · 加 --task recall 换成找钥匙
🎉 稀疏注意力 · 通关
你亲手拆开了 NSA 的三条支路:压缩给每块留摘要,挑选只细看几块,滑窗兜住最近的字,门控决定各听几成。K/V 全存、只读一部分,找钥匙照样答对。