第 12 章的注意力推理时要把每个 token 的 K/V 都存下来,上下文越长,
KV cache
越大。线性注意力
只留一块 head_size × head_size 的矩阵,每来一个 token 往里写一次,大小和上下文长度无关。
本关从最朴素的写法出发,加上 delta rule 和遗忘门,得到 Qwen3-Next 用的
Gated DeltaNet,
再用 9 次真实训练看它在语言模型和精确检索上各输赢在哪。配套代码 phase4-efficiency/11_linear_attn.py。
全注意力层每个 token 要存一份 K 和 V,上下文每长一个 token,cache 就多一份。 记忆板层每个头只存一块 head_size × head_size 的矩阵,读多长的文本都是这么大。 拖滑块改上下文长度,比较三种层配方要存多少个数:
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 的输入,量级很小,这里没有计入。
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 去读:
RecurrentMixer.forward(mode == "linear" 那一支)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 之间正交来避免串扰。
linear 模式里 q、k 先过 elu + 1,所有分量都是正数,两个全正向量的点积不会小于 0,更难做到正交。
只加不减的写法擦不掉旧值。delta rule 先读出这个 key 当前存的值,只补差额;Gated DeltaNet 再加一个遗忘门 α。 下面是一串三次写入:先把 A 写成旧值,再写 B,最后把 A 改成新值。然后分别用 kA、kB 去读。 切换三种写法,拖 β、α 看读出值怎么变:
代码里 α 的偏置初始化成 3.0,开训时 α = sigmoid(3) ≈ 0.95。下面是训练 5000 步后,第一层 4 个头在 val 文本上的平均 α, 画成"一条记忆过 n 个 token 后还剩几成"(按平均 α 粗算 αⁿ;实际 α 每个 token 各算一次):
RecurrentMixer(b_proj 算 β,a_proj 算 α)· 开关 --preset linear / delta / gdnS_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)。
a_proj 各算一次,个别位置可以接近 1;
② 只统计了第一层,代码没有导出后面三层的 α。开训时 α ≈ 0.95,训练把第一层压到了 0.55–0.66,
说明对字符级语言模型的第一层来说,忘得快的记忆板让 loss 更低。更长的依赖交给哪几层、α 是多少,本实验没有测。
同一份 tiny shakespeare、同一套超参(4 层、5000 步、batch 64、lr 1e-3、seed 1337),只换每层用什么: 全注意力、全 linear、全 delta、全 GDN、3 GDN + 1 注意力。训练长度 64,训完再拿 128 / 256 / 512 长的片段测 val loss。 切换指标,点一组看细节:
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,幅度大得多,更可能是真的差距。
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)键值检索任务:序列前半段是 N 对 (key, value),后半段把这 N 个 key 打乱再出一遍,每个 key 后面要答出它当初配的 value。 只有答案位置算 loss。点后半段任意一个 key,看它要去前面找哪一对:
get_batch 同样的规则现场生成;训练用 N = 32,序列长 128)elu + 1 全是正数,key 之间点积都大于 0,读任何一个 key 都会混进其他 value。
训练曲线上它的 loss 从 2.60 降到 2.47 就停住了,没有出现另外三组那样的骤降。
公开的线性注意力大模型都用混合排布:每 3 层固定状态的层配 1 层保留 K/V 的注意力层。选一个模型看它的 48 或 60 层怎么排:
| 模型 | 记忆板更新 |
|---|---|
| 线性注意力 | S_t = S_{t-1} + v_t k_tᵀ |
| Mamba2 | S_t = α_t S_{t-1} + v_t k_tᵀ |
| DeltaNet | S_t = S_{t-1}(I − β_t k_t k_tᵀ) + β_t v_t k_tᵀ |
| Gated DeltaNet | S_t = S_{t-1}(α_t(I − β_t k_t k_tᵀ)) + β_t v_t k_tᵀ |
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 三种写法都加了它,所以对照里的差别只来自写规则。
get_batch(recall 分支)· 跑法 python 11_linear_attn.py --preset gdn --task recall