← 返回首页
13_linear_attention_viz.html

线性注意力:把 KV cache 换成一块固定大小的记忆板

第 12 章的注意力推理时要把每个 token 的 K/V 都存下来,上下文越长, KV cache 越大。线性注意力 只留一块 head_size × head_size 的矩阵,每来一个 token 往里写一次,大小和上下文长度无关。 本关从最朴素的写法出发,加上 delta rule 和遗忘门,得到 Qwen3-Next 用的 Gated DeltaNet, 再用 9 次真实训练看它在语言模型和精确检索上各输赢在哪。配套代码 phase4-efficiency/11_linear_attn.py。

STEP 1
KV cache 越存越长
STEP 2
记忆板:写入是加外积
STEP 3
先擦再写:delta 与遗忘门
STEP 4
五种层配方 · 真实对照
STEP 5
精确检索与 3:1 混合
这一关 n_embd 128 主干维度 n_head 4 头数 head_size 32 记忆板 32×32 n_kv_head 2 只在注意力层 block_size 64 训练长度 层配方 3:1 3 GDN + 1 注意力 4 层 · 其余零件同第 12 章四件齐上 · 训练在 RTX 4090 / CUDA / PyTorch 2.4.1

① KV cache 越存越长

全注意力层每个 token 要存一份 K 和 V,上下文每长一个 token,cache 就多一份。 记忆板层每个头只存一块 head_size × head_size 的矩阵,读多长的文本都是这么大。 拖滑块改上下文长度,比较三种层配方要存多少个数:

4K
玩具这三个数是怎么算出来的?
11_linear_attn.py 训练完把两项写进 json: kv_per_token = 2 × n_kv_head × head_size × 注意力层数, state_const = n_head × head_size × head_size × 记忆板层数。 全注意力(4 层):2 × 2 × 32 × 4 = 512 个数 / token; 全 GDN(4 层):4 × 32 × 32 × 4 = 16,384 个数,与长度无关; 3:1 混合:注意力 1 层 = 128 个数 / token,记忆板 3 层 = 12,288 个数。 三组数分别来自 runs/cuda/ch13_lm_attn.json、ch13_lm_gdn.json、ch13_lm_hybrid.json。 记忆板层还带一个 kernel=4 的短卷积,推理时要缓存最近 3 个 token 的输入,量级很小,这里没有计入。
Qwen3-Next 那一栏按什么配置估算?
出自 HF 模型卡 Qwen/Qwen3-Next-80B-A3B-Instruct:48 层,布局 12 × (3 × (Gated DeltaNet → MoE) → 1 × (Gated Attention → MoE))。 Gated Attention:16 个 Q 头、2 个 KV 头、head dim 256,所以 12 层注意力每 token 存 2 × 12 × 2 × 256 = 12,288 个数。 Gated DeltaNet:V 32 头、QK 16 头、head dim 128,本页按每层 32 × 128 × 128 个数估算记忆板,36 层共 18,874,368 个数。 "48 层全用注意力"和"48 层全用 GDN"两行是按同一套头配置算出的假想对照,用来和实际布局比。字节数一律按每个数 2 字节(BF16)折算,不含短卷积缓存。
↳ 下一步:一块固定大小的矩阵,怎么装得下越来越长的历史?先看怎么往里写。

② 记忆板:写入就是加一个外积

把 (k, v) 写进记忆板:S ← S + v kᵀ。读的时候拿 query 去乘:o = S q。 如果 q 恰好等于某个写过的 k,而且所有 k 互相正交、长度为 1,读出来的就是那个 k 对应的 v。 下面用 2 维向量演示(只为了上屏;真实每个头 head_size=32,记忆板 32×32)。 点按钮写入或撤掉一对,再选一个 key 去读:

37°
记忆板 S(2×2)
橙 = 正,蓝 = 负,颜色越深绝对值越大
–
读出 o = S q
–
这个 key 写入时的 v
–
误差 ‖o − v‖
↳ 代码:11_linear_attn.py 的 RecurrentMixer.forward(mode == "linear" 那一支)
这和第 2 章的 softmax 注意力是什么关系?
去掉 softmax 之后,注意力的输出是 o_t = Σ_i v_i (k_iᵀ q_t) = (Σ_i v_i k_iᵀ) q_t。 括号里那一坨就是 S:把所有历史 (k, v) 的外积加起来。先求和再乘 q,和逐个打分再加权,算出来一样。 所以不需要留着每个 k_i、v_i,留着它们的和就够了,这就是 Katharopoulos et al. 2020 "Transformers are RNNs"(arXiv 2006.16236)的出发点,论文把复杂度从 O(N²) 降到 O(N)。 代价是 softmax 没了:softmax 能让最匹配的那个 key 压过其余所有 key,线性版只能靠 key 之间正交来避免串扰。
32 维的记忆板能装多少对?
32 维空间里最多有 32 个互相正交的单位向量,所以一个头最多无串扰地存 32 对;再多就必然有 key 不正交,读出时混进别的 value。 本页的 2 维版本只能装 2 对,第三对一写进去就开始互相干扰。 Schlag et al. 2021(arXiv 2102.11174)从"快速权重"的角度推出了同一个结论,原文称之为 "memory capacity limitation"。 本项目的 linear 模式里 q、k 先过 elu + 1,所有分量都是正数,两个全正向量的点积不会小于 0,更难做到正交。
↳ 下一步:同一个 key 写两次,只加不减会把新旧值叠在一起。能不能先擦掉旧的再写?

③ 先擦再写:delta rule 与遗忘门

只加不减的写法擦不掉旧值。delta rule 先读出这个 key 当前存的值,只补差额;Gated DeltaNet 再加一个遗忘门 α。 下面是一串三次写入:先把 A 写成旧值,再写 B,最后把 A 改成新值。然后分别用 kA、kB 去读。 切换三种写法,拖 β、α 看读出值怎么变:

1.00
0.80
–
用 kA 读(新值应为 (−0.70, 0.40))
–
用 kB 读(写入时为 (−0.30, 0.80))
真实训完之后,第一层的 α 学成了多少

代码里 α 的偏置初始化成 3.0,开训时 α = sigmoid(3) ≈ 0.95。下面是训练 5000 步后,第一层 4 个头在 val 文本上的平均 α, 画成"一条记忆过 n 个 token 后还剩几成"(按平均 α 粗算 αⁿ;实际 α 每个 token 各算一次):

↳ 代码:11_linear_attn.py 的 RecurrentMixer(b_proj 算 β,a_proj 算 α)· 开关 --preset linear / delta / gdn
为什么叫 delta rule?
写入量正比于"想存的值 v"和"当前读出的值 S·k"之差,也就是误差 delta。 Yang et al. 2024(arXiv 2406.06484)把它写成 S_t = S_{t-1} − β_t (S_{t-1} k_t − v_t) k_tᵀ, 并解释为对在线回归损失 ½‖S k_t − v_t‖² 做一步 SGD,β_t 就是学习率。 k 做了 L2 归一化时,β_t = 1 让 I − k_t k_tᵀ 成为投影矩阵:先把 k_t 方向上的旧内容整个投影掉,再写新值。 最早把 delta rule 用到线性 Transformer 上的是 Schlag et al. 2021(arXiv 2102.11174)。
第一层的 α 只有 0.6 左右,是不是忘得太快了?
按平均 α 粗算,一条记忆过 1.2 到 1.7 个 token 衰减到一半,第一层的记忆板主要在记最近几个字符。 这里有两点要分开看:① 这是 64 段 × 256 个 val token 上的平均值,α 每个位置由 a_proj 各算一次,个别位置可以接近 1; ② 只统计了第一层,代码没有导出后面三层的 α。开训时 α ≈ 0.95,训练把第一层压到了 0.55–0.66, 说明对字符级语言模型的第一层来说,忘得快的记忆板让 loss 更低。更长的依赖交给哪几层、α 是多少,本实验没有测。
↳ 下一步:三种写法在玩具上的差别看清楚了。真训一遍语言模型,val loss 差多少?

④ 五种层配方的真实对照

同一份 tiny shakespeare、同一套超参(4 层、5000 步、batch 64、lr 1e-3、seed 1337),只换每层用什么: 全注意力、全 linear、全 delta、全 GDN、3 GDN + 1 注意力。训练长度 64,训完再拿 128 / 256 / 512 长的片段测 val loss。 切换指标,点一组看细节:

训练长度 64 上,五组 val loss 落在 1.5357–1.5722 之间;拉到 512,全注意力涨到 3.5171,GDN 与 3:1 混合停在 1.532 / 1.536。
为什么全注意力拉长到 512 就崩了?
全注意力层用 RoPE(第 12 章):位置越靠后转的角度越大。训练只见过 0–63 的位置, 到 128、256、512 时,q·k 里出现的相对距离和转角组合模型从没见过,val loss 从 1.539 涨到 1.941、2.855、3.517。 记忆板层没有位置编码,顺序靠"一个一个往里写"自带,所以 GDN 在更长的片段上 loss 反而略降(1.569 → 1.517 @256),多出来的上文有用。 3:1 混合的第 4 层也是带 RoPE 的注意力,512 时却是 1.536,没有崩。本实验没有单独做消融,这个现象的原因不在这里下结论。
@64 这些差距,哪些是真的、哪些是噪声?
每组只跑了一个种子。页面上同一个模型有两个 @64 的数:训练末尾的 val(200 个 batch)和外推测试里的 64(40 个 batch), 两者最多差 0.0069(3:1 混合:1.5446 对 1.5515)。这就是这套评测的抽样噪声量级。 所以 GDN 与 delta 的 0.0005、全注意力与 3:1 混合的 0.0089 都在噪声附近,不能据此排先后。 全注意力比全 linear 低 0.0365,比全 GDN 低 0.0289,幅度大得多,更可能是真的差距。
linear 和 delta 为什么越长越差?
点上面的 linear 看它的记忆板:读完 256 个字符后,第一层 head 0 左上 8×8 的最大绝对值是 215.6; delta 是 0.518,GDN 是 0.121。linear 只加不减,越长累积越多,512 时 loss 涨到 1.673。 delta 的写入会替换旧值,状态没有发散,但它没有遗忘门,没被重新写到的方向一直留着,512 时涨到 1.629。 加了 α 的 GDN 在 512 时是 1.532。
记忆板这几组为什么训得慢这么多?
RecurrentMixer.forward 用 Python for 循环逐 token 更新 S,每层每步要循环 64 次,全注意力是一次矩阵乘。 实测全注意力 493 秒,全 GDN 1603 秒。另外这 9 组训练和第 14 章的 3 组是在同一张 RTX 4090 上同时跑的, 互相抢 GPU,秒数只能看个大概。 真实实现(flash-linear-attention 等)把序列切块并行计算,和逐 token 循环数学上等价; 代码里写循环,是为了让"它是个 RNN"一眼看得出来。
↳ 跑法:python 11_linear_attn.py --preset attn / linear / delta / gdn / hybrid(默认 --task lm)
↳ 下一步:loss 只差百分之几,记忆板还省了一大截 cache。它输在哪?换一个要精确找回的任务。

⑤ 精确检索与 3:1 混合

键值检索任务:序列前半段是 N 对 (key, value),后半段把这 N 个 key 打乱再出一遍,每个 key 后面要答出它当初配的 value。 只有答案位置算 loss。点后半段任意一个 key,看它要去前面找哪一对:

一条 N = 8 的样本(页面按 get_batch 同样的规则现场生成;训练用 N = 32,序列长 128)
key(token 0–63) value(token 64–127,显示成 V00–V63) 要答的 value(只在这里算 loss)
真实四种层配方答对多少
32
竖线 = 随机猜的准确率 1/64 ≈ 1.56%(value 有 64 种)
同样 4 层、同样 32×32 的记忆板,只加不减的 linear 在 N=32 时答对 16.6%;带 delta rule 和遗忘门的 GDN 答对 99.9%,和全注意力的 100% 接近。
3:1 混合只有 88–93%,是不是说明混合不如纯 GDN?
这组数不能这么读,它没训完。切到"训练曲线"看:全注意力在 2000 → 2500 步之间 loss 从 2.60 骤降到 0.33, GDN 在 1500 → 2500 步从 1.77 降到 0.07;混合这组平台期拖到很后面,5000 步时还是 1.35,5500 步 0.65, 最后一次评估(5999 步)train 0.1475 / val 0.1462,仍在往下走。训练预算固定 6000 步、只跑一个种子, 它刚好在骤降的半路上停了。要比较混合和纯 GDN,需要训到 loss 都走平、再多跑几个种子;本页不给这个排序。
纯 GDN 这里也 99%,工业界为什么还要每 4 层留 1 层注意力?
要看条件。本实验最多记 N = 32 对,序列 128 个 token,每层 4 个头各有一块 32×32 的记忆板,要记的对数和 head_size 同一量级, 纯 GDN 装得下。真实模型的上下文是几万到几十万 token(Qwen3-Next 原生 262,144),要精确找回的内容远多于记忆板的维度, 固定大小的状态总有装不下的时候;全注意力层把每个 token 的 K/V 都留着,这类检索不受容量限制。 Zoology(Arora et al.,arXiv 2312.04927)对比注意力与多种高效架构,原文:"82% of the gap is explained by each model's ability to recall information that is previously mentioned in-context"。 Kimi Linear(arXiv 2510.26692)§5.2 的消融里,3 层 KDA 配 1 层 MLA 的 3:1 比例在质量与吞吐之间取舍最好。
linear 为什么只答对 6–17%?
它比随机猜(1.56%)高,但离检索还很远。两个原因叠在一起:S 只加不减,32 对写进去全部叠加; q、k 过了 elu + 1 全是正数,key 之间点积都大于 0,读任何一个 key 都会混进其他 value。 训练曲线上它的 loss 从 2.60 降到 2.47 就停住了,没有出现另外三组那样的骤降。
工业真实模型怎么排层

公开的线性注意力大模型都用混合排布:每 3 层固定状态的层配 1 层保留 K/V 的注意力层。选一个模型看它的 48 或 60 层怎么排:

Mamba2 和 Gated DeltaNet 是什么关系?
Gated DeltaNet 论文(Yang, Kautz, Hatamizadeh,arXiv 2412.06464,ICLR 2025)标题就是 "Improving Mamba2 with Delta Rule"。§2.1 把几种写法摆在一起:
模型记忆板更新
线性注意力S_t = S_{t-1} + v_t k_tᵀ
Mamba2S_t = α_t S_{t-1} + v_t k_tᵀ
DeltaNetS_t = S_{t-1}(I − β_t k_t k_tᵀ) + β_t v_t k_tᵀ
Gated DeltaNetS_t = S_{t-1}(α_t(I − β_t k_t k_tᵀ)) + β_t v_t k_tᵀ
Mamba2 有遗忘门 α、没有 delta rule;DeltaNet 有 delta rule、没有遗忘门;Gated DeltaNet 两个都要。 α 负责整块按比例清空,β 负责对单个 key 精确改写。Kimi Linear 的 KDA 在 Gated DeltaNet 基础上把门控做得更细。
为什么 delta 和 GDN 的 key 要做 L2 归一化?
delta rule 写完一次后,S' k = S k + β (v − S k) ‖k‖²。 ‖k‖ = 1、β = 1 时,S'k 正好等于 v,"读出 S·k"才等于"这个 key 当前存的值"。 如果 ‖k‖² = 4,同样的 β 会把差额放大 4 倍,读出值越过 v,反复写同一个 key 时误差按 |1 − β‖k‖²| = 3 倍往上翻,状态会发散。 代码里 F.normalize(q, dim=-1)、F.normalize(k, dim=-1) 两行就是为此;linear 模式不归一化,用 elu + 1。
短卷积在这里干什么用?
RecurrentMixer 在算 q/k/v 之前先过一个 kernel=4 的因果卷积,让每个位置的 q/k/v 看到自己和前 3 个 token。 检索任务里这一点很直接:读到 value 那一格时,应该以"前一个 token 的 key"为地址写入"这个 token 的 value"; 没有卷积,第 t 个位置的 k 只来自第 t 个 token 自己。 Gated DeltaNet 与 Qwen3-Next 都带这个卷积(Qwen3.8-Flash-Next 的 config 里 conv kernel 也是 4)。 代码对 linear / delta / gdn 三种写法都加了它,所以对照里的差别只来自写规则。
达标本关实测(RTX 4090 · CUDA · PyTorch 2.4.1 · 单种子)
↳ 代码:11_linear_attn.py 的 get_batch(recall 分支)· 跑法 python 11_linear_attn.py --preset gdn --task recall
🎉 线性注意力 · 通关
你已经看过记忆板的三种写法:只加的 linear、先擦再写的 delta、再加遗忘门的 Gated DeltaNet,也看到了它在长度外推上的优势和精确检索上的容量边界。下一关换一条路:K/V 全存,但每个 query 只读一部分。