AIGC标识 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()
posted @ 2026-09-30 19:48  Dsp Tian  阅读(28)  评论(0)    收藏  举报