GDN(Gated DeltaNet)详解
Gated DeltaNet —— Qwen3-Next 等新一代大模型中取代了大部分 self-attention 的线性注意力层。
本文从零推导 GDN 的两套等价实现(逐时间步递归 / 分块并行),讲清「delta rule 到底在解决什么问题、分块为什么能并行、训练时为什么必须用分块」,并给出与 transformers 官方实现逐位对齐的实测验证。文末附核心实现代码(纯 PyTorch,不依赖任何 CUDA kernel)。
一、数学原理
这一节把 GDN 的完整数学形式一次性摆出来:记号、状态递推、delta rule 的来历、门控,以及它为什么没法并行。后面各节再逐块展开实现细节。
1.1 记号与问题设定
设序列长度 T,每个时间步 t 提供三个向量:
| 符号 | 形状 | 含义 |
|---|---|---|
q_t |
R^dk |
query |
k_t |
R^dk |
key |
v_t |
R^dv |
value |
标准注意力的第 t 步输出是
$$ o_t = \sum_{s \le t} a_{t,s}\, v_s, \qquad a_{t,s} = \operatorname{softmax}\!\left( \frac{q_t \cdot k_s}{\sqrt{d_k}} \right) $$
每个 o_t 都要回头扫描全部 s ≤ t,所以存储 O(T)、计算 O(T²)。
1.2 线性注意力:把 softmax 拿掉
去掉归一化项,直接令 a_{t,s} = q_t · k_s:
$$ o_t = \sum_{s \le t} (q_t \cdot k_s)\, v_s = \Bigl( \sum_{s \le t} v_s k_s^\top \Bigr) q_t $$
括号里的量只依赖 s ≤ t,与 q_t 无关——把它记成一个矩阵:
$$ S_t := \sum_{s \le t} v_s k_s^\top \in \mathbb{R}^{d_k \times d_v} $$
S_t 就是状态矩阵:它把「截至时刻 t 的全部历史」压进了一个固定大小的量,于是
$$ o_t = S_t^\top q_t $$
存储和计算都变成 O(dk·dv),与 T 无关。这是线性注意力(以及所有 SSM)的出发点。
1.3 递推形式
把上面的求和写成增量式:
$$ S_t = S_{t-1} + v_t k_t^\top, \qquad o_t = S_t^\top q_t $$
每一步只做一次外积加法。这就是 3.1 节要讲的朴素线性注意力,也是理解 delta rule 的起点。
1.4 问题:这个递推只增不减
注意 v_t k_tᵀ 前面的符号永远是加号。同一批 key 反复出现时,旧关联不会被冲掉,只会层层叠加:
t=0: (k̂, v=1) → S = k̂ ⊗ 1
t=1: (k̂, v=2) → S = k̂ ⊗ 3 ← 期望读到 2,实际读到 3
S 一旦被写满就永久污染。这正是早期线性注意力在检索类任务上打不过标准注意力的原因。要修掉它,必须给递推加上「减法」。
1.5 Delta rule:来自在线最小二乘的梯度下降
GDN 的数学出发点是一句很朴素的话:把 S 看成一个待学习的线性映射,让 S k_t 尽量等于 v_t。
对第 t 个时间步定义平方损失
$$ L_t(S) = \frac{1}{2} \| S k_t - v_t \|^2 $$
(相当于把 S 的 dv 行各看成一个独立的线性回归,样本是 (k_t, v_t)。)对 S 求梯度:
$$ \nabla_S L_t = (S k_t - v_t)\, k_t^\top $$
取一步梯度下降,步长 β_t:
$$ \begin{aligned} S_t &= S_{t-1} - \beta_t \nabla_S L_t \\ &= S_{t-1} + \beta_t (v_t - S_{t-1} k_t) k_t^\top \end{aligned} $$
这就是 delta rule(Widrow-Hoff 规则)。括号里的
$$ \delta_t := v_t - S_{t-1} k_t $$
是预测误差——状态对当前 key 的预测与真实 value 差多少。整步更新的含义是「只在预测错的那部分上做修正」。
把它整理成矩阵形式,几何意义更清楚:
$$ S_t = S_{t-1} \bigl( I - \beta_t k_t k_t^\top \bigr) + \beta_t v_t k_t^\top $$
k_t 归一化后(‖k_t‖ = 1),k_t k_tᵀ 是往 k_t 方向的投影算子,I − k_t k_tᵀ 则是把该方向分量抹掉的投影。于是上式读作:
先把状态里关于
k_t的那份记忆擦掉β_t比例,再写入新的v_t。
验证一步。取 β_t = 1,两边右乘 k_t:
$$ \begin{aligned} S_t k_t &= S_{t-1} k_t - S_{t-1} k_t (k_t^\top k_t) + v_t (k_t^\top k_t) \\ &= v_t \end{aligned} $$
S_t k_t 恰好等于 v_t——新的关联被精确写入,旧的被彻底擦除。3.4 节那个 T=2 的最小例子验证的就是这一步。
1.6 门控:再加一个时间上的遗忘
上面的 S 仍会无限累积:β 只控制单次写入的强度,不负责让「过于久远的信息」整体退场。再引入标量门控
$$ \alpha_t = \exp(g_t) \in (0,\, 1] $$
把整个状态按时间指数衰减:
$$ S_t = \alpha_t S_{t-1} + k_t \otimes \Bigl( \beta_t \bigl( v_t - \alpha_t S_{t-1} k_t \bigr) \Bigr) $$
注意 α_t 出现了两次,且 S_{t-1} 要先乘 α_t、再拿去算误差(4.2 节会说明这个顺序为什么不能反)。
g_t 取对数域有个实际好处:连续衰减等于指数相加,
$$ \alpha_{t_1} \cdot \alpha_{t_2} = \exp(g_{t_1} + g_{t_2}) $$
第五节那套分块推导整个建立在这条性质上——没有它,块内的累积衰减就没法写成一次 cumsum。
1.7 和在线梯度下降对照
把 GDN 和在线梯度下降摆在一起,对应关系一目了然:
| 在线梯度下降 | GDN |
|---|---|
待学参数 W |
状态矩阵 S_t |
学习率 η |
写入强度 β_t |
第 t 个样本 (x_t, y_t) |
(k_t, v_t) |
目标 ‖W x_t − y_t‖² |
‖S k_t − v_t‖² |
一步更新 W ← W − η(W x_t − y_t) x_tᵀ |
S ← S + β_t (v_t − S k_t) k_tᵀ |
| —(无对应项) | 门控衰减 α_t = exp(g_t) |
GDN 的整个前向过程,就是对一个不断到来的 (k, v) 序列做在线最小二乘。 门控是额外加上的「遗忘」,让有限大小的状态不至于被过于久远的信息占满。
1.8 剩下的问题:这个式子没法并行
上面每一步都在时间维串行:S_t 依赖 S_{t-1},S_{t-1} 又依赖 S_{t-2}……训练时序列长度 T 就是串行步数,GPU 完全跑不满(2.4 节实测慢了 25 倍)。
所以真正要解决的是:在保持上面这套语义的前提下,把串行循环改写成矩阵乘。这是第五节的主题,也是 GDN 能上生产的关键。
二、整体架构总览
2.1 递归形式的数据流
GDN 的核心是一个状态矩阵 S 的递推。下面是 gdn.py:234-337 的前向数据流:
query (B,T,H,dk) key (B,T,H,dk) value (B,T,H,dv) g (B,T,H) beta (B,T,H)
│ │ │ │ │
l2norm l2norm │ exp() sigmoid()
▼ ▼ ▼ ▼ ▼
q̂ (单位向量) k̂ (单位向量) v α = e^g ≤ 1 β ∈ (0,1)
│ │ │ 「衰减」 「写入强度」
│ × 1/√dk │ │ │ │
└────────┬─────────┴────────┬─────────┴────────┬─────────┴───────────────┘
│ │ │
▼ ▼ ▼
┌────────────────────────────────────────────────────────────┐
│ for t = 0 … T-1 ← 时间维串行,这就是慢的原因 │
│ │
│ S ← α_t · S ① 遗忘 │
│ kv_mem ← S k_t ② 回忆 │
│ δ ← β_t (v_t − kv_mem) ③ 误差 │
│ S ← S + k_t ⊗ δ ④ 写入 │
│ o_t ← Sᵀ q_t ⑤ 读出 │
└────────────────────────────────────────────────────────────┘
│
▼
o (B,T,H,dv) + final_state S (B,H,dk,dv)
核心思想一句话:把「历史」压缩成一个固定大小的状态矩阵 S ∈ R^(dk×dv),每一步用 delta rule(先擦除旧关联、再写入新关联)更新它,而不是像标准注意力那样把全部 K/V 都存下来。于是推理成本与上下文长度解耦。
2.2 关键张量形状(gdn.py:234-243)
| 张量 | 形状 | 含义 |
|---|---|---|
query / key |
(B, T, H, dk) |
查询 / 键 |
value |
(B, T, H, dv) |
值 |
g |
(B, T, H) |
对数衰减,恒 ≤ 0 |
beta |
(B, T, H) |
写入强度,通常在 (0,1) |
S(状态) |
(B, H, dk, dv) |
与 T 无关,这是全部价值的来源 |
o(输出) |
(B, T, H, dv) |
2.3 两套实现,一个数学
| 函数 | 文件位置 | 串行步数 | 用途 |
|---|---|---|---|
recurrent_gated_delta_rule |
gdn.py:234 |
T | 最好懂;也是分块版的裁判 |
chunk_gated_delta_rule |
gdn.py:499 |
T/C | Qwen3-Next 训练实际走的路径 |
两套数学完全等价。递归版是「照数学定义直译」,几乎不可能写错;分块版是为了提速重写的等价形式,所以递归版天然就是分块版的测试基准。
2.4 为什么还需要分块版
递归版每一步只做几个 O(dk·dv) 的小运算,在 GPU 上单次 kernel 启动开销远大于计算本身。T=4096 就要发 4096 次 kernel,GPU 全程在等。实测(RTX 3050 Ti Laptop):
[T=2048, dk=dv=64, H=4]
递归版 (逐时间步) : 430.0 ms 2048 次串行 kernel
分块版 (C=64) : 17.2 ms 32 次串行 kernel 25.1x
所以分块不是为了「算术量更少」(其实两者 FLOPs 同量级),而是为了把串行步数从 T 降到 T/C,让 GPU 有活干。
三、从线性注意力到 Delta Rule
3.1 线性注意力:把历史压成一个矩阵
标准注意力要保存全部历史 K/V,显存随上下文线性增长,计算随长度平方增长。线性注意力换了个思路:把历史压成一个矩阵 S ∈ R^(dk×dv),用递推维护:
$$ S_t = S_{t-1} + v_t k_t^\top, \qquad o_t = S_t^\top q_t $$
展开就是 S_T = Σ_t v_t k_tᵀ —— 等价于把每条 (k, v) 关联原样堆进去。这个形式下 o_t = S_tᵀ q_t = Σ_{s≤t} (q_t · k_s) v_s,把 softmax 抽掉后剩下的「未归一化注意力」。
对应代码(gdn.py:326-330):
state = state + k_t[..., :, None] * delta[..., None, :] # 写入
o[:, :, t] = (state * q_t[..., :, None]).sum(dim=-2) # 读出
3.2 问题:只增不减
这个递推有个硬伤:没有任何机制删除旧内容。
同一个 key 反复出现时,旧值不会消失,只会不断往上叠。比如先存 (k, v=1) 再存 (k, v=2),读出来是 3 而不是 2。换句话说,状态会被历史「填满」并永久污染——这正是早期线性注意力/SSM 在检索类任务上打不过标准注意力的主要原因。
3.3 Delta rule:先擦除,再写入
GDN 借鉴在线梯度下降(Widrow-Hoff / LMS)的思路:写入前先看看状态对当前 key 已经记了什么,只把误差补上去。
S_t = S_{t-1}(I − β_t k_t k_tᵀ) + β_t v_t k_tᵀ
= S_{t-1} + β_t (v_t − S_{t-1} k_t) k_tᵀ
└────────┬─────────┘
预测误差 δ
关键在 k_t 经过 L2 归一化,于是 k_t k_tᵀ 是一个投影算子:(I − k kᵀ) 会把任意向量中平行于 k 的分量抹掉。所以 S(I − β k kᵀ) 的语义就是「把状态里关于 k 的那份记忆擦掉 β 比例」。
β_t ∈ (0,1) 是写入强度,作用完全等价于学习率:
| β | 行为 |
|---|---|
| β = 1 | 完全替换旧值(标准 delta rule) |
| β = 0 | 完全不写入(状态冻结) |
| 0 < β < 1 | 新旧值按比例混合 |
对应代码(gdn.py:316-328):
kv_mem = (state * k_t[..., :, None]).sum(dim=-2) # ② 回忆:S k_t
delta = (v_t - kv_mem) * beta_t[..., None] # ③ 误差
state = state + k_t[..., :, None] * delta[..., None, :] # ④ 写入:S + k ⊗ δ
3.4 最小例子:T=2 看清 delta rule 和朴素累加的区别
这是理解 delta rule 最快的方式。设 g = 0(不衰减)、β = 1(完全写入)、q = k̂(读出同一个方向),两个时间步用同一个 key、不同的 value:
| 时刻 | 朴素线性注意力 | Delta rule |
|---|---|---|
| t=0 后 | S = k̂ ⊗ v₀ |
S = k̂ ⊗ v₀ |
| t=1 后 | S = k̂ ⊗ (v₀ + v₁) |
S = k̂ ⊗ v₁ |
读出 o₁ |
(v₀ + v₁)/√dk ← 旧值残留 |
v₁/√dk ← 旧值被擦掉 |
实测(gdn.py 演示脚本第 5 节):
delta rule 读出 o[1] 与 v₁/√d 的差 : 4.843e-08 ← 吻合
朴素累加会读出 (v₀+v₁)/√d,与 v₁/√d 的差 : 3.450e-01 ← 明显不同
这是 GDN 相比「Mamba2 + 纯累加」最关键的一步。两套代码长得几乎一模一样,只差一个
S k项——这正是它容易被写错、也最需要单独测试的地方。
四、门控衰减 α
4.1 公式
再引入一个标量门控 α_t = exp(g_t) ≤ 1,让整个状态按时间指数遗忘:
$$ S_t = \alpha_t S_{t-1} + k_t \otimes \Bigl( \beta_t \bigl( v_t - \alpha_t S_{t-1} k_t \bigr) \Bigr) $$
g 取对数域的好处是累加即相乘(α₁α₂ = exp(g₁+g₂)),分块推导里会大量用到。
4.2 ⚠️ 一个容易写错的顺序问题
α 必须先乘到 S 上,再用衰减之后的 S 去算预测误差。 顺序反了数值就对不上:
# ✅ 正确(官方实现,也是本仓库)
state = state * alpha_t[..., None, None] # 先遗忘
kv_mem = (state * k_t[..., :, None]).sum(dim=-2) # 再基于「已遗忘」的状态回忆
delta = (v_t - kv_mem) * beta_t[..., None]
# ❌ 错误
kv_mem = (state * k_t[..., :, None]).sum(dim=-2) # 先回忆(基于未遗忘的状态)
state = state * alpha_t[..., None, None] # 再遗忘
delta = (v_t - kv_mem) * beta_t[..., None]
直觉:衰减代表「遗忘」,那「我还记得什么」就必须基于遗忘之后的状态判断。否则会误以为某些旧信息还在,从而少写一部分新内容。
这个细节自洽性测试(递归 vs 分块)永远发现不了——两套实现同时写错的话,它们会一致地错。必须靠和官方实现对比才能抓到,见第八节。
五、分块并行:把串行循环变成矩阵乘
5.1 思路
序列 T ──┬── chunk 0 ──┬── chunk 1 ──┬── … ──┬── chunk nc-1 ──┬──
│ C 个位置 │ │ │ │
▼ ▼ ▼ ▼ ▼
块内:稠密矩阵乘 ── 全部可并行
块间:只有状态 S 串行传递 ── 串行步数 T/C
5.2 块内:WY 表示(Neumann 级数)
只看一个块。设块内累积对数衰减 G_i = Σ_{s≤i} g_s,位置 i 的状态可展开为
$$ S_i = \sum_{j \le i} \exp(G_i - G_j)\, k_j \otimes \delta_j, \qquad \delta_j = \beta_j (v_j - S_{j-1} k_j) $$
麻烦在 δ_j 自己又依赖 S_{j-1},直接展开是个嵌套递推。但写成矩阵形式后会发现,块内所有 δ 构成一个单位下三角线性方程组:
$$ (I - L)\, \delta = \mathrm{rhs}, \qquad L[i, j] = -\beta_i (k_i \cdot k_j) \exp(G_i - G_j) \quad (j < i) $$
所以 δ = (I − L)⁻¹ · rhs。用 Neumann 级数展开:
$$ (I - L)^{-1} = I + L + L^2 + L^3 + \cdots $$
因为 L 严格下三角(L^C = 0),这个级数是有限的,逐行前代就能精确算出来。结果记作
$$ T = I + L + L^2 + \cdots + L^{C-1} $$
这个 T 就是 WY 表示(名字来自 Woodbury 恒等式)。它在块内扮演「因果关系修正器」:让每个位置只减掉它前面那些位置已经写进状态的东西。
对应代码(gdn.py:411-420):
# L[i, j] = -β_i (k_i · k_j) exp(G_i - G_j),j < i;上三角与对角线清零
L = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(upper_and_diag, 0)
# 逐行前代:把 L 就地累加成 L + L² + …
for i in range(1, C):
row = L[..., i, :i].clone()
sub = L[..., :i, :i].clone()
L[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
return L + torch.eye(C, dtype=L.dtype, device=L.device)
第
i行只读下标< i的行,所以行与行之间没有循环依赖——这既是它能顺序算完的原因,也是因果性在数值上严格成立的原因(见 8.3)。
拿到 T 之后,块内所有位置一次性矩阵乘完(gdn.py:569-572):
wy = _wy_representation(k_beta, key, decay_mask, chunk_size)
w = wy @ v_beta # 块内修正后的 value
u = wy @ (k_beta * g_cum.exp().unsqueeze(-1)) # 块内修正后的 key
5.3 块间状态扫描
_chunk_state_pass(gdn.py:467-495)串行遍历 nc 个块,每块内部是稠密矩阵乘:
intra = (q_i @ k_i.transpose(-1, -2) * decay_i).masked_fill(strict_upper, 0)
v_prime = u_i @ state # 块外状态里已存的部分
w_new = w_i - v_prime # ← delta 的「擦除」发生在块级别
inter = (q_i * g_i.exp().unsqueeze(-1)) @ state
o[:, :, i] = inter + intra @ w_new
state = state * g_i[:, :, -1, None, None].exp() + (
k_i * (g_i[:, :, -1, None] - g_i).exp().unsqueeze(-1)
).transpose(-1, -2) @ w_new
对应关系:
| 块内符号 | 递归版里的对应物 |
|---|---|
intra |
块内位置间的注意力(含衰减 exp(G_i − G_j)) |
v_prime |
「回忆」S k,只不过在块级别一次算完 |
w_new |
「误差」v − S k |
inter |
「读出」Sᵀ q |
state 更新 |
「遗忘 + 写入」 |
5.4 两个辅助函数的职责
| 函数 | 位置 | 作用 |
|---|---|---|
_chunk_decay_mask |
gdn.py:350 |
算块内累积衰减 G 和衰减掩码 exp(G_i − G_j) |
_wy_representation |
gdn.py:378 |
前代求解 (I − L)⁻¹,得到 WY 表示 |
_chunk_state_pass |
gdn.py:423 |
块间状态扫描 |
_chunk_decay_mask 里有个细节值得一提(gdn.py:365-370):
diff = g_cum.unsqueeze(-1) - g_cum.unsqueeze(-2)
decay_mask = diff.tril().exp() # 必须先 tril 再 exp
decay_mask = decay_mask.tril() # 第一次 tril 让上三角变成 0,exp(0)=1,还要再清一次
顺序不能反:上三角处 G_i − G_j > 0,直接 exp 可能溢出成 inf,先置零才安全。
六、Qwen3-Next 的参数化细节
GDN 是个通用的算子,「g 和 β 从哪来」由具体模型决定。Qwen3-Next 的做法:
6.1 g 和 β 的构造
beta = sigmoid(b) # b = in_proj_ba(hidden) 的前半
g = -exp(A_log) * softplus(a + dt_bias) # a = in_proj_ba(hidden) 的后半
A_log和dt_bias是可学习参数,借用了 Mamba2 里 SSM 离散化的写法;- 负号 +
softplus(恒正)保证g ≤ 0,即α = exp(g) ≤ 1—— 状态只会遗忘,不会指数爆炸; - 每个 value head 有各自独立的
A和dt_bias,不同头可以学到不同的遗忘速率(有的头记长程,有的头只看近期)。
6.2 L2 归一化与缩放
if use_qk_l2norm_in_kernel:
query = l2norm(query, dim=-1, eps=1e-6)
key = l2norm(key, dim=-1, eps=1e-6)
query = query * (k_head_dim ** -0.5) # 只缩放 q
q、k都归一化到单位球面 —— 只有k是单位向量,(I − k kᵀ)才是严格意义的投影算子,delta rule 的几何解释才成立;- 只缩放 q,不缩放 k。因为 k 要保持单位长度,缩放会破坏投影性质;
l2norm的eps加在平方和上而不是模长上:
inv_norm = torch.rsqrt((x * x).sum(dim=-1, keepdim=True) + eps) # ✅
# x / (x.norm(dim=-1, keepdim=True) + eps) # ❌ 数值上会差一点
这是为了和 FLA 库的 CUDA kernel 逐位对齐。看着是吹毛求疵,但它决定了第八节的「0.000e+00」能不能成立。
6.3 GQA
Qwen3-Next 的 GDN 层是 32 个 value head、16 个 key head,每个 key head 被 2 个 value head 共享。q/k 沿 head 维展开:
q = q.repeat_interleave(2, dim=2)
k = k.repeat_interleave(2, dim=2)
注意必须是 repeat_interleave:
repeat_interleave -> [h0 h0 h1 h1 h2 h2 ...] ✅ 官方用法
tile / repeat -> [h0 h1 h2 ... h0 h1 h2 ...] ❌ 结果不同
这是个不报错、不崩溃、只是精度悄悄变差的典型错误。
6.4 短卷积
GDN 之前还有一层 kernel_size = 4 的 depthwise causal conv1d,在 q/k/v 拼接后的 conv_dim 上做。作用是给每个位置提供局部上下文——纯线性注意力对「紧邻的几个 token」这种局部模式不敏感,短卷积补上了这一块。
6.5 3:1 混合结构
Qwen3-Next 不是纯 GDN,而是 full_attention_interval = 4:
层 0 层 1 层 2 层 3 层 4 层 5 层 6 层 7 ...
GDN GDN GDN 注意力 GDN GDN GDN 注意力
即 每 4 层里 3 层 GDN + 1 层标准注意力。GDN 负责廉价地处理长上下文,注意力层负责精确的全局检索(精确复制、长距离指代这类 GDN 不擅长的任务)。
七、复杂度与显存对比
7.1 计算复杂度
| 形式 | 计算量 | 串行步数 | 训练可行性 |
|---|---|---|---|
| 递归(逐时间步) | O(T · dk · dv) |
T | 太慢 |
| 分块并行(C=64) | O(T·C·(dk+dv) + T·dk·dv) |
T/C | ✅ Qwen3-Next 走这条 |
| 朴素二次(单块) | O(T² · (dk+dv)) |
1 | 只在 T 很小时可用 |
有意思的是递归和分块的 FLOPs 同量级,分块赢的是串行步数。
把 chunk_size ≥ T 传进去,分块实现就退化成最后那行的二次形式——这可以拿来直观感受「分块到底省了什么」(gdn.py:674 有对应测试)。
7.2 显存对比:GDN 的真正卖点
Qwen3-Next 全注意力层的配置是 num_attention_heads=16、num_key_value_heads=2、head_dim=256(8:1 的 GQA)。
单层、单条序列的存储开销:
| 层类型 | 存什么 | T = 1 | T = 4096 |
|---|---|---|---|
| 标准注意力 | KV cache | 2 KiB | 8 MiB |
| GDN | 状态矩阵 S |
1 MiB | 1 MiB(不变) |
(按 bf16 计算。全注意力 KV:2 × 2 heads × 256 dim × 2 B = 2 KiB/token;GDN 状态:32 × 128 × 128 × 2 B = 1 MiB)
交叉点约在 512 个 token:短于此时注意力更省,长于此时 GDN 更省。到 4K 上下文,GDN 比同层注意力便宜 8 倍,而且这个倍数随上下文继续线性增长。
这才是 GDN 的意义:不是为了在短序列上打败注意力,而是让超长上下文的推理成本变得可接受。
八、正确性验证
分块 GDN 有几类 bug 长得几乎一样但坏在不同地方。本仓库用 29 个用例分四道防线,下面是实测结果。
8.1 第一道:两套实现互相印证
递归版是「照数学定义直译」的,几乎不可能写错;它能当分块版的裁判。
chunk_size 输出 max|Δ| 状态 max|Δ|
8 4.470e-08 2.384e-07
16 6.706e-08 6.557e-07
32 2.161e-07 1.907e-06
64 3.129e-07 1.848e-06
128 4.992e-07 4.232e-06
256 4.992e-07 4.232e-06
误差在 fp32 的累积舍入量级(~1e-7),说明分块那套矩阵推导是对的;且 chunk_size 从 8 变到 256 结果一致,说明它只是实现细节。
8.2 第二道:与官方实现逐位对齐
第一道防线有个盲区:两套实现同时错同一个地方,它发现不了。比如第四节那个「α 先乘还是后乘」的顺序问题——如果两套都写反了,自洽性测试会一致通过。
所以必须有个外部标准。对照 transformers 里 Qwen3-Next 的官方实现:
递归版 vs torch_recurrent_gated_delta_rule : max|Δ| = 0.000e+00
分块版 (C=16) vs torch_chunk_gated_delta_rule : max|Δ| = 0.000e+00
分块版 (C=64) vs torch_chunk_gated_delta_rule : max|Δ| = 0.000e+00
分块版 (C=128) vs torch_chunk_gated_delta_rule : max|Δ| = 0.000e+00
差值为 0 表示逐位完全相同。 这不是巧合:运算顺序、float32 升位策略、l2norm 的 eps 位置都刻意对齐了。测试还覆盖了 T ∈ {1, 7, 63, 64, 65}——专门盯骑在 chunk 边界上的 off-by-one。
8.3 第三道:因果性
past_diff = (o1[:, :split+1] - o2[:, :split+1]).abs().max().item()
assert past_diff == 0.0
在 t=40 之后替换全部输入,t ≤ 40 的输出必须逐位不变(实测 0.000e+00)。断言用的是精确相等而不是近似——因为分块版块内是稠密矩阵乘,掩码写错就会让未来信息漏进来。
这是最阴险的一类 bug:未来泄漏不会崩、不会 NaN,loss 曲线看着完全正常,甚至更低。只有推理时才暴露。
8.4 第四道:性质测试
| 用例 | 抓什么 bug |
|---|---|
test_delta_rule_erases_old_value |
只实现了线性注意力,没实现 delta rule |
test_beta_zero_writes_nothing |
公式里 β 乘的位置错 |
test_gqa_repeat_kv_heads_is_interleave |
用了 tile 而不是 interleave |
test_prefill_decode_consistency |
final_state 语义错(多衰减/少衰减一次) |
test_initial_state_equals_prefix |
initial_state 进入公式的方式错 |
test_strong_decay_forgets_the_past |
门控方向反了(g 符号错),变成「不遗忘」 |
test_long_sequence_stability |
T=2048 溢出 NaN/Inf |
test_gradients_match_between_implementations |
前向对但反向路径错 |
其中 test_delta_rule_erases_old_value 价值最高:用 3.4 节那个最小例子直接验数值,还带了个反例断言(确认结果不等于 v₀+v₁),防止测试本身没有鉴别力。
8.5 性能实测
[T=2048, dk=dv=64, H=4, device=cuda] (RTX 3050 Ti Laptop)
递归版 (逐时间步) : 430.0 ms
分块版 (C=64) : 17.2 ms 加速比 25.1x
九、代价与局限(诚实地说)
1. 固定大小的状态是有损压缩。 S 只有 dk × dv 这么大,塞不下无限多的历史。精确复制、长距离指代这类需要「原样取回某个 token」的任务,GDN 天然弱于标准注意力。这正是 Qwen3-Next 保留 1/4 全注意力层的原因——GDN 不是注意力的替代品,而是它的廉价补充。
2. 短序列上不划算。 第 7.2 节的交叉点在 ~512 token。上下文短于此时,GDN 的状态反而比 KV cache 占更多显存。
3. 本实现是教学版,不是生产版。 分块版会显式物化 [b, h, T, C] 量级的中间张量(decay_mask、WY 矩阵 T)。以 b=1, h=32, T=4096, C=64 为例,decay_mask 就是 1×32×4096×64×4B = 32 MB,WY 矩阵同理。生产的 FLA 实现用 Triton kernel 做分块融合,避免物化这些中间量。本仓库的 25x 加速是相对自己的递归版,离生产 kernel 还有很大差距。
4. 本仓库只做 core 算法。 距一个能跑的完整层还差:in_proj_qkvz / in_proj_ba 投影、短卷积、A_log + dt_bias 参数化、Gated RMSNorm、out_proj,以及增量解码缓存。initial_state / output_final_state 这套接口已经为缓存预留好了。
5. chunk_size 是需要调的。 太大则块内 O(C²) 主导,太小则串行步数多、并行度不足。官方默认 64。
十、一句话总结
| GDN 做了什么 | 结果 |
|---|---|
| 把历史压成固定大小的状态矩阵 | 推理显存与上下文长度解耦 |
| 用 delta rule(先擦再写)替代纯累加 | 同一个 key 重复出现能覆盖而非污染 |
加标量门控 α = exp(g) |
可以学习「记多久」,且保证数值稳定 |
| 用 WY 表示把块内串行变矩阵乘 | 训练串行步数从 T 降到 T/C |
| 和标准注意力 3:1 混用 | 长上下文成本大降,精确检索能力不丢 |
附录:gdn.py 完整代码
以下是 gdn.py 的全部内容(1217 行):模块 docstring 里那份数学推导、工具函数、
两套实现、29 个 pytest 用例,以及演示脚本 main()。
"""
Gated DeltaNet (GDN) 从零手写 —— 单文件版
================================================================================
Qwen3-Next 里取代了大部分 self-attention 的那个线性注意力层。
纯 PyTorch 实现,不依赖 fla / triton 的任何 CUDA kernel,CPU 或 GPU 都能跑。
包含两套等价实现:
* ``recurrent_gated_delta_rule`` 逐时间步递归 —— 最好懂,训练时太慢
* ``chunk_gated_delta_rule`` 分块并行 —— Qwen3-Next 训练实际走的路径
以及 30 个 pytest 用例(含与 transformers 官方实现的逐位对比)。
运行方式::
# 完整演示 + 自检(推荐先跑这个)
python gdn.py
# 跑 pytest 用例
python -m pytest gdn.py -v
================================================================================
一、从线性注意力说起
================================================================================
标准注意力要保存全部历史 K/V,显存和计算随序列长度平方增长。线性注意力把历史
压成一个矩阵 ``S ∈ R^{d_k × d_v}``,用递推方式维护:
S_t = S_{t-1} + v_t k_tᵀ o_t = S_tᵀ q_t
展开就是 ``S_T = Σ_t v_t k_tᵀ``,等价于把每条 (k, v) 关联原样堆进去。
**问题**:只增不减。同一个 key 反复出现时,旧值永远不会消失,只会不断往上叠。
比如先存 (k, v=1) 再存 (k, v=2),读出来是 3 而不是 2。
================================================================================
二、Delta rule:先擦除,再写入
================================================================================
借鉴在线梯度下降(Widrow-Hoff / LMS):**写入前先看看状态对当前 key 已经记了什么,
只把误差补上去**。
S_t = S_{t-1}(I - β_t k_t k_tᵀ) + β_t v_t k_tᵀ
= S_{t-1} + β_t (v_t - S_{t-1} k_t) k_tᵀ
└────────┬─────────┘
预测误差 delta
* ``k_t`` 经过 **L2 归一化**,所以 ``k_t k_tᵀ`` 是投影算子,
``(I - k kᵀ)`` 把向量中平行于 k 的分量抹掉;
* ``β_t ∈ (0,1)`` 是写入强度,作用等价于学习率:
β=1 完全替换,β=0 完全不写。
回到上面的例子:β=1 时第二次写入后状态恰好是 ``k ⊗ v₂``,旧值被干净地擦掉了。
================================================================================
三、加门控衰减 α
================================================================================
再加一个标量门控 ``α_t = exp(g_t) ≤ 1``,让状态按时间指数遗忘:
S_t = α_t · S_{t-1} + k_t ⊗ ( β_t (v_t - α_t · S_{t-1} k_t) )
o_t = S_tᵀ q_t
⚠️ **最容易写错的细节**:α 必须**先乘到 S 上**,再用**衰减之后**的 S 去算预测误差。
官方实现就是这个顺序,顺序反了数值对不上。
直觉:衰减代表"遗忘",那"我还记得什么"就必须基于遗忘之后的状态判断。否则会误以为
某些旧信息还在,从而少写一部分新内容。
================================================================================
四、分块并行:把串行循环变成矩阵乘
================================================================================
**为什么需要**:递归版在时间维串行,T=4096 就要发 4096 次 kernel,GPU 全程在等,
实测比递归版慢 23 倍。
**思路**:把序列切成 C 大小的块,块内一次性矩阵乘完(可并行),块间只串行传递状态 S。
串行长度从 T 降到 T/C。
**块内怎么变成矩阵乘?**
设块内累积对数衰减 ``G_i = Σ_{s≤i} g_s``,位置 i 的状态可写成
S_i = Σ_{j≤i} exp(G_i - G_j) · k_j ⊗ δ_j , δ_j = β_j (v_j - S_{j-1} k_j)
麻烦在于 ``δ_j`` 自己又依赖 ``S_{j-1}``。但写成矩阵形式后会发现,块内所有 δ 构成一个
**单位下三角线性方程组**:
(I - L) · δ = rhs , 其中 L[i, j] = -β_i (k_i · k_j) exp(G_i - G_j) ( j < i )
所以 ``δ = (I - L)⁻¹ · rhs``。用 **Neumann 级数**:
(I - L)⁻¹ = I + L + L² + L³ + …
因为 L 严格下三角(``L^C = 0``),这个级数是**有限**的,逐行前代就能精确算出来。
这个 ``T = I + L + L² + … + L^{C-1}`` 就叫 **WY 表示**(名字来自 Woodbury 恒等式)。
它在块内扮演"因果关系修正器":让每个位置只减掉**它前面**那些位置已经写进状态的东西。
拿到 T 之后,块内所有位置一次性矩阵乘完:
w = T @ (β ⊙ v) 块内修正后的 value
u = T @ (β ⊙ k ⊙ exp(G)) 块内修正后的 key
**块间状态扫描**,对第 i 个块:
intra = (Q Kᵀ) ⊙ decay_mask 块内因果注意力(下三角掩码)
w_new = w - u @ S 减掉块外状态里已存的部分 ← delta 的"擦除"
inter = (Q ⊙ exp(G)) @ S 从块外状态直接读出
O = inter + intra @ w_new
S ← S · exp(G_C) + (K ⊙ exp(G_C - G))ᵀ @ w_new 衰减到块末 + 写入本块
================================================================================
五、复杂度对比
================================================================================
形式 计算量 串行步数 推理状态大小
--------------------------------------------------------------------------
递归(逐时间步) O(T · d_k · d_v) T O(d_k·d_v)
分块并行(C=64) O(T·C·(d_k+d_v) + T·d_k·d_v) T/C O(d_k·d_v)
朴素二次(单块) O(T² · (d_k+d_v)) 1 O(T²)
关键:分块版算术量只线性增长,**串行步数降到 T/C**,而状态大小始终与 T 无关 ——
这就是 GDN 能处理超长上下文的根本原因。
================================================================================
六、Qwen3-Next 的参数化细节
================================================================================
beta = sigmoid(b) # b = in_proj_ba(hidden) 的前半
g = -exp(A_log) * softplus(a + dt_bias) # a = in_proj_ba(hidden) 的后半
``A_log`` 和 ``dt_bias`` 是可学习参数,借用了 Mamba2 里 SSM 离散化的写法,保证 g ≤ 0
(衰减系数 exp(g) ≤ 1,状态只会遗忘不会爆炸)。每个 value head 有各自独立的 A 和
dt_bias,不同头可以学到不同的遗忘速率。
其它要点:
* **q/k 的 L2 归一化在 kernel 内做**(``use_qk_l2norm_in_kernel=True``),且 q 额外乘
``1/sqrt(d_k)``。**只缩放 q,不缩放 k** —— k 要保持单位长度,``(I - k kᵀ)`` 才是投影算子。
* **GQA**:``num_v_heads=32``、``num_k_heads=16``,q/k 沿 head 维用 ``repeat_interleave``
展开成 32 个 head。注意是 ``[h0 h0 h1 h1 ...]``(interleave),**不是**
``[h0 h1 ... h0 h1 ...]``(tile),两者结果不同。
* **短卷积**:GDN 之前还有一层 ``kernel_size=4`` 的 depthwise causal conv1d
(在完整层里,本文件只做 core 算法)。
* 输出再经过 **Gated RMSNorm**(用 z 做门控)和 ``out_proj``。
默认配置:``hidden_size=2048``、``linear_num_key_heads=16``、
``linear_num_value_heads=32``、``linear_key_head_dim=128``、
``linear_value_head_dim=128``、``linear_conv_kernel_dim=4``。
================================================================================
七、本文件**没有**包含的部分
================================================================================
1. ``Qwen3NextGatedDeltaNet`` 整层:in_proj_qkvz / in_proj_ba 投影、短卷积 conv1d、
A_log + dt_bias 门控参数化、Qwen3NextRMSNormGated、out_proj。
2. 增量解码缓存:conv_state + recurrent_state(torch_causal_conv1d_update)。
3. 玩具任务训练,对比 GDN 层 vs 标准注意力层。
``initial_state`` / ``output_final_state`` 这套接口已经为第 2 项预留好了。
================================================================================
八、参考资料
================================================================================
* 论文:Gated Delta Networks: Improving Mamba2 with Delta Rule (Yang et al., ICLR 2025)
* flash-linear-attention: https://github.com/fla-org/flash-linear-attention
* transformers 源码:transformers/models/qwen3_next/modeling_qwen3_next.py
(本文件的数值对齐金标准,看 torch_chunk_gated_delta_rule /
torch_recurrent_gated_delta_rule / Qwen3NextGatedDeltaNet)
"""
from __future__ import annotations
import time
from typing import Optional, Tuple
import torch
import torch.nn.functional as F
__all__ = [
"l2norm",
"repeat_kv_heads",
"recurrent_gated_delta_rule",
"chunk_gated_delta_rule",
]
# =============================================================================
# 工具函数
# =============================================================================
def l2norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-6) -> torch.Tensor:
"""沿 ``dim`` 做 L2 归一化。
注意 eps 加在**平方和**上,而不是加在模长上 —— 这是为了和 FLA 库 CUDA kernel
里的实现逐位对齐。写成 ``x / (x.norm(dim=dim, keepdim=True) + eps)``
数值上会差一点点,对齐测试就过不了。
"""
inv_norm = torch.rsqrt((x * x).sum(dim=dim, keepdim=True) + eps)
return x * inv_norm
def repeat_kv_heads(x: torch.Tensor, num_repeats: int, dim: int = 2) -> torch.Tensor:
"""GQA:把 query/key 的 head 维从 ``num_k_heads`` 展开到 ``num_v_heads``。
Qwen3-Next 默认 num_k_heads=16、num_v_heads=32,即每个 key head 被 2 个 value
head 共享。展开必须用 ``repeat_interleave``:
repeat_interleave -> [h0 h0 h1 h1 h2 h2 ...] ✅ 官方用法
tile / repeat -> [h0 h1 h2 ... h0 h1 h2 ...] ❌ 结果不同
``dim`` 默认 2,对应 ``[batch, seq_len, num_heads, head_dim]`` 布局。
"""
if num_repeats == 1:
return x
return x.repeat_interleave(num_repeats, dim=dim)
# =============================================================================
# 实现一:逐时间步递归
# =============================================================================
def recurrent_gated_delta_rule(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
initial_state: Optional[torch.Tensor] = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
"""逐时间步执行 gated delta rule。
这是最贴近数学定义的一版:一个 ``for t in range(T)`` 循环,每一步把状态矩阵 S
更新一次。好懂,但训练时不能用 —— 循环在时间维上串行,GPU 只能干等。
参数
----
query, key : [b, T, h, d_k]
value : [b, T, h, d_v]
g : [b, T, h] 对数衰减,<= 0
beta : [b, T, h] 写入强度,通常在 (0, 1)
initial_state : [b, h, d_k, d_v] 或 None
上一段序列留下的状态。推理时 prefill 出一份状态,decode 阶段每步带上它继续算,
这就是 GDN 的 O(1) 增量解码。
output_final_state : bool
是否返回走完这段序列后的状态。
返回
----
(o, final_state)
o : [b, T, h, d_v]
final_state : [b, h, d_k, d_v],未请求时为 None
签名与 transformers 里的 ``torch_recurrent_gated_delta_rule`` 完全一致,
方便直接互换对比。
"""
initial_dtype = query.dtype
# q/k 归一化到单位球面。只有归一化后,(I - k kᵀ) 才是严格意义的投影算子,
# delta rule 的几何解释才成立。
if use_qk_l2norm_in_kernel:
query = l2norm(query, dim=-1, eps=1e-6)
key = l2norm(key, dim=-1, eps=1e-6)
# [b, T, h, d] -> [b, h, T, d],方便按 (batch, head) 并行、只串行时间维。
# 统一升到 float32 计算:状态是长时间累积的量,半精度下误差会滚雪球。
query, key, value, beta, g = [
x.transpose(1, 2).contiguous().to(torch.float32) for x in (query, key, value, beta, g)
]
batch_size, num_heads, sequence_length, k_head_dim = key.shape
v_head_dim = value.shape[-1]
# 标准的 1/sqrt(d_k) 缩放,作用在 query 上。
# 注意:只缩放 q,不缩放 k —— 因为 k 要保持单位长度,投影算子才有意义。
query = query * (k_head_dim**-0.5)
if initial_state is None:
state = torch.zeros(
batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device
)
else:
state = initial_state.to(value.dtype)
o = torch.zeros(
batch_size, num_heads, sequence_length, v_head_dim, dtype=value.dtype, device=value.device
)
for t in range(sequence_length):
q_t = query[:, :, t] # [b, h, d_k]
k_t = key[:, :, t] # [b, h, d_k]
v_t = value[:, :, t] # [b, h, d_v]
# alpha_t ∈ (0, 1]:本步的遗忘系数。g 是对数域,取 exp 回到线性域。
alpha_t = g[:, :, t].exp() # [b, h]
beta_t = beta[:, :, t] # [b, h]
# --- 1) 遗忘:S <- α · S -------------------------------------------
state = state * alpha_t[..., None, None]
# --- 2) 回忆:当前 key 在(已遗忘的)状态里对应什么值 ----------------
# kv_mem = S k_t ∈ R^{d_v}
# [b,h,d_k,d_v] × [b,h,d_k] -> [b,h,d_v]
# 这一步等价于"用 k 去查表",看状态当前给出的答案。
kv_mem = (state * k_t[..., :, None]).sum(dim=-2) # [b, h, d_v]
# --- 3) 误差:期望值与回忆值之差,按写入强度缩放 --------------------
delta = (v_t - kv_mem) * beta_t[..., None] # [b, h, d_v]
# --- 4) 写入:S <- S + k_t ⊗ delta ---------------------------------
# 外积 [b,h,d_k,1] × [b,h,1,d_v] -> [b,h,d_k,d_v]
# 沿着 k 的方向把误差补进去;正交分量不受影响(因为 k 是单位向量)。
state = state + k_t[..., :, None] * delta[..., None, :]
# --- 5) 读出:o_t = S_tᵀ q_t ---------------------------------------
# 注意用的是**更新后**的 S_t,所以当前 token 自己写入的内容也能被自己读到。
o[:, :, t] = (state * q_t[..., :, None]).sum(dim=-2)
final_state = state if output_final_state else None
# 出口转回原始精度和 [b, T, h, d] 布局。final_state 保持 float32 不转
# (官方实现也是如此:它要在多步解码之间反复传递,转换会掉精度)。
o = o.transpose(1, 2).contiguous().to(initial_dtype)
return o, final_state
# =============================================================================
# 实现二:分块并行
# =============================================================================
def _chunk_pad_size(seq_len: int, chunk_size: int) -> int:
"""需要补多少个位置才能整除成整数个 chunk。"""
return (chunk_size - seq_len % chunk_size) % chunk_size
def _chunk_decay_mask(g: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""由块内逐位置的 ``g`` 算出累积衰减和衰减掩码。
参数
----
g : [b, h, nc, C] 每个位置的对数衰减(<= 0)
返回
----
g_cum : [b, h, nc, C]
``G_i = Σ_{s≤i} g_s``,块内累积对数衰减。
decay_mask : [b, h, nc, C, C]
``mask[i, j] = exp(G_i - G_j)``(``i >= j``),否则 0。
含义:位置 j 写入的内容,到位置 i 时还残留多少比例。
"""
g_cum = g.cumsum(dim=-1)
diff = g_cum.unsqueeze(-1) - g_cum.unsqueeze(-2) # diff[i,j] = G_i - G_j
# 先 .tril() 再 .exp(),顺序不能反:
# 上三角(i < j)处 G_i - G_j > 0,直接 exp 可能溢出成 inf,先置零就安全了。
decay_mask = diff.tril().exp()
# 第一次 tril 把上三角变成 0,exp(0) = 1,所以还要再 tril 一次把它们清干净。
decay_mask = decay_mask.tril()
return g_cum, decay_mask
def _wy_representation(
k_beta: torch.Tensor,
key: torch.Tensor,
decay_mask: torch.Tensor,
chunk_size: int,
) -> torch.Tensor:
"""块内 delta 修正的闭式解 ``T = (I - L)^{-1}``(WY 表示)。
参数
----
k_beta : [b, h, nc, C, d_k] ``key * beta``
key : [b, h, nc, C, d_k]
decay_mask : [b, h, nc, C, C]
返回
----
T : [b, h, nc, C, C] 单位下三角(对角线为 1)
L 严格下三角,所以逐行前代就是精确的 Neumann 求和:
T = I + L + L² + … + L^{C-1}
循环里第 i 行做的 ``row + (row ⊗ sub).sum(-2)`` 就是"把前面已经算好的部分再乘
一遍 L",等价于把级数多展开一项。
"""
C = chunk_size
device = k_beta.device
# 上三角 + 对角线:这些位置要清零,只留严格下三角
upper_and_diag = torch.triu(torch.ones(C, C, dtype=torch.bool, device=device), diagonal=0)
# L[i, j] = -β_i (k_i · k_j) exp(G_i - G_j),j < i
# k_beta @ keyᵀ 的第 (i,j) 项就是 β_i (k_i · k_j)
L = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(upper_and_diag, 0)
# 逐行前代。注意第 i 行只读下标 < i 的行,所以行与行之间没有循环依赖,
# 而且这也是"因果性"在数值上严格成立的原因。
for i in range(1, C):
row = L[..., i, :i].clone()
sub = L[..., :i, :i].clone()
L[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
return L + torch.eye(C, dtype=L.dtype, device=L.device)
def _chunk_state_pass(
query: torch.Tensor,
key: torch.Tensor,
w: torch.Tensor,
u: torch.Tensor,
g_cum: torch.Tensor,
decay_mask: torch.Tensor,
initial_state: Optional[torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
"""块间状态扫描:串行遍历 nc 个块,每块内部是稠密矩阵乘。
参数
----
query : [b, h, nc, C, d_k]
key : [b, h, nc, C, d_k]
w : [b, h, nc, C, d_v] 块内修正后的 value(``T @ (β⊙v)``)
u : [b, h, nc, C, d_k] 块内修正后的 key(``T @ (β⊙k⊙exp(G))``)
g_cum : [b, h, nc, C]
decay_mask : [b, h, nc, C, C]
initial_state : [b, h, d_k, d_v] 或 None
返回
----
(o, final_state)
o : [b, h, nc, C, d_v]
final_state : [b, h, d_k, d_v]
"""
batch_size, num_heads, num_chunks, C, v_head_dim = w.shape
k_head_dim = key.shape[-1]
device = w.device
if initial_state is None:
state = torch.zeros(
batch_size, num_heads, k_head_dim, v_head_dim, dtype=w.dtype, device=device
)
else:
state = initial_state.to(w.dtype)
o = torch.empty_like(w)
# 块内因果掩码:清零严格上三角(i < j),保留下三角含对角线。
# 与 recurrent 版"位置 t 只看 s <= t"一一对应。
strict_upper = torch.triu(torch.ones(C, C, dtype=torch.bool, device=device), diagonal=1)
for i in range(num_chunks):
q_i, k_i = query[:, :, i], key[:, :, i] # [b, h, C, d_k]
w_i, u_i = w[:, :, i], u[:, :, i] # [b, h, C, d_v] / [b, h, C, d_k]
g_i = g_cum[:, :, i] # [b, h, C]
decay_i = decay_mask[:, :, i] # [b, h, C, C]
# --- 块内因果注意力:位置 i 对块内位置 j<=i 的权重 -------------------
# 就是 scaled q·k,再乘上 j -> i 之间的衰减残留 exp(G_i - G_j)
intra = (q_i @ k_i.transpose(-1, -2) * decay_i).masked_fill(strict_upper, 0)
# --- 块外状态里已经存过的部分,先减掉(delta rule 的"擦除")----------
# u 是块内 key 经过修正后的版本,u @ S 得到"状态对块内每个 key 的回忆值"
v_prime = u_i @ state # [b, h, C, d_v]
w_new = w_i - v_prime
# --- 块外状态直接读出的部分 -----------------------------------------
# q 乘 exp(G_i) 表示把状态从块首衰减到位置 i
inter = (q_i * g_i.exp().unsqueeze(-1)) @ state # [b, h, C, d_v]
# --- 块内 + 块外合并 -------------------------------------------------
o[:, :, i] = inter + intra @ w_new
# --- 状态推进 --------------------------------------------------------
# 第一项:整块衰减到块末(G_C 是块内最后一个位置的累积衰减)
# 第二项:写入本块修正后的内容;key 也要按"从 j 到块末"的衰减加权
state = state * g_i[:, :, -1, None, None].exp() + (
k_i * (g_i[:, :, -1, None] - g_i).exp().unsqueeze(-1)
).transpose(-1, -2) @ w_new
return o, state
def chunk_gated_delta_rule(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
chunk_size: int = 64,
initial_state: Optional[torch.Tensor] = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
"""分块并行执行 gated delta rule。
参数与 :func:`recurrent_gated_delta_rule` 相同,多一个 ``chunk_size``。
``chunk_size >= 序列长度`` 时只有一个块,退化成"整段一次算完"的二次形式 ——
数值上仍然正确,但复杂度是 O(T²),可以拿来对照理解分块到底省了什么。
返回
----
(o, final_state),形状与递归版一致。
"""
initial_dtype = query.dtype
if use_qk_l2norm_in_kernel:
query = l2norm(query, dim=-1, eps=1e-6)
key = l2norm(key, dim=-1, eps=1e-6)
# 与递归版相同的布局变换和 float32 升位
query, key, value, beta, g = [
x.transpose(1, 2).contiguous().to(torch.float32) for x in (query, key, value, beta, g)
]
batch_size, num_heads, sequence_length, k_head_dim = key.shape
v_head_dim = value.shape[-1]
# ---- 0) 补齐到 chunk 整数倍 ---------------------------------------------
# 补进来的位置 k=v=0、β=0、g=0,对真实位置没有任何影响(因果性保证),
# 最后再切掉就行。
pad_size = _chunk_pad_size(sequence_length, chunk_size)
if pad_size:
query = F.pad(query, (0, 0, 0, pad_size))
key = F.pad(key, (0, 0, 0, pad_size))
value = F.pad(value, (0, 0, 0, pad_size))
beta = F.pad(beta, (0, pad_size))
g = F.pad(g, (0, pad_size))
total_len = sequence_length + pad_size
num_chunks = total_len // chunk_size
# ---- 1) 缩放 query -------------------------------------------------------
# 只缩放 q,且必须在切块前完成(和官方一致,避免切块改变数值)
query = query * (k_head_dim**-0.5)
# ---- 2) 切块:[b, h, T, d] -> [b, h, nc, C, d] ---------------------------
def to_chunks(x: torch.Tensor) -> torch.Tensor:
return x.reshape(batch_size, num_heads, num_chunks, chunk_size, x.shape[-1])
# β⊙k / β⊙v 必须在切块**之前**算好:此时 key/value 还是 [b,h,T,d],
# 才能和 beta 的 [b,h,T] 广播。切块后 T 已经被拆成 (nc, C),就对不上了。
k_beta = key * beta.unsqueeze(-1)
v_beta = value * beta.unsqueeze(-1)
query, key, value, k_beta, v_beta = [
to_chunks(x) for x in (query, key, value, k_beta, v_beta)
]
g = g.reshape(batch_size, num_heads, num_chunks, chunk_size)
# ---- 3) 块内累积衰减与 decay_mask ---------------------------------------
g_cum, decay_mask = _chunk_decay_mask(g)
# ---- 4) WY 表示,得到块内"已修正"的 value / key --------------------------
wy = _wy_representation(k_beta, key, decay_mask, chunk_size)
w = wy @ v_beta # 对应官方实现里的 `value = attn @ v_beta`
u = wy @ (k_beta * g_cum.exp().unsqueeze(-1)) # 对应官方实现里的 `k_cumdecay`
# ---- 5) 块间状态扫描 -----------------------------------------------------
o, final_state = _chunk_state_pass(query, key, w, u, g_cum, decay_mask, initial_state)
if not output_final_state:
final_state = None
# ---- 6) 还原形状:切掉 padding,转回 [b, T, h, d_v] 和原始精度 ------------
o = o.reshape(batch_size, num_heads, total_len, v_head_dim)[:, :, :sequence_length]
o = o.transpose(1, 2).contiguous().to(initial_dtype)
return o, final_state
# =============================================================================
# 测试
# =============================================================================
# ---------------------------------------------------------------------------
# 官方实现(对齐金标准)。装不上 transformers 时相关用例自动 skip。
# ---------------------------------------------------------------------------
try:
from transformers.models.qwen3_next.modeling_qwen3_next import (
torch_chunk_gated_delta_rule as _ref_chunk,
torch_recurrent_gated_delta_rule as _ref_recurrent,
)
HAS_TRANSFORMERS = True
except Exception: # pragma: no cover - 取决于环境
HAS_TRANSFORMERS = False
import pytest # noqa: E402
requires_transformers = pytest.mark.skipif(
not HAS_TRANSFORMERS, reason="需要 transformers(含 qwen3_next)才能做官方对齐"
)
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
# 与官方实现当前是逐位相同(max diff 恰好为 0),留一点余量防止版本/设备差异。
ATOL_OFFICIAL = 1e-6
# 递归 vs 分块:数学等价但求和顺序不同,容差放大到 1e-4
ATOL_CROSS = 1e-4
def make_inputs(batch=2, seq_len=128, num_heads=4, k_dim=32, v_dim=None, seed=0, dtype=torch.float32):
"""造一组随机输入。
``g`` 必须 <= 0(它是**对数**衰减,模型里由 ``-exp(A_log)*softplus(...)`` 得到),
否则状态会指数爆炸。这里用 ``-softplus(randn)`` 模拟同样的取值范围。
"""
if v_dim is None:
v_dim = k_dim
torch.manual_seed(seed)
q = torch.randn(batch, seq_len, num_heads, k_dim, device=DEVICE, dtype=dtype)
k = torch.randn(batch, seq_len, num_heads, k_dim, device=DEVICE, dtype=dtype)
v = torch.randn(batch, seq_len, num_heads, v_dim, device=DEVICE, dtype=dtype)
beta = torch.rand(batch, seq_len, num_heads, device=DEVICE, dtype=dtype)
g = -F.softplus(torch.randn(batch, seq_len, num_heads, device=DEVICE, dtype=dtype))
return q, k, v, g, beta
# --------------------------- 1. 自洽性:递归 vs 分块 ---------------------------
@pytest.mark.parametrize("chunk_size", [16, 32, 64, 128])
def test_recurrent_matches_own_chunk(chunk_size):
"""递归版与分块版必须一致 —— 这是分块推导正确最直接的证据。"""
q, k, v, g, beta = make_inputs(seq_len=128)
o_rec, s_rec = recurrent_gated_delta_rule(
q, k, v, g, beta, None, True, use_qk_l2norm_in_kernel=True
)
o_chk, s_chk = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=chunk_size, output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
assert o_rec.shape == o_chk.shape == q.shape[:-1] + (v.shape[-1],)
assert torch.allclose(o_rec, o_chk, atol=ATOL_CROSS, rtol=ATOL_CROSS)
assert torch.allclose(s_rec, s_chk, atol=ATOL_CROSS, rtol=ATOL_CROSS)
def test_chunk_size_does_not_change_result():
"""chunk_size 只是实现细节,不应该影响数学结果。"""
q, k, v, g, beta = make_inputs(seq_len=128)
sizes = [8, 16, 32, 64, 128]
outputs = [
chunk_gated_delta_rule(q, k, v, g, beta, chunk_size=cs, use_qk_l2norm_in_kernel=True)[0]
for cs in sizes
]
base = outputs[0]
for chunk_size, o in zip(sizes[1:], outputs[1:]):
assert torch.allclose(base, o, atol=ATOL_CROSS, rtol=ATOL_CROSS), (
f"chunk_size 改变后结果不一致: chunk_size={chunk_size}"
)
def test_single_chunk_equals_quadratic_form():
"""chunk_size >= T 时只有一个块,退化成"整段一次算完"的二次形式。
补零到 256(chunk_size=256,T=128)与正好一块(chunk_size=128)应当完全一致 ——
说明补进来的位置确实没有污染真实输出。
"""
q, k, v, g, beta = make_inputs(seq_len=128)
o_one_chunk, _ = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=128, use_qk_l2norm_in_kernel=True
)
o_padded_one_chunk, _ = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=256, use_qk_l2norm_in_kernel=True
)
o_rec, _ = recurrent_gated_delta_rule(q, k, v, g, beta, use_qk_l2norm_in_kernel=True)
assert torch.allclose(o_one_chunk, o_padded_one_chunk, atol=ATOL_CROSS, rtol=ATOL_CROSS)
# 单块 = 完整 O(T²) 因果注意力,数学上与递归版等价
assert torch.allclose(o_one_chunk, o_rec, atol=ATOL_CROSS, rtol=ATOL_CROSS)
def test_non_divisible_sequence_length():
"""序列长度不是 chunk_size 整数倍时的补零逻辑。"""
q, k, v, g, beta = make_inputs(seq_len=100)
o_rec, _ = recurrent_gated_delta_rule(q, k, v, g, beta, use_qk_l2norm_in_kernel=True)
for chunk_size in [16, 32, 64]:
o_chk, _ = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=chunk_size, use_qk_l2norm_in_kernel=True
)
assert o_chk.shape == o_rec.shape
assert torch.allclose(o_rec, o_chk, atol=ATOL_CROSS, rtol=ATOL_CROSS)
# ---------------------- 2. 对齐金标准:Qwen3-Next 官方实现 ----------------------
@requires_transformers
def test_matches_qwen3next_recurrent():
"""递归版 vs 官方 torch_recurrent_gated_delta_rule。"""
q, k, v, g, beta = make_inputs(seq_len=128)
mine, s_mine = recurrent_gated_delta_rule(
q, k, v, g, beta, None, True, use_qk_l2norm_in_kernel=True
)
ref, s_ref = _ref_recurrent(q, k, v, g, beta, None, True, use_qk_l2norm_in_kernel=True)
assert torch.allclose(mine, ref, atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
assert torch.allclose(s_mine, s_ref, atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
@requires_transformers
@pytest.mark.parametrize("chunk_size", [16, 64, 128])
def test_matches_qwen3next_chunk(chunk_size):
"""分块版 vs 官方 torch_chunk_gated_delta_rule。"""
q, k, v, g, beta = make_inputs(seq_len=128)
mine, s_mine = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=chunk_size, output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
ref, s_ref = _ref_chunk(
q, k, v, g, beta, chunk_size=chunk_size, initial_state=None,
output_final_state=True, use_qk_l2norm_in_kernel=True,
)
assert torch.allclose(mine, ref, atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
assert torch.allclose(s_mine, s_ref, atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
@requires_transformers
def test_matches_qwen3next_with_initial_state():
"""带 initial_state 时也要和官方一致(增量解码路径)。"""
q, k, v, g, beta = make_inputs(seq_len=64, seed=7)
state0 = torch.randn(2, 4, 32, 32, device=DEVICE) * 0.1
mine, s_mine = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=16, initial_state=state0,
output_final_state=True, use_qk_l2norm_in_kernel=True,
)
ref, s_ref = _ref_chunk(
q, k, v, g, beta, chunk_size=16, initial_state=state0,
output_final_state=True, use_qk_l2norm_in_kernel=True,
)
assert torch.allclose(mine, ref, atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
assert torch.allclose(s_mine, s_ref, atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
@requires_transformers
@pytest.mark.parametrize("seq_len", [1, 7, 63, 64, 65])
def test_matches_qwen3next_odd_lengths(seq_len):
"""短序列 / 刚好卡在 chunk 边界上的长度。"""
q, k, v, g, beta = make_inputs(seq_len=seq_len, seed=11)
mine, _ = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=64, use_qk_l2norm_in_kernel=True
)
ref, _ = _ref_chunk(q, k, v, g, beta, chunk_size=64, use_qk_l2norm_in_kernel=True)
assert torch.allclose(mine, ref, atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
# ------------------------------ 3. 性质测试 ------------------------------
def test_causality():
"""改动 t 时刻**之后**的输入,t 及之前的输出必须逐位不变。
分块实现里块内是稠密矩阵乘,很容易不小心让未来信息漏进来,所以这个测试很关键。
"""
q, k, v, g, beta = make_inputs(seq_len=64)
split = 40
q2, k2, v2 = (x.clone() for x in (q, k, v))
for x in (q2, k2, v2):
x[:, split + 1 :] = torch.randn_like(x[:, split + 1 :])
beta2 = beta.clone()
beta2[:, split + 1 :] = torch.rand_like(beta2[:, split + 1 :])
g2 = g.clone()
g2[:, split + 1 :] = -F.softplus(torch.randn_like(g2[:, split + 1 :]))
# chunk_size=16 保证 split 落在块**内部**,能真正检验块内掩码
o1, _ = chunk_gated_delta_rule(q, k, v, g, beta, chunk_size=16, use_qk_l2norm_in_kernel=True)
o2, _ = chunk_gated_delta_rule(q2, k2, v2, g2, beta2, chunk_size=16, use_qk_l2norm_in_kernel=True)
past_diff = (o1[:, : split + 1] - o2[:, : split + 1]).abs().max().item()
assert past_diff == 0.0, f"过去位置的输出被未来的输入影响了,max diff = {past_diff}"
# 未来位置确实变了 —— 否则说明测试构造有问题,没真正改到东西
assert not torch.allclose(o1[:, split + 1 :], o2[:, split + 1 :])
def test_prefill_decode_consistency():
"""prefill 出一份状态,再逐步 decode,结果应与整段一次算完一致。
这正是 GDN 相比标准注意力省显存的地方:状态大小恒定,与已生成长度无关。
"""
q, k, v, g, beta = make_inputs(seq_len=64)
split = 48
o_full, _ = recurrent_gated_delta_rule(q, k, v, g, beta, use_qk_l2norm_in_kernel=True)
o_pre, state = recurrent_gated_delta_rule(
q[:, :split], k[:, :split], v[:, :split], g[:, :split], beta[:, :split],
None, True, use_qk_l2norm_in_kernel=True,
)
o_dec, _ = recurrent_gated_delta_rule(
q[:, split:], k[:, split:], v[:, split:], g[:, split:], beta[:, split:],
state, False, use_qk_l2norm_in_kernel=True,
)
assert torch.allclose(o_pre, o_full[:, :split], atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
assert torch.allclose(o_dec, o_full[:, split:], atol=ATOL_OFFICIAL, rtol=ATOL_OFFICIAL)
def test_initial_state_equals_prefix():
"""传 initial_state ≡ 把产生该状态的那段前缀拼到序列前面。"""
qp, kp, vp, gp, bp = make_inputs(seq_len=32, seed=1)
q, k, v, g, beta = make_inputs(seq_len=48, seed=2)
prefix_len = 32
_, state = recurrent_gated_delta_rule(
qp, kp, vp, gp, bp, None, True, use_qk_l2norm_in_kernel=True
)
o_with_state, _ = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=16, initial_state=state, use_qk_l2norm_in_kernel=True
)
def cat(*xs):
return torch.cat(xs, dim=1)
o_concat, _ = chunk_gated_delta_rule(
cat(qp, q), cat(kp, k), cat(vp, v), cat(gp, g), cat(bp, beta),
chunk_size=16, use_qk_l2norm_in_kernel=True,
)
assert torch.allclose(o_with_state, o_concat[:, prefix_len:], atol=1e-5, rtol=1e-5)
def test_gqa_repeat_kv_heads_is_interleave():
"""GQA 展开必须是 interleave(h0 h0 h1 h1),不是 tile(h0 h1 h0 h1)。"""
x = torch.arange(2 * 3).reshape(1, 1, 2, 3).float() # [b=1, T=1, h=2, d=3]
out = repeat_kv_heads(x, num_repeats=2, dim=2)
assert out.shape == (1, 1, 4, 3)
assert torch.equal(out[0, 0, 0], x[0, 0, 0])
assert torch.equal(out[0, 0, 1], x[0, 0, 0]) # h0 连续重复
assert torch.equal(out[0, 0, 2], x[0, 0, 1])
assert torch.equal(out[0, 0, 3], x[0, 0, 1])
# 反例:tile 的结果不同,确认测试真的能区分两者
tiled = x.repeat(1, 1, 2, 1)
assert not torch.equal(out, tiled)
def test_gqa_end_to_end():
"""num_v_heads = 2 * num_k_heads 时,q/k 展开后两套实现仍然一致。"""
num_k_heads, num_v_heads = 2, 4
q_wide, k_wide, v, g, beta = make_inputs(
seq_len=64, num_heads=num_v_heads, k_dim=32, v_dim=32, seed=3
)
# 只取前 num_k_heads 个 head,当作"还没做 GQA 展开"的 q/k
q, k = q_wide[:, :, :num_k_heads], k_wide[:, :, :num_k_heads]
q_e = repeat_kv_heads(q, num_v_heads // num_k_heads, dim=2)
k_e = repeat_kv_heads(k, num_v_heads // num_k_heads, dim=2)
assert q_e.shape[2] == k_e.shape[2] == num_v_heads
o_rec, _ = recurrent_gated_delta_rule(q_e, k_e, v, g, beta, use_qk_l2norm_in_kernel=True)
o_chk, _ = chunk_gated_delta_rule(
q_e, k_e, v, g, beta, chunk_size=16, use_qk_l2norm_in_kernel=True
)
assert torch.allclose(o_rec, o_chk, atol=ATOL_CROSS, rtol=ATOL_CROSS)
def test_beta_zero_writes_nothing():
"""β=0 表示"完全不写入",状态恒为零,输出也应为零。"""
q, k, v, g, beta = make_inputs(seq_len=32)
beta = torch.zeros_like(beta)
o, s = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=8, output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
assert torch.allclose(o, torch.zeros_like(o), atol=1e-6)
assert torch.allclose(s, torch.zeros_like(s), atol=1e-6)
def test_delta_rule_erases_old_value():
"""delta rule 的核心语义:同一个 key 重复写入会**覆盖**,而不是累加。
构造:g=0(不衰减)、β=1(完全写入)、q=k̂(读出同一个方向),
两个时间步用同一个 key、不同的 value。
- delta rule -> 第二步之后状态 = k̂ ⊗ v1,读出 v1/√d_k(旧值被擦掉)
- 朴素线性注意力 -> 状态 = k̂ ⊗ (v0+v1),读出 (v0+v1)/√d_k(旧值残留)
这就是 GDN 比"Mamba2 + 纯累加"强的地方。
"""
seq_len, num_heads, dim = 2, 1, 16
torch.manual_seed(0)
# 两个时间步用**同一个** key
k_vec = torch.randn(1, 1, num_heads, dim, device=DEVICE)
q = k_vec.expand(1, seq_len, num_heads, dim).contiguous() # q = k,读出该 key 方向
k = k_vec.expand(1, seq_len, num_heads, dim).contiguous()
v = torch.randn(1, seq_len, num_heads, dim, device=DEVICE)
v[:, 1] = torch.randn_like(v[:, 1]) # 第二步写入一个明显不同的值
beta = torch.ones(1, seq_len, num_heads, device=DEVICE)
g = torch.zeros(1, seq_len, num_heads, device=DEVICE) # alpha = 1,不衰减
o, _ = recurrent_gated_delta_rule(q, k, v, g, beta, use_qk_l2norm_in_kernel=True)
scale = dim**-0.5 # query 的 1/sqrt(d_k) 缩放
# q̂ · k̂ = 1,所以读出恰为 v1 * scale
assert torch.allclose(o[:, 1], v[:, 1] * scale, atol=1e-5), "delta rule 没有正确覆盖旧值"
# 反例:朴素累加会得到 (v0 + v1) * scale,确认这个测试有鉴别力
naive = (v[:, 0] + v[:, 1]) * scale
assert not torch.allclose(o[:, 1], naive, atol=1e-3), "输出与朴素累加相同,测试没有区分力"
def test_backward_gradients():
"""反向传播可通,且梯度有限(没有 NaN/Inf)。"""
q, k, v, g, beta = make_inputs(seq_len=64)
for x in (q, k, v, g, beta):
x.requires_grad_(True)
o, s = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=16, output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
(o.sum() + s.sum()).backward()
for name, x in zip(["q", "k", "v", "g", "beta"], (q, k, v, g, beta)):
assert x.grad is not None, f"{name} 没有梯度"
assert torch.isfinite(x.grad).all(), f"{name} 的梯度出现 NaN/Inf"
assert x.grad.abs().sum() > 0, f"{name} 的梯度恒为零"
def test_gradients_match_between_implementations():
"""两套实现的反向梯度也应一致 —— 说明分块推导的伴随(backward)也是对的。"""
q, k, v, g, beta = make_inputs(seq_len=64, seed=5)
impls = [
lambda *a: recurrent_gated_delta_rule(*a, use_qk_l2norm_in_kernel=True),
lambda *a: chunk_gated_delta_rule(*a, chunk_size=16, use_qk_l2norm_in_kernel=True),
]
grad_sets = []
for fn in impls:
inputs = [x.clone().detach().requires_grad_(True) for x in (q, k, v, g, beta)]
out, _ = fn(*inputs)
out.sum().backward()
grad_sets.append([x.grad for x in inputs])
for label, gr, gc in zip(["q", "k", "v", "g", "beta"], *grad_sets):
assert torch.allclose(gr, gc, atol=1e-4, rtol=1e-3), f"{label} 的梯度两套实现不一致"
def test_long_sequence_stability():
"""长序列下不出现 NaN/Inf,且与递归版仍然对齐。"""
q, k, v, g, beta = make_inputs(batch=1, seq_len=2048, num_heads=2, k_dim=32, seed=13)
o_chk, s_chk = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=64, output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
o_rec, s_rec = recurrent_gated_delta_rule(
q, k, v, g, beta, None, True, use_qk_l2norm_in_kernel=True
)
assert torch.isfinite(o_chk).all(), "输出出现 NaN/Inf"
assert torch.isfinite(s_chk).all(), "最终状态出现 NaN/Inf"
assert (o_chk - o_rec).abs().max().item() < ATOL_CROSS
assert (s_chk - s_rec).abs().max().item() < ATOL_CROSS
def test_strong_decay_forgets_the_past():
"""衰减极强时(g 很负),输出应当只看最近几步 —— 验证门控确实在起作用。"""
q, k, v, g, beta = make_inputs(seq_len=64, seed=17)
g = torch.full_like(g, -20.0) # alpha = exp(-20) ≈ 2e-9,几乎瞬间遗忘
o, _ = recurrent_gated_delta_rule(q, k, v, g, beta, use_qk_l2norm_in_kernel=True)
# 重算一遍:只保留最后一步的输入,前面的全部置零,结果应当非常接近
mask = torch.zeros_like(v)
mask[:, -1] = 1.0
o_last, _ = recurrent_gated_delta_rule(
q, k, v * mask, g, beta, use_qk_l2norm_in_kernel=True
)
assert torch.allclose(o[:, -1], o_last[:, -1], atol=1e-4, rtol=1e-4)
# ------------------------------ 4. 性能对比 ------------------------------
def test_chunked_faster_than_recurrent_on_long_sequence():
"""分块版在长序列上应显著快于逐时间步版本。
这也是 Qwen3-Next 训练必须用分块形式的原因:递归版把 GPU 全耗在 kernel 启动上了。
"""
q, k, v, g, beta = make_inputs(batch=1, seq_len=2048, num_heads=4, k_dim=64, seed=19)
def timed(fn, repeat=3):
fn() # warmup(CUDA 首次调用含初始化开销)
if DEVICE == "cuda":
torch.cuda.synchronize()
t0 = time.perf_counter()
for _ in range(repeat):
fn()
if DEVICE == "cuda":
torch.cuda.synchronize()
return (time.perf_counter() - t0) / repeat
t_rec = timed(
lambda: recurrent_gated_delta_rule(q, k, v, g, beta, use_qk_l2norm_in_kernel=True)
)
t_chk = timed(
lambda: chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=64, use_qk_l2norm_in_kernel=True
)
)
print(f"\n [T=2048, d_k=d_v=64, h=4, device={DEVICE}]")
print(f" 递归版 (逐时间步): {t_rec * 1e3:8.1f} ms")
print(f" 分块版 (C=64) : {t_chk * 1e3:8.1f} ms 加速比 {t_rec / t_chk:.1f}x")
# 宽松上界:分块版不该比递归版慢多少。真正做到"更快"与否取决于设备,
# 具体数字看上面的打印。
assert t_chk < t_rec * 3, "分块版比递归版慢了 3 倍以上,可能存在性能问题"
# =============================================================================
# 演示:直接 python gdn.py 时跑的完整自检
# =============================================================================
_SEP = "=" * 78
def _banner(title: str) -> None:
print(f"\n{_SEP}\n{title}\n{_SEP}")
def main() -> None:
print(f"GDN 单文件演示 torch={torch.__version__} device={DEVICE}")
if DEVICE == "cuda":
print(f"GPU: {torch.cuda.get_device_name(0)}")
# ------------------------------------------------------------------
_banner("1. 形状约定")
b, T, h, dk, dv = 2, 128, 4, 32, 32
q, k, v, g, beta = make_inputs(b, T, h, dk, dv, seed=0)
print(f" query/key : {tuple(q.shape)} [batch, seq_len, num_heads, head_dim]")
print(f" value : {tuple(v.shape)}")
print(f" g, beta : {tuple(g.shape)} [batch, seq_len, num_heads]")
print(f" g 的范围 : [{g.min():.3f}, {g.max():.3f}] (对数衰减,恒 <= 0)")
print(f" beta 范围 : [{beta.min():.3f}, {beta.max():.3f}]")
# ------------------------------------------------------------------
_banner("2. 递归版 vs 分块版:不同 chunk_size 下的最大误差")
o_rec, s_rec = recurrent_gated_delta_rule(
q, k, v, g, beta, None, True, use_qk_l2norm_in_kernel=True
)
print(f" 输出形状 {tuple(o_rec.shape)},最终状态 {tuple(s_rec.shape)}")
print()
print(f" {'chunk_size':>12} {'输出 max|Δ|':>16} {'状态 max|Δ|':>16}")
print(f" {'-' * 12} {'-' * 16} {'-' * 16}")
for cs in [8, 16, 32, 64, 128, 256]:
o_chk, s_chk = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=cs, output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
print(
f" {cs:>12} {(o_rec - o_chk).abs().max().item():>16.3e}"
f" {(s_rec - s_chk).abs().max().item():>16.3e}"
)
print("\n 注:chunk_size >= T(=128) 时只有一个块,退化成 O(T²) 的二次形式。")
# ------------------------------------------------------------------
_banner("3. 与 transformers 里 Qwen3-Next 官方实现的对比")
if not HAS_TRANSFORMERS:
print(" 未安装 transformers,跳过。")
else:
mine, s_mine = recurrent_gated_delta_rule(
q, k, v, g, beta, None, True, use_qk_l2norm_in_kernel=True
)
ref, s_ref = _ref_recurrent(q, k, v, g, beta, None, True, use_qk_l2norm_in_kernel=True)
print(f" 递归版 vs torch_recurrent_gated_delta_rule : "
f"max|Δ| = {(mine - ref).abs().max().item():.3e}")
for cs in [16, 64, 128]:
mine, s_mine = chunk_gated_delta_rule(
q, k, v, g, beta, chunk_size=cs, output_final_state=True,
use_qk_l2norm_in_kernel=True,
)
ref, s_ref = _ref_chunk(
q, k, v, g, beta, chunk_size=cs, initial_state=None,
output_final_state=True, use_qk_l2norm_in_kernel=True,
)
print(f" 分块版(C={cs:<3}) vs torch_chunk_gated_delta_rule : "
f"max|Δ| = {(mine - ref).abs().max().item():.3e}")
print("\n 差值为 0 表示逐位完全相同 —— 运算顺序、float32 升位策略、")
print(" l2norm 的 eps 位置都刻意和官方对齐了。")
# ------------------------------------------------------------------
_banner("4. 因果性:改动未来的输入,过去输出必须逐位不变")
split = 40
q2, k2, v2 = (x.clone() for x in (q, k, v))
for x in (q2, k2, v2):
x[:, split + 1 :] = torch.randn_like(x[:, split + 1 :])
beta2 = beta.clone()
beta2[:, split + 1 :] = torch.rand_like(beta2[:, split + 1 :])
g2 = g.clone()
g2[:, split + 1 :] = -F.softplus(torch.randn_like(g2[:, split + 1 :]))
o1, _ = chunk_gated_delta_rule(q, k, v, g, beta, chunk_size=16, use_qk_l2norm_in_kernel=True)
o2, _ = chunk_gated_delta_rule(q2, k2, v2, g2, beta2, chunk_size=16, use_qk_l2norm_in_kernel=True)
past = (o1[:, : split + 1] - o2[:, : split + 1]).abs().max().item()
future = (o1[:, split + 1 :] - o2[:, split + 1 :]).abs().max().item()
print(f" 在 t={split} 之后替换输入(chunk_size=16,切分点在块内部)")
print(f" 过去位置 (t <= {split}) 的 max|Δ| = {past:.3e} ← 必须为 0")
print(f" 未来位置 (t > {split}) 的 max|Δ| = {future:.3e} ← 必须明显 > 0")
# ------------------------------------------------------------------
_banner("5. Delta rule 的'擦除'语义(最小例子:T=2, g=0, β=1)")
dim, ns = 16, 1
torch.manual_seed(0)
k_vec = torch.randn(1, 1, ns, dim, device=DEVICE)
qq = k_vec.expand(1, 2, ns, dim).contiguous()
kk = k_vec.expand(1, 2, ns, dim).contiguous()
vv = torch.randn(1, 2, ns, dim, device=DEVICE)
vv[:, 1] = torch.randn_like(vv[:, 1])
o, _ = recurrent_gated_delta_rule(
qq, kk, vv,
torch.zeros(1, 2, ns, device=DEVICE),
torch.ones(1, 2, ns, device=DEVICE),
use_qk_l2norm_in_kernel=True,
)
scale = dim**-0.5
print(" 同一个 key 连续写入 v0、v1,然后读出:")
print(f" delta rule 读出 o[1] 与 v1/√d 的差 : "
f"{(o[:, 1] - vv[:, 1] * scale).abs().max().item():.3e} ← 旧值被擦掉")
print(f" 朴素累加会读出 (v0+v1)/√d,与 v1/√d 的差 : "
f"{((vv[:, 0] + vv[:, 1]) * scale - vv[:, 1] * scale).abs().max().item():.3e}"
f" ← 旧值残留")
# ------------------------------------------------------------------
_banner("6. 增量解码:prefill + decode ≡ 整段一次算完")
qq, kk, vv, gg, bb = make_inputs(seq_len=64, seed=0)
o_full, _ = recurrent_gated_delta_rule(qq, kk, vv, gg, bb, use_qk_l2norm_in_kernel=True)
o_pre, state = recurrent_gated_delta_rule(
qq[:, :48], kk[:, :48], vv[:, :48], gg[:, :48], bb[:, :48],
None, True, use_qk_l2norm_in_kernel=True,
)
o_dec, _ = recurrent_gated_delta_rule(
qq[:, 48:], kk[:, 48:], vv[:, 48:], gg[:, 48:], bb[:, 48:],
state, False, use_qk_l2norm_in_kernel=True,
)
print(" 状态大小 : "
f"{tuple(state.shape)} = {state.numel() * 4 / 1024:.1f} KB"
f" (与已处理长度无关,标准注意力这里要存 48 步的 KV)")
print(f" prefill 段 (t<48) 的一致性 max|Δ| : {(o_pre - o_full[:, :48]).abs().max().item():.3e}")
print(f" decode 段 (t>=48) 的一致性 max|Δ| : {(o_dec - o_full[:, 48:]).abs().max().item():.3e}")
# ------------------------------------------------------------------
_banner("7. 性能对比:为什么训练必须用分块版")
ql, kl, vl, gl, bl = make_inputs(batch=1, seq_len=2048, num_heads=4, k_dim=64, seed=19)
def timed(fn, repeat=3):
fn()
if DEVICE == "cuda":
torch.cuda.synchronize()
t0 = time.perf_counter()
for _ in range(repeat):
fn()
if DEVICE == "cuda":
torch.cuda.synchronize()
return (time.perf_counter() - t0) / repeat
t_rec = timed(lambda: recurrent_gated_delta_rule(ql, kl, vl, gl, bl, use_qk_l2norm_in_kernel=True))
t_chk = timed(lambda: chunk_gated_delta_rule(ql, kl, vl, gl, bl, chunk_size=64, use_qk_l2norm_in_kernel=True))
print(f" [T=2048, d_k=d_v=64, h=4, device={DEVICE}]")
print(f" 递归版 (逐时间步) : {t_rec * 1e3:8.1f} ms (2048 次串行 kernel)")
print(f" 分块版 (C=64) : {t_chk * 1e3:8.1f} ms (32 次串行 kernel)")
print(f" 加速比 : {t_rec / t_chk:.1f}x")
# ------------------------------------------------------------------
_banner("下一步")
print(" 本文件只做 core 算法。要变成能跑的完整层,还需要:")
print(" 1. 投影层 in_proj_qkvz / in_proj_ba")
print(" 2. 短卷积 conv1d (depthwise, kernel_size=4)")
print(" 3. 门控参数化 A_log + dt_bias")
print(" 4. Gated RMSNorm + out_proj")
print(" 5. 增量解码缓存 conv_state + recurrent_state")
print("\n 跑测试: python -m pytest gdn.py -v")
print()
if __name__ == "__main__":
main()

浙公网安备 33010602011771号