Gated Delta Network(GDN)

本页专题分析 Gated Delta Network(Gated DeltaNet, GDN) [1]递推公式、并行化训练算法,并回答一个关键问题: 长序列上的递推到底是逐 token 串行累积,还是另有并行方法?

结论先行:GDN 是线性注意力 + 门控遗忘 + Delta 更新规则的组合。 数学上它是逐 token 的一阶线性递推(串行),但训练与 prefill 用的是 chunkwise(分块并行)算法——块内用矩阵乘法一次算完,只在块间串行传递状态; 只有自回归解码(decode,逐 token 生成)时才真正回到逐 token 的串行递推。

1. 出发点:线性注意力的状态递推

线性注意力 [2] 维护一个矩阵状态 StRdv×dkS_t \in \mathbb{R}^{d_v \times d_k}:第 tt 个 token 的 key ktRdkk_t \in \mathbb{R}^{d_k} 与 value vtRdvv_t \in \mathbb{R}^{d_v}外积 vtktv_t k_t^\top 写入状态, 再用 query qtRdkq_t \in \mathbb{R}^{d_k} 读出:

St=St1+vtkt,ot=Stqt S_t = S_{t-1} + v_t k_t^\top, \qquad o_t = S_t q_t

式 1|线性注意力的状态递推:外积写入 vtktv_tk_t^\top,query 读出 StqtS_tq_t

  • 它等价于一个 key-value 关联记忆(outer-product memory);
  • 能精确存储的正交 key-value 对数受维度限制,序列一长就发生 memory collision(记忆碰撞),检索能力下降。

这个朴素递推有两个层次的根本问题,后续改进都围绕它们展开:

  • 无法修改过去写下的记忆:状态更新 St=St1+vtktS_t = S_{t-1} + v_t k_t^\top 只增不改, 一旦某个 key-value 被写入,之后没有任何机制擦除或修正它;想覆盖一条旧关联, 只能等同一个 key 再次出现,且新值仍以叠加方式加入,旧值依然留在状态里。 这引出第一个改进方向——gating(门控遗忘),即"该遗忘什么"。
  • 数值无上界、状态会发散StS_t累加和,若 vtktv_t k_t^\top 的幅度不衰减, 随序列增长 St\|S_t\| 单调增大,容易数值溢出、训练不稳定;同时在信息层面, 状态被越写越满,多条信息叠加(superposition)相互干扰,精确检索变难。 这引出第二个改进方向——delta rule,即"该以多强的方式覆盖旧记忆": 既能定向覆盖单条关联,又能避免状态无限累积。

2. 两条改进路线

2.1 Mamba2:门控衰减

Mamba2 [3](其前身为 Mamba [4])给状态加一个 数据相关的标量衰减 αt(0,1)\alpha_t \in (0,1)St=αtSt1+vtkt,ot=Stqt S_t = \alpha_t S_{t-1} + v_t k_t^\top, \qquad o_t = S_t q_t

式 2|Mamba2 的门控衰减递推:标量 αt\alpha_t 对整体状态做乘性衰减。

αt\alpha_t 控制"忘记多少"。缺点:对所有 key-value 关联一视同仁地衰减—— 想忘记某一条旧关联,只能把全部关联一起衰减,不精准

2.2 DeltaNet:Delta 更新规则

Delta 规则来自在线学习(Widrow-Hoff)[5],对应的 线性 Transformer 即 DeltaNet [6]。其核心是:写入前,先用当前 key 把状态中对应的旧值读出并擦掉,再写入新值;新值是当前输入值与旧值的线性组合, 组合比例由写入强度 βt(0,1)\beta_t \in (0,1) 决定:

旧值是 vtold=St1ktv_t^{\text{old}} = S_{t-1} k_t;把旧值替换为 新值 vtnew=βtvt+(1βt)St1ktv_t^{\text{new}} = \beta_t v_t + (1-\beta_t)\,S_{t-1} k_t (当前输入值 vtv_t 与旧值的线性组合),得到:

St=St1(Iβtktkt)+βtvtkt S_t = S_{t-1}\big(I - \beta_t k_t k_t^\top\big) + \beta_t v_t k_t^\top

式 3|DeltaNet 的 delta 更新规则:先擦旧值再写新值。

其中转移矩阵 (Iβtktkt)(I - \beta_t k_t k_t^\top)广义 Householder 矩阵ktk_t 已归一化时 ktktk_t k_t^\top 是到 ktk_t 方向的投影), 实现"定向擦除 + 定向写入"。

  • 优点:能定向修改单条 key-value,在 in-context retrieval 类任务上很强;
  • 缺点:一次只改一条,无法快速清空整段过期记忆(例如上下文切换时)。

3. Gated Delta Rule:统一两者

GDN 的递推公式把门控衰减delta 更新乘到一起 [1]

St=St1(αt(Iβtktkt))+βtvtkt S_t = S_{t-1}\Big(\alpha_t\big(I - \beta_t k_t k_t^\top\big)\Big) + \beta_t v_t k_t^\top

式 4|Gated Delta Rule:门控衰减与 delta 更新的统一递推。

  • αt(0,1)\alpha_t \in (0,1)整体遗忘门,让状态快速擦除(如上下文切换);
  • βt(0,1)\beta_t \in (0,1)写入强度,用 delta 结构对单条关联做精准覆盖。

两种机制互补:门控负责"大范围重来",delta 负责"局部精修"。

3.1 αt\alpha_tβt\beta_t 如何从 hidden state 算出

二者都是当前 token 的 hidden state xtx_t 经线性投影 + 激活得到的逐头标量 (每个 head 一个值),不是固定超参数:

βt=σ(Wβxt)(0,1)αt=exp(softplus(Wαxt+bα)A)(0,1) \beta_t = \sigma\!\big(W_\beta x_t\big) \in (0,1) \qquad \alpha_t = \exp\!\big(-\mathrm{softplus}(W_\alpha x_t + b_\alpha)\cdot A\big) \in (0,1)

式 5|αt\alpha_tβt\beta_t 的逐头参数化:均由 hidden state 线性投影 + 激活得到。

各量的 shapeBB 为 batch,TT 为序列长度,dd 为 hidden size,HH 为 head 数):

xtRd,Wβ,WαRH×d,bαRH,ARH x_t \in \mathbb{R}^{d},\qquad W_\beta, W_\alpha \in \mathbb{R}^{H \times d},\qquad b_\alpha \in \mathbb{R}^{H},\qquad A \in \mathbb{R}^{H}

  • 线性投影后得到 HH 维向量(每个 head 一个分量), Wβxt,WαxtRHW_\beta x_t,\; W_\alpha x_t \in \mathbb{R}^{H}
  • 激活是逐元素的,故 βt,αtRH\beta_t, \alpha_t \in \mathbb{R}^{H},即 per-head 标量
  • 展到整个 batch/序列则是 β,αRB×T×H\beta, \alpha \in \mathbb{R}^{B \times T \times H}
  • 递推式(Eq. 4)中 αt\alpha_t 作用在状态 StRdv×dkS_t \in \mathbb{R}^{d_v \times d_k} 上时按 head 广播 (每个 head 的 SS 乘各自的标量),βt\beta_t 同样逐 head 作用于 delta 项。

  • βt\beta_t(写入强度):即 xtx_t 经过与权重 WβW_\beta(实现中的 b_proj)的线性投影后, 经过 sigmoid,天然落在 (0,1)(0,1),表示"这条新 key-value 该以多强的力度覆盖旧值"。 (官方实现 gated_delta_net.py#L164

  • αt\alpha_t(衰减门):沿用 Mamba2 的参数化(论文脚注 4 明确说明)。 即 xtx_t 经过与权重 WαW_\alpha(实现中的 a_proj / gk_proj)的线性投影后, 加上偏置权重 bαb_\alpha(含 dt_bias),最后经过 softplus; 再乘一个可学习的逐头负参数 eAlog-e^{A_{\log}},最后取指数: (官方实现 gated_delta_net.py#L159

    logαt=eAlogsoftplus(Wαxt+bα),Alog=logARH,AU(0,16) \log \alpha_t = -e^{A_{\log}} \cdot \mathrm{softplus}(W_\alpha x_t + b_\alpha), \qquad A_{\log} = \log A \in \mathbb{R}^{H},\; A \sim \mathcal{U}(0,16)

    这样 logαt0\log\alpha_t \le 0,故 αt(0,1]\alpha_t \in (0,1]:输入决定"衰减多快", 而 AA 控制该 head 衰减速率的基础量级(初始 A(0,16)A \in (0,16),训练可学)。

直觉:βt1\beta_t \to 1 时近乎"硬覆盖",βt0\beta_t \to 0 时几乎不写入; αt1\alpha_t \to 1 表示不遗忘,αt0\alpha_t \to 0 表示整块状态快速清空。 二者都由当前 token 内容动态决定,因此 GDN 能按语义切换/插入等场景自适应地 选择"保留、覆盖还是清空"。

在实现上(官方 gated_delta_net.py),q/k/v 还要额外经过短卷积 + SiLU, q,kq,k 做 L2 归一化以稳定训练;而 α,β\alpha,\beta 只经线性投影,不卷积。

4. 核心问题:长序列递推如何并行?

4.1 两种计算形式

同一个 GDN 算子有两种在数学上等价、但性能差异巨大的计算分解 [10]

形式 计算方式 并行度 适合阶段
Recurrent(递推) 逐 token 更新 StS_t 序列维完全串行 自回归推理 / 解码
Parallel / Chunkwise(并行) 展开成带衰减掩码的注意力矩阵,一次矩阵乘算完;Chunkwise 即先把序列切成大小 CC 的块,块内C×CC\times C 的 Parallel 形式、块间传状态 SS Parallel 全序列并行(O(L2)O(L^2) 中间量);Chunkwise 块内并行、块间串行 Parallel 适合短序列与理论对照;Chunkwise 是训练/prefill 默认
  • Recurrent 与 Parallel 的等价即 state space duality(SSD) [3]
  • Parallel 是"整段不切块"的极端情形(CLC \to L:所有 token 一批并行, 但 L×LL \times L 中间量随序列长度平方增长,长序列下撑不住;
  • Chunkwise 是 Parallel 与 Recurrent 的折中:Parallel 只在块内用 (C×CC\times C,避免平方爆炸),块间靠传状态保持因果——训练/prefill 的实际方案。

(a) 问题:C 长度 token 序列的内部递推是 Householder 连乘

只看一段长度为 CC 的 token 序列内部如何递推(暂不谈分块、不谈块间), 用局部索引 i=1,,Ci = 1, \cdots, C 给这段序列里的 token 编号,并把它的 初始状态记作 S0S_0(这段序列开始前的状态)。每个 token 带自己的参数: key kik_i、value viv_i、遗忘门 αi\alpha_i、写入强度 βi\beta_i

逐 token 的 GDN 递推(初始状态 S0S_0,读完第 ii 个 token 后状态为 SiS_i):

Si=Si1(αi(Iβikiki))+βiviki S_i = S_{i-1}\Big(\alpha_i\big(I - \beta_i k_i k_i^\top\big)\Big) + \beta_i v_i k_i^\top

式 6|块内逐 token 的 GDN 递推(局部索引 ii,初始状态 S0S_0)。

记第 ii 个 token 的转移矩阵(dk×dkd_k \times d_k):

Ai:=αi(Iβikiki)Si=Si1Ai+βiviki A_i := \alpha_i\big(I - \beta_i k_i k_i^\top\big) \qquad\Longleftrightarrow\qquad S_i = S_{i-1}A_i + \beta_i v_i k_i^\top

式 7|单 token 转移矩阵 AiA_i:标量门控 × 广义 Householder 矩阵。

S0S_0 起逐步代入,观察 S0S_0 前面乘了什么:

  • i=1i=1S1=S0A1+β1v1k1S_1 = S_0A_1 + \beta_1 v_1 k_1^\top
  • i=2i=2S2=S0A1A2+β1v1k1A2+β2v2k2S_2 = S_0A_1A_2 + \beta_1 v_1 k_1^\top A_2 + \beta_2 v_2 k_2^\top
  • i=3i=3S3=S0A1A2A3+β1v1k1A2A3+β2v2k2A3+β3v3k3S_3 = S_0A_1A_2A_3 + \beta_1 v_1 k_1^\top A_2A_3 + \beta_2 v_2 k_2^\top A_3 + \beta_3 v_3 k_3^\top

归纳即得([1] 的带门控版本,读完前 rr 个 token,1rC1 \le r \le C):

Sr=S0Pr+Hr S_r = S_0\,P_r + H_r

Pr=i=1rαi(Iβikiki),Hr=i=1rβivikij=i+1rαj(Iβjkjkj) P_r = \prod_{i=1}^{r}\alpha_i\big(I - \beta_i k_i k_i^\top\big), \qquad H_r = \sum_{i=1}^{r}\beta_i v_i k_i^\top\prod_{j=i+1}^{r}\alpha_j\big(I - \beta_j k_j k_j^\top\big)

式 8|块内归纳形式:初始状态项 S0PrS_0P_r 与写入累积项 HrH_rPrP_rrrAiA_i 的连乘。

关键观察:初始状态 S0S_0rrAiA_i 依次作用, 每个 Ai=αi(Iβikiki)A_i = \alpha_i(I - \beta_i k_i k_i^\top) = 标量门控 × 广义 Householder 矩阵kik_i 归一化时 kikik_ik_i^\top 是到 kik_i 方向的投影)。 所以"delta 段内是 Householder 连乘" = 递推展开后初始状态被一串 (带标量权重的)Householder 变换累积作用。 其中 βi\beta_i 是每个 Householder 的"强度"(写入强度), αi\alpha_i 是每个变换的"权重"(遗忘强度)。

为什么这让并行变难Pr=A1A2ArP_r = A_1A_2\cdots A_r 表面上看必须按顺序连乘 rr 次, 且每次都是 dk×dkd_k \times d_k 矩阵乘 → 段内是纯串行、且昂贵。 这正是 delta 规则没法像纯累加情形(见下 (b1))那样直接写成一次因果注意力的根源。

这段 CC 长度的序列,放到完整序列里就是"第 tt 块"(chunk): S0S_0 对应块初始状态 S[t]S_{[t]}(= 上一块末状态),局部索引 ii 对应 全局时间步 j=(t1)C+ij = (t-1)\cdot C + i。下文 (b)(c)(d) 恢复 [t][t] 记号。

(b) 解法:C 内部如何从串行变成矩阵并行

(b1) 对照组:纯累加为何"天然"可并行——Parallel 形式的由来

同样的 CC 个 token,如果状态更新是纯累加(普通线性注意力 Si=Si1+vikiS_i = S_{i-1} + v_ik_i^\top,没有擦除项),CC 步串行展开后是一个 外积求和

SC=S0+i=1Cviki=S0+VK S_C = S_0 + \sum_{i=1}^{C} v_ik_i^\top = S_0 + V^\top K

式 9|纯累加的并行形式:外积求和与顺序无关,一步折叠成 VKV^\top K

为什么它能一步并行:求和的每一项 vikiv_ik_i^\top处理顺序无关—— 标量加法满足交换律,先加谁后加谁结果一样。于是 CC 步串行折叠成 一次矩阵乘 VKV^\top K:所有 token 的外积同时算、最后一起加。 这就是 §4.1 里 Parallel 形式的来源Eq. 9)——把 L×LL\times L 的注意力矩阵 (QKΓ)V(QK^\top \odot \Gamma)V 整个一次算出来。

放到完整序列(分块视角)里,输出分解为两项:

  • 块间项 Q[t]S[t]Q_{[t]}S_{[t]}^\top:用上一块传来的状态读出, 把块之前的全部上下文带进来;
  • 块内项 (Q[t]K[t]M)V[t](Q_{[t]}K_{[t]}^\top \odot M)V_{[t]}:块内因果注意力, token rr 能看到块内 iri \le r 的所有 token(由 MM 保证)。

这正是 Mamba2/SSD 的分块算法 [3]但 GDN 走不到这一步—— delta 的擦除项打破了这个美好的交换律,见 (b2)。

(b2) delta 的障碍:擦除操作不可交换,叠加顺序 = 因果顺序

delta 规则多了"定向擦除":每个 token 的 Householder 变换 AiA_i 作用在之后所有 token 的写入项上((a) 的展开式里 k1k^1 后面跟着 A2A3A_2A_3)。 矩阵乘法不交换——A1A2A2A1A_1A_2 \ne A_2A_1——所以不能像累加那样随便重排。 CC 步的因果顺序"藏"在连乘的顺序里,这就是串行的本质障碍。

(b3) 关键一步:把顺序依赖"装进"一个 C×CC\times C 三角矩阵

观察 (a) 展开式的结构:影响第 ii 个写入项的,只有排在它之后的变换 Ai+1ArA_{i+1}\cdots A_r;而这些变换的交叉效应全部由两两内积 kakbk_a^\top k_baba \ne b)组合而成。也就是说:

递推里真正"串行"的只是系数的累积,而系数累积 = 一个 带因果掩码的 C×CC \times C 下三角系统的解。

这正是 WY representation(WY 表示)[7] 的内容:存在一组只依赖本块 k,βk,\beta(不依赖 S0S_0、不依赖顺序求解)的 辅助向量 wi,uiw_i, u_i,使得

P=IWKRdk×dk,H=UKRdv×dk P = I - W^\top K \in \mathbb{R}^{d_k \times d_k}, \qquad H = U^\top K \in \mathbb{R}^{d_v \times d_k}

式 10|WY 表示:把 Pr,HrP_r,H_r 改写为辅助矩阵 W,UW,U 与 key 矩阵 KK 的乘积。

其中 WRC×dkW \in \mathbb{R}^{C \times d_k}URC×dvU \in \mathbb{R}^{C \times d_v} 按块内 token 顺序堆叠。为什么这就算并行了wi,uiw_i, u_i 的定义虽然带递推形式, 但它只依赖本块内部ki,vi,βik_i, v_i, \beta_i,与 S0S_0 无关—— 于是整个"解系数"的过程变成解一个 C×CC\times C单位下三角线性系统

T=[I+strictLower(diag(β)KK)]1diag(β),W=TK,U=TV T = \Big[I + \mathrm{strictLower}\big(\mathrm{diag}(\beta)\,K K^\top\big)\Big]^{-1}\mathrm{diag}(\beta), \qquad W = T K, \quad U = T V

式 11|辅助矩阵 TT 由单位下三角系统解出,进而得到 W=TKW = TKU=TVU = TV

"串行 → 并行"的三步对应

串行视角(逐 token) 并行视角(整块矩阵)
ii 步先读 Si1S_{i-1}、擦除、再写入 系数解 TT 一次算出(下三角系统)
擦除顺序 = 因果顺序 因果性由 strictLower 掩码保证
Pr=A1ArP_r = A_1\cdots A_r 连乘 W=TKW = TKU=TVU = TV 两次矩阵乘

其中因果掩码 strictLower(严格下三角)保证每个 token 只吸收之前 token 的 擦除效应——用矩阵结构替代了时间顺序。加 II 得单位下三角,求逆 数值稳定且可并行(solve_tril)。

KKKK^\top 是块内 Gram 矩阵,与 KV cache 无关: 它是本块 CC 个 token 的 key 两两内积C×CC\times C), 把 delta 连乘展开产生的所有交叉项 kakbk_a^\top k_b 一次算齐。 KK 是每次前向重新投影算出的中间量,算完即丢; 不是缓存历史 key 的 KV cache——GDN 推理时没有随长度增长的 key cache, 只有固定大小的递推状态 SSdv×dkd_v\times d_k)和短卷积窗口状态。

真实实现(FLA / Qwen3.x)把"KK^\top + 三角求逆"融合成一个 kernel:

关键收益——运算数量对比(每个 head、每块 CC 个 token):

串行基线 Chunkwise 并行 阶段
CC 步串行状态更新 KKKK^\top(Gram 矩阵,C×CC\times C)× 1 intra,全块并行
(串行深度 = CC ② 三角求逆得 TTC×CC\times C)× 1 intra
W=TKW = TK × 1 intra
U=TVU = TV(蕴含 H=UKH = U^\top K)× 1 intra
Vnew=UWS[t]V^{\text{new}} = U - WS_{[t]} × 1 inter
S[t+1]=γCS[t]+KVnewS_{[t+1]} = \gamma^C S_{[t]} + K^\top V^{\text{new}} × 1 inter
  • 串行深度从 CC 降到:intra 的 ①–④(与 S[t]S_{[t]} 无关,所有块并行)+ inter 每块 2 次矩阵乘、块间 L/CL/C 步串行
  • 注意 PP 不需要显式算出——P=IWKP = I - W^\top K 的作用已蕴含在 WW 里,inter 循环只消费 W,UW, U
  • 求逆规模只有 C×CC \times CC64C \le 64,而非 d×dd \times d),且只依赖本块, 所以可以对所有块并行完成——这就是"块内并行"的具体来源。

效率结论(以 C=64C=64、head_dim d=128d=128 的一段为例):

指标 串行递推 Chunkwise 变化
关键路径(串行深度) C=64C = 64 步逐 token 更新 ≈ 2(intra 并行 1 批 + inter 2 次矩阵乘) ~32×
FLOPs 总量 2Cd22.1M\approx 2Cd^2 \approx 2.1\text{M} 7.4M\approx 7.4\text{M}(含 C3/3C^3/3 求逆与 Gram) 多算 ~3.5×
算子形态 每步 state(d×dd\times d向量 全部是矩阵 × 矩阵 tensor core 可用

解读:chunkwise 用 ~3.5× 的额外 FLOPs(可全并行、tensor core 高利用率) 换来 ~32× 的关键路径缩短。串行递推每步是"矩阵乘向量",既无法吃满 tensor core,也无法并行;chunkwise 把所有重活换成大矩阵乘, FLOPs 变多但每一步都高效——这正是"理论复杂度更高、实际吞吐更快"的典型例子。

(c) 两阶段结构:块内并行 vs 块间串行

把上面的准备和状态推进摆在一起,就得到 GDN 训练/prefill 的完整两阶段结构:

阶段 做什么 含求逆吗 并行性
块内准备(intra) kkt → 求逆得 TTW,UW,U C×CC\times C 所有块并行,无跨块依赖
块间推进(inter) W,UW,U 逐块更新状态 SS 块间串行L/CL/C 步)
  • 求逆只出现在 intra,且与循环体解耦:进入块间循环前,所有块的 T,W,UT, W, U 已算好。
  • 块间循环本身不含求逆,只做矩阵乘与状态累加,所以很轻。

(d) 块间状态递推(真实循环体)

有了 W,UW, U,块间状态递推在内核 chunk_gated_delta_rule_fwd_h 中完成(Triton kernel chunk_gated_delta_rule_fwd_kernel_h_blockdim64)。 它在 for i_t in range(NT) 循环里逐块做三件事:

  1. 用上一块状态算新 valueV[t]new=U[t]W[t]S[t]V^{\text{new}}_{[t]} = U_{[t]} - W_{[t]}\,S_{[t]} (代码 b_v = u - dot(w, h))。
    • 写入强度 β\beta 在这里不再显式出现:它已在 intra 阶段融进 W,UW, Urecompute_w_ub_b = tl.load(p_b)β[t]\beta_{[t]},见 wy_fast.py#L77, 应用于 wy_fast.py#L88-L89), 所以循环体里用 W[t]S[t]W_{[t]}S_{[t]} 就完成了"定向擦除"。
  2. 施加门控衰减并更新状态
    • 衰减强度在内核里是 g(对数域),本块末值 b_g_last、块内各位置 b_g, 均以 exp2 还原成 γ\gamma(见 b_g_last / b_g 的加载): γ[t]r=j=1rα[t]j=2g[t]r,γ[t]C=2g[t]last \gamma^{r}_{[t]} = \prod_{j=1}^{r}\alpha^j_{[t]} = 2^{\,g^r_{[t]}},\qquad \gamma^{C}_{[t]} = 2^{\,g^{\text{last}}_{[t]}}
    • 先对状态乘块末衰减因子 b_h1 *= b_g_last, 再以 KK^\top 写入 b_h1 = tl.dot(b_k, b_v, b_h1)S[t+1]=diag(γ[t]C)S[t]+K[t]V[t]new S_{[t+1]} = \mathrm{diag}(\gamma^{C}_{[t]})\,S_{[t]} + K_{[t]}^\top V^{\text{new}}_{[t]}
  3. 存下本块末状态(同时把 VnewV^{\text{new}} 写入 v_newSAVE_NEW_VALUE),供后续读出与下一块使用。

变量对照表(内核 chunk_delta_h.py):

数学量 内核变量 说明
写入强度 β[t]\beta_{[t]} b_b(intra)/ b_proj → sigmoid 已在 recompute_w_u 进入 W,UW,U,循环体无
衰减 α[t]r\alpha^r_{[t]} b_g(log2 域) 块内各位置
块末累积衰减 γ[t]C\gamma^C_{[t]} b_g_last exp2\exp_2 后乘到状态上
状态 S[t]S_{[t]} b_h1(K≤64;K>64 拆 b_h2..4 分块于 K 维
新 value V[t]newV^{\text{new}}_{[t]} b_v / v_new u - w·h
定向擦除 W[t]S[t]W_{[t]}S_{[t]} dot(w, h) Vnew=UWhV^{\text{new}} = U - Wh 体现
写入 KVnewK^\top V^{\text{new}} tl.dot(b_k, b_v, b_h1) 累加进状态

直观理解:β\beta 是"写入强度"、α\alpha 是"遗忘强度"。 在块间循环里,β\beta 已"固化"进 W,UW,U(所以公式里只看到 W,UW,U), α\alpha 则以 g(log2 域累积)的形式显式对状态做乘性衰减。

注意:真实实现的 chunk size 限制在 {16,32,64}\{16,32,64\}(默认 64,见 chunk.py#L539), q,kq,k 每个 head 维度 KKvv 的 head 维度 VV256\le 256; 状态存为 [B, HV, K, V](每个 value head 一份)。

4.4 prefill 不逐 token,只有 decode 才逐 token

推理要分成两个阶段,不能一概说"推理就是逐 token 递推"

  • Prefill(处理整段 prompt):输入 {x1,,xL}\{x_1,\cdots,x_L\} 一次性已知, 因此直接用 chunkwise:把 prompt 切成 L/CL/C 个 chunk, 块内矩阵乘并行、块间串行传状态。例如 prompt 长 256256、chunk size C=64C=64, 串行步数只有 256/64=4256/64 = 4 步,不是 256 步
  • Decode(自回归逐 token 生成):新 token 一个一个来,没有"未来 token"可凑成整块, 只能退回到 GDN 的原始 recurrent 形式,每来一个 token 更新一次状态:

St=St1(αt(Iβtktkt))+βtvtkt,ot=Stqt S_t = S_{t-1}\Big(\alpha_t(I - \beta_t k_t k_t^\top)\Big) + \beta_t v_t k_t^\top, \qquad o_t = S_t q_t

式 12|decode 阶段的原始 recurrent 形式:每来一个 token 更新一次状态。

每生成一个 token:读一个 kt,vt,qtk_t,v_t,q_t → 更新状态 StS_t → 读出 oto_t。 单步计算量与已生成长度无关(不随序列增长),因此 decode 是 O(1)O(1) 每 token、 显存恒定(不需要随长度增长的 KV cache)。

官方实现里也一眼可见: mode = 'fused_recurrent' if hidden_states.shape[1] == 1 else self.mode—— 只有序列长度为 1(decode)才走 recurrent,否则走 chunk

所以对"长序列上的递推过程"的准确回答: 训练 / prefill:chunkwise,块内并行、块间串行,不是纯逐 token; decode:确实逐 token 串行更新状态,但每步是常数开销, 且状态大小固定(dv×dkd_v \times d_k),与序列长度无关。

TODO

待核实(尚未在论文中找到明确分析):

  • [ ] CC 的选取是否有理论分析? 已核对 GDN 论文 [1]
    全文,**未见**对 chunk/block size 的定量分析(延迟、吞吐、精度随 CCC 的变化)。
    DeltaNet 并行论文 <sup class="cite"><a href="#ref-8">[8]</a></sup> 与 Mamba2/SSD <sup class="cite"><a href="#ref-3">[3]</a></sup>
    **尚未逐页核对**,需进一步确认是否讨论 block size。
    目前关于"CCC 太大求逆贵、太小矩阵乘碎"的说法主要来自实现经验与 kernel 约束,
    相关推导见本页
    [§4.2 块内如何并行](gated-delta-network.md#42-块内如何并行wy-representationwy-表示把递推连乘变成少数几次矩阵乘法)。
    
  • [ ] 若找到,补充:CC 的最优取值是否与 head dim、序列长度、硬件相关。

5. 架构与实现细节(Qwen3.x / FLA)

  • Block 设计:沿用 Llama 宏观结构(token mixer + SwiGLU MLP), 用 gated delta rule 替换 self-attention。其中 q,k,vq,k,v 由 hidden state xtx_t 经线性投影 Wq,Wk,WvW_q,W_k,W_v、短卷积、SiLU 得到, 并对 q,kq,k 做 L2 归一化以稳定训练; α,β\alpha,\beta 则由 Wα,WβW_\alpha,W_\beta 只做线性投影(不卷积), 再按 §3.1 的激活得到。输出经归一化与门控后过 WoW_o
  • GVA(Grouped Value Attention)q,kq,kHH 个 head,vvHVHH_V \ge H 个 head(HVH_VHH 的整数倍),类似 GQA; 状态按 value head[B, H_V, K, V]
  • 门控融合α\alpha-exp(A_log)·softplus(g+dt_bias) 与块内 cumsum 在 kernel 内融合(use_gate_in_kernel),β\beta 的 sigmoid 也可融合 (use_beta_sigmoid_in_kernel);β(0,2)\beta \in (0,2) 时启用 allow_neg_eigval 以允许负特征值 [9]
  • 混合模型:GDN 层与 sliding window attention(SWA)Mamba2 层堆叠, 兼顾训练效率与长上下文表现(GatedDeltaNet-H1/H2)。Qwen3.x 系列即采用 GDN + 全注意力/门控注意力的混合层。

本页公式与实现细节对齐以下固定 commit 的代码,而非仅凭论文伪代码:

上面正文中的行号链接均基于上述 commit;若仓库更新,行号可能漂移。

6. 与相关工作的关系

模型 状态更新 遗忘 更新方式
Linear Attention St1+vtktS_{t-1} + v_t k_t^\top 纯累加
Mamba2 αtSt1+vtkt\alpha_t S_{t-1} + v_t k_t^\top 全局标量衰减 直接写入
DeltaNet St1(Iβtktkt)+βtvtktS_{t-1}(I-\beta_t k_t k_t^\top) + \beta_t v_t k_t^\top delta 定向覆盖
Gated DeltaNet St1αt(Iβtktkt)+βtvtktS_{t-1}\alpha_t(I-\beta_t k_t k_t^\top) + \beta_t v_t k_t^\top 全局衰减 delta 定向覆盖

GDN 是 DeltaNet 的推广:αt1\alpha_t \equiv 1 时退化为 DeltaNet; 去掉 delta 项(βtktkt\beta_t k_t k_t^\top 部分)则退化为 Mamba2 式门控线性注意力。

术语澄清

  • GDN 的递推是逐 token 的(数学定义如此),但训练与 prefill 用 chunkwise 并行, 只有 decode(逐 token 生成)时才是严格逐 token 串行。
  • 块内(chunk)与块间(cross-chunk):块内用矩阵乘并行; 块间必须串行传状态,这是因果性的必然代价。
  • chunk boundary ≠ sequence boundary:同一序列内部切块要传状态; packing 的不同独立序列之间必须在边界重置状态,否则信息泄漏。
  • 理论复杂度低 ≠ 实际延迟低:WY representation(WY 表示)/ UT transform 的意义正是把 数学上串行或昂贵的运算搬上 tensor core。

参考文献

[1] YANG S, KAUTZ J, HATAMIZADEH A. Gated delta networks: improving Mamba2 with delta rule[C]//Proceedings of the 13th International Conference on Learning Representations (ICLR). 2025. arXiv:2412.06464. https://arxiv.org/abs/2412.06464

[2] KATHAROPOULOS K, VYAS A, PAPPAS D, et al. Transformers are RNNs: fast autoregressive transformers with linear attention[C]//Proceedings of the 37th International Conference on Machine Learning (ICML). 2020. arXiv:2006.16236. https://arxiv.org/abs/2006.16236

[3] DAO T, GU A. Transformers are SSMs: generalized models and efficient algorithms through structured state space duality[C]//Proceedings of the 41st International Conference on Machine Learning (ICML). 2024. arXiv:2405.21060. https://arxiv.org/abs/2405.21060

[4] GU A, DAO T. Mamba: linear-time sequence modeling with selective state spaces[J/OL]. arXiv preprint arXiv:2312.00752, 2023. https://arxiv.org/abs/2312.00752

[5] WIDROW B, HOFF M E. Adaptive switching circuits[C]//IRE WESCON Convention Record. 1960.

[6] SCHLAG I, IRIE K, SCHMIDHUBER J. Linear transformers are secretly fast weight programmers[C]//Proceedings of the 38th International Conference on Machine Learning (ICML). 2021. arXiv:2102.11174. https://arxiv.org/abs/2102.11174

[7] BISCHOF C H, VAN LOAN C. The WY representation for products of Householder matrices[C]//SIAM Conference on Parallel Processing for Scientific Computing. 1985.

[8] YANG S, WANG B, ZHANG Y, et al. Parallelizing linear transformers with the delta rule over sequence length[C]//Advances in Neural Information Processing Systems (NeurIPS). 2024. arXiv:2406.06484. https://arxiv.org/abs/2406.06484

[9] GRAZZI R, SIEMS J, ZELA A, et al. Unlocking state-tracking in linear RNNs through negative eigenvalues[C]//Proceedings of the 13th International Conference on Learning Representations (ICLR). 2025. arXiv:2411.12537. https://arxiv.org/abs/2411.12537

[10] YANG S, WANG B, SHEN Y, et al. Gated linear attention transformers with hardware-efficient training[C]//Proceedings of the 41st International Conference on Machine Learning (ICML). 2024. arXiv:2312.06635. https://arxiv.org/abs/2312.06635


© 2026 Yang Huan · yanghuan9812@qq.com

results matching ""

    No results matching ""