引言#
自 Qwen3-Next 发布以来,Gated Delta Network (GDN) 已经成为 Qwen 全系列的主力注意力层 —— 从 Qwen3-Next-80B-A3B 一路延伸到后续推出的 Qwen3.5 / Qwen3.6 系列。随着模型规模扩展到 397A17B / 122A10B / 35B / 27B,上下文长度突破 256K,GDN 这一模块在端到端训练与推理中的开销也变得不可忽视。
今天我们正式开源 FlashQLA : 一个基于 TileLang 实现的高性能线性注意力算子库。FlashQLA 将 GDN Chunked Prefill 的前向和反向 进行了合理的算子融合与性能优化,在 NVIDIA Hopper 上实现多场景相较于 FLA triton Kernel 2-3× 前向加速 和 2× 反向加速。对于预训练场景和端侧 agentic 推理效率提升明显。
本次发布的核心亮点:
Gate 驱动的自动化卡内序列并行。利用 GDN gate 的指数衰减性质,FlashQLA 在 TP、长序列、小头数等场景下自动开启卡内序列并行,提高 GPU SM 利用率。
硬件友好的代数改写。对 GDN Chunked Prefill 的前向和反向流程进行一定程度的改写,在不影响数值精度的前提下有效降低了 Tencosr Core、 CUDA Core 及 SFU 开销。
Tilelang fused warp-specialized kernels。我们没有沿用每步一个独立 kernel 的拆分方式,也没有将整个计算流程融合为一个 kernel,而是考虑序列并行和 backward 的需求,使用 Tilelang 语言构建出几个关键的 fused kernel,并通过手动的 warpgroup specialization 实现数据搬运、Tensor Core 计算与 CUDA Core 计算的重叠。
FlashQLA 代码 和 benchmark 均开源在 github.com/QwenLM/FlashQLA。
FLA GDN Chunked Prefill 面临的主要问题#
首先让我们回顾一下 GDN Chunked Prefill 的前向计算流程,以 chunk idx $i$ 为例:
- $A_i \gets \left(I+\mathrm{StrictLower}\left( \mathrm{diag}(\beta_i)(\Gamma_i \odot K_iK_i^\intercal) \right)\right)^{-1}$
- $\left\{\begin{aligned} W_i &\gets A_i\mathrm{diag}(\beta_i)\mathrm{diag}(\gamma_i)K_i \\ U_i &\gets A_i\mathrm{diag}(\beta_i)V_i \end{aligned}\right.$
- $\left\{\begin{aligned} V_i’ &\gets U_i-W_iS_i \\ S_{i+1} &\gets \gamma_{i,C-1}S_i + K_i^\intercal\mathrm{diag}\left(\frac{\gamma_{i,C-1}}{\gamma_i}\right)V_i’ \end{aligned}\right.$
- $O_i \gets \mathrm{diag}({\gamma})Q_iS_i + \left(\mathrm{Lower}(\Gamma_i) \odot Q_iK_i^\intercal\right)V_i'$
不考虑 gate 的预处理和 CP,该流程的每一步在 FLA 中都对应一个 kernel。该流程在效率上有两个主要的问题:
- 以上大多为 memory-bounded kernel,在流程中需要反复读取 $K$、$V$ 等数据,而 $W$、$U$、$S$ 作为中间变量也需要写入 HBM 再由下一个 kernel 读取,访存开销较大。
- SSM state 的递推性质导致对应第三步
chunk_gated_delta_rule_fwd_kernel能同时开出的 thread block 数量仅为batch_size * num_heads,在小模型、小 batch 或 TP 场景下 GPU 利用率较低。
这两个问题的解法是相互矛盾的。对于第一个问题,最直观的解法是写一个 fully-fused kernel,所有数据只做一次访存,所有的中间变量也都收到片上,在 batch_size * num_heads 足够大时这一定是最优的。但这样的方案显然会遇到第二个问题,对于一些端侧用小尺寸模型 batch_size=1 的推理场景,或者对于大模型线上部署开 TP 遇到 coding agent 等长序列输入做 chunked prefill 开不出足够大的 batch 的工况,fully-fused kernel 相比于 FLA 原版实现的加速是有限的。
而第二个问题最早的解决方案来源于DeltaNet如何做序列并行,将长序列拆分为多个子序列,使用 $S_0=0$ 并行递推,再计算一个额外的 $M$ 矩阵用于校正递推结果。这个方案后来被优化为在递推 kernel 前插入一步计算每个子序列的 $S_0$,目前已被合并到 FLA 仓库中。对于 CP rank $j$,其具体的预处理流程为:
- $\left\{\begin{aligned} S^\ast_{j,i+1} &\gets \gamma_{j,i,C-1}S^\ast_{j,i} + K_{j,i}^\intercal \mathrm{diag}\left(\frac{\gamma_{j,i,C-1}}{\gamma_{j,i}}\right) V’_{j,i} \\ M_{j,i+1} &\gets \left( \gamma_{j,i,C-1} I -K_{j,i}^\intercal \mathrm{diag}\left(\frac{\gamma_{j,i,C-1}}{\gamma_{j,i}}\right) W_{j,i} \right) M_{j,i} \end{aligned}\right.$
- $S_{j,0} \gets S^\ast_{j,0} + M_{j,0}S_{j-1,0}$
然而这样序列并行的方案也有其弊端:一是引入的额外计算量较大,递推 $M$ 矩阵的时间复杂度甚至高于 $S$ 矩阵;二是和 fully-fused kernel 相性不好,因为需要先经过矩阵求逆等步骤才能计算每个子序列的 $S_0$。
兼顾访存开销与序列并行的解法#
基于上述两个问题,可以得到一个折中的解法:将 GDN Chunked Prefill 前向计算流程拆分为两个 fused kernel,在其中插入 CP 相关的预处理步骤。再经过一些变换和化简,得到如下计算流程:
- $A_i \gets \left(I+\mathrm{StrictLower}\left( \mathrm{diag}(\beta_i)K_iK_i^\intercal\right)\right)^{-1}$
- CP Preprocess
- 2.1. $\left\{\begin{aligned} X_{j,i} &\gets -\beta_{j,i} A_{j,i}’^\intercal K_{j,i} \\ Y_{j,i} &\gets \gamma_{j,i,C-1} K_{j,i} S^\ast_{j,i} - \mathrm{diag}\left(\frac{\gamma_{j,i,C-1}}{\gamma_{j,i}}\right) V_{j,i} \\ Z_{j,i} &\gets K_{j,i} M_{j,i} \\ S^\ast_{j,i+1} &\gets \gamma_{j,i,C-1}S^\ast_{j,i} + X_{j,i}^\intercal Y_{j,i} \\ M_{j,i+1} &\gets \gamma_{j,i,C-1} \left( M_{j,i} + X_{j,i}^\intercal Z_{j,i} \right) \end{aligned}\right.$
- 2.2. $S_{j,0} \gets S^\ast_{j,0} + M_{j,0}S_{j-1,0}$
- $\left\{\begin{aligned} V_i^\Delta &\gets V_i - \mathrm{diag}(\gamma_i)K_iS_i \\ V_i’ &\gets \left(\Gamma_i \odot A_i\right)\mathrm{diag}(\beta_i)V_i^\Delta \\ S_{i+1} &\gets \gamma_{i,C-1}S_i + K_i^\intercal\mathrm{diag}\left(\frac{\gamma_{i,C-1}}{\gamma_i}\right)V_i’ \\ O_i &\gets \mathrm{diag}({\gamma_i})Q_iS_i + \left(\mathrm{Lower}(\Gamma_i) \odot Q_iK_i^\intercal\right)V_i’ \end{aligned}\right.$
我们还设计了一个简单的数学模型自动计算并行度。设一个序列上的 chunk 数量为 $N$,每个 CP rank 上的 chunk 数量为 $L$。易知步骤 2.1 和 3 的运行时间正比于 $L$,而步骤 2.2 的运行时间正比于 $\frac NL$,因此我们可以取 $L=\lambda \sqrt N$ 使得总时间最短,其中 $\lambda$ 为 batch_size、num_heads 等其他超参数组成的系数。
实际生产中并不总是需要开启卡内序列并行。参考 FLA 原版实现,步骤 3 也可以通过切分 v_head_dim 增加 2-4x 并行度,代价是对 Q 和 K 的冗余访存。根据实测数据,我们仅在 batch_size * num_heads <= 40 和 batch_size * num_heads <= 56 && seq_len >= 8192 这两种情况下开启序列并行。
利用 Gate 衰减性质进一步优化#
回看 GDN 递推公式:
$$S_{i+1} = \alpha_iS_i(I-\beta_ik_ik_i^\intercal)+\beta_iv_ik_i^\intercal$$
对于 $\alpha_i\in(0,1)$,每个 $S_i$ 对后续状态的影响呈指数衰减,因此具备滑动窗口的性质。对于足够长的窗口尺寸 $W$,从 $S_{i-W}=0$ 开始递推即可获取精确的 $S_i$,而不必从 $S_0$ 开始递推。我们将这一过程称为 warmup。在真实数据上,我们发现 60-80% 的线性注意力头上 $\alpha_i$ 不恒为 1,6~8 个 chunk 的 warmup 就足以将 $S_i$ 的误差压低到噪声以下。
由此我们可以针对具备滑窗性质的线性注意力头设计一套更轻量级的 CP preprocess 流程,舍弃对修正量 $M$ 的计算,直接通过 warmup 获得同样精确的子序列 $S_0$:
| C0 | C1 | C2 | C3 | C4 | C5 | C6 | C7 | C8 | C9 | C10 | C11 | C12 | |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| R1 | O | O | O | O | O | ||||||||
| R2 | X | X | O | O | O | O | |||||||
| R3 | X | X | O | O | O | O |
X 表示用零初始状态做 warmup 直到 gate 衰减到足够小之后写出该 CP rank 的 $S_0$,O 表示后续正常的递推计算。每个 rank 的 warmup 长度由一个独立的 kernel 通过统计 gate 决定,该步骤的耗时是可以忽略的。
Tilelang Warp-Specialized Kernel#
我们基于 TileLang,采用 warpgroup specialization 的方式实现:同一个 SM 内包含一个生产者 warpgroup 和三个消费者 warpgroup,通过 shared memory 交换数据,并通过 mbarrier 同步。
前向#
在前向流程中,我们让三个消费者 warpgroup 分别计算 $V’$、$S$ 和 $O$,并通过 ping-pong 结构遮盖计算与访存。
| WG3 | WG2 | WG1 | WG0 | |||
|---|---|---|---|---|---|---|
| WG3/0 | WG3/1 | WG3/2 | ||||
| BAR 0 | LD$Q$ | LD$\gamma$ | ST$O$ | $\gamma, \gamma_{C-1}\gamma^{-1}$ | TC$P = Q K^\intercal$ | |
| BAR 1 | LD$K$ | LD$\beta$ | ST$S_i$ | TC$U = K S_i$ | $\Gamma = L(\gamma I \gamma^{-1})$ $A_\gamma = \Gamma \odot A$ $P_\gamma = s\Gamma \odot P$ | $S_{i+1} = \gamma_{C-1} S_i$ |
| BAR 2 | LD$V$ | $W = \beta (V - \gamma U)$ | TC$O = Q S_i$ | |||
| BAR 3 | LD$A$ | TC$V^\Delta = A_\gamma W$ | $O = s\gamma O$ | |||
| BAR 4 | $V’ = \gamma_{C-1}\gamma^{-1} V^\Delta$ | TC$O = O + P_\gamma V^\Delta$ | ||||
| BAR 5 | TC$S_{i+1} = S_{i+1} + K^\intercal V'$ | |||||
注意:
- 每个 chunk 上输出 $S$ 仅作 debug 用,一般只输出 $O$ 和最后一个 chunk 的 $S$。
序列并行预处理#
刚才说到序列并行的预处理分为原始做法(计算 $M$ 和 $S$)和滑动窗口(仅计算 $S$)两种情况。我们设计了一个 fused kernel 可以同时处理这两种情况:
| WG3 | WG2 | WG1 | WG0 | |||
|---|---|---|---|---|---|---|
| WG3/0 | WG3/1 | WG3/2 | ||||
| BAR 0 | LD$K$ | LD$\gamma$ | $\gamma_{C-1}\gamma^{-1}$ | TC$X = A^\intercal K$ | ||
| BAR 1 | LD$V$ | LD$\beta$ | ST$S_i$ | TC$U = K S_i$ $Y = -\gamma_{C-1}\gamma^{-1} V + \gamma_{C-1} U$ | $X = -\beta X$ | $S_{i+1} = \gamma_{C-1} S_i$ |
| BAR 2 | LD$A$ | $\gamma^\pi = \gamma^\pi \gamma_{C-1}$ | $\gamma^\pi = \gamma^\pi \gamma_{C-1}$ | TC$S_{i+1} = S_{i+1} + X^\intercal Y$ | ||
| BAR 3 | TC$Z^L = K M^L$ TC$M^L = M^L + X^\intercal Z^L$ | TC$Z^R = K M^R$ TC$M^R = M^R + X^\intercal Z^R$ | ||||
注意:
- WG1 和 WG2 的最后两步为计算 $M$ 矩阵的过程,仅在需要时触发。
- 仅反向重算时在每个 chunk 上输出 $S$。
反向#
在反向流程中,我们可以直接套用上一节的序列并行预处理 kernel 重算 $S$ 矩阵;之后把 bwd_dv bwd_dhu bwd_dqkwg bwd_wy 融合到一个 kernel 里,并作相应的代数优化。受片上资源限制,反向 kernel 不设置 multi-stage,而是利用长计算流程遮盖访存。完整的流水线可以在 FlashQLA 仓库 中查看。
| WG3 | WG2 | WG1 | WG0 | ||
|---|---|---|---|---|---|
| BAR 00 | ST$dK$ | TC $P=QK^\intercal$ | $\gamma, \gamma_{C-1}\gamma^{-1}$ | ||
| BAR 01 | TC $dV’=KdS_{i+1}$ $dV’=\gamma_{C-1}\gamma^{-1}dV'$ | $\Gamma=\gamma I \gamma^{-1}$ $P_\gamma=sL(\Gamma)\odot P$ | $dS_i = \gamma_{C-1} dS_{i+1}$ | ||
| BAR 02 | TC $dV’=dV’+P_\gamma^\intercal dO$ | $A_\beta = A \beta$ $A_\gamma = \Gamma \odot A_\beta$ | |||
| BAR 03 | TC$U=KS_i$ | ||||
| BAR 04 | TC $dV=A_\gamma^\intercal dV’$ | $W=V-\gamma U$ | $d\gamma_{C-1}=\sum S_i \odot dS_{i+1}$ | ||
| BAR 05 | ST$dV$ | LD$V$ | $dV_\gamma = -\gamma dV$ $d\gamma = \sum_i dV_\gamma \odot U$ | TC $dA_\gamma = dV’W^T$ TC $V’=A_\gamma W$ | |
| BAR 06 | TC $dP_\gamma = dO V’^\intercal$ | ||||
| BAR 07 | LD$K$ | TC $dK=V’dS_{i+1}^\intercal$ | $dA_\beta = \Gamma \odot dA_\gamma$ $d\gamma = d\gamma + \sum_i dP_\gamma \odot L(P)$ $d\gamma = d\gamma - \sum_j dP_\gamma \odot L(P)$ $dP = sL(\Gamma)\odot dP_\gamma$ | ||
| BAR 08 | $dK=\gamma_{C-1}\gamma^{-1}dK$ $d\gamma_{C-1}=\sum K \odot dK$ $d\gamma = -\sum_i K \odot dK$ | TC $dQ=dOS_i^T$ | |||
| BAR 09 | LD$Q$ | TC $dK=dK+dV_\gamma S_i^\intercal$ | $dQ=s\gamma dQ$ $d\gamma = \sum Q \odot dQ$ | ||
| BAR 10 | LD$S$ | TC $dQ=dQ+dPK$ | |||
| BAR 11 | ST$dQ$ | $d\gamma = d\gamma + \sum_i dA_\beta \odot A \beta$ $d\gamma = d\gamma - \sum_j dA_\beta \odot A \beta$ $d\beta = \sum_j dA_\beta \odot A$ $dA=dA_\beta \beta$ | TC $dS_i = dS_i + K^\intercal dV_\gamma$ | ||
| BAR 12 | TC $dK=dK+dP^\intercal Q$ | ||||
| BAR 13 | TC $dA = -A^\intercal dA A^\intercal$ TC $A_T = KK^\intercal$ | $dO_\gamma=s\gamma dO$ | |||
| BAR 14 | LD$dO$ LD$A$ | $d\beta = d\beta + \sum_i dA \odot A_T$ $dA_T = \beta dA$ $dA_S = dA_T + dA_T^\intercal$ | TC $dS_0 = dS_0 + Q^\intercal dO_\gamma$ | ||
| BAR 15 | TC $dK=dK+dA_S K$ | ||||
Benchmark#
我们在 Qwen3.5 / Qwen3.6 系列的 head 配置上 —— h_v ∈ {64, 48, 32, 24, 16, 8},对应 TP1 至 TP8 —— 与 FLA Triton and FlashInfer baseline(FLA 0.5.0,Triton 3.5.1, FlashInfer 0.6.9, TileLang 0.1.8)做了全面对比。
其中FWD 中测试了不同模型、TP setting下对于不同batch 长度下单Kernel latency,BWD中测试了单次更新中batch内不同总token number与latency的关系。
部分 H200 单层前向结果:
| 模型 / TP | Seqlen | $h_{qk}$ | $h_v$ | FlashQLA | FlashInfer | FLA | vs FLA | vs FI |
|---|---|---|---|---|---|---|---|---|
| 397B/122B TP8 | 1x32768 | 2 | 8 | 0.310ms | 1.653ms | 0.913ms | 2.95× | 5.33× |
| 397B/122B TP8 | 1x16384 | 2 | 8 | 0.184ms | 0.833ms | 0.465ms | 2.53× | 4.53× |
| 397B/122B TP8 | 24576+8192 | 2 | 8 | 0.302ms | 1.242ms | 0.767ms | 2.54× | 4.11× |
| 397B/122B TP4 | 1x32768 | 4 | 16 | 0.486ms | 1.654ms | 1.250ms | 2.57× | 3.40× |
| 397B/122B TP4 | 1x16384 | 4 | 16 | 0.292ms | 0.832ms | 0.623ms | 2.13× | 2.85× |
| 27B TP2 | 1x32768 | 8 | 24 | 0.659ms | 1.616ms | 1.564ms | 2.37× | 2.45× |
| 2B/0.8B TP1 | 1x32768 | 16 | 16 | 0.493ms | 1.640ms | 1.285ms | 2.60× | 3.33× |
| Sym h32 | 1x32768 | 32 | 32 | 0.877ms | 1.554ms | 1.952ms | 2.23× | 1.77× |
加速比随 TP 增大而提升,这是因为 FlashQLA 能够通过卡内的 AutoCP 提高TP,小 num_heads 等场景下 SM 利用率。
使用方式#
FlashQLA 同时提供了对齐 FLA 签名的 high-level API 与底层 fwd / bwd 入口:
import torch
from qla import chunk_gated_delta_rule
o, final_state = chunk_gated_delta_rule(
q=q, # [B, T, H_q, K]
k=k, # [B, T, H_q, K]
v=v, # [B, T, H_v, V]
g=g, # [B, T, H_v]
beta=beta, # [B, T, H_v]
scale=scale,
initial_state=initial_state, # 可选, [B, H_v, K, V]
output_final_state=True,
cu_seqlens=cu_seqlens, # 可选, varlen 支持
)
环境要求:SM90,CUDA 12.8+,PyTorch 2.8+。安装:
git clone https://github.com/QwenLM/FlashQLA.git
cd FlashQLA && pip install -v .
感谢#
FlashQLA 的实现受到 Flash Linear Attention,FlashInfer 与 TileLang 的启发,感谢社区的参考实现。
引用#
如果 FlashQLA 对你的研究有所帮助,欢迎引用:
@misc{flashqla2026,
title = {FlashQLA: Flash Qwen Linear Attention},
author = {Zhang, Chengruidong and Lin, Xi and Jiang, Huiqiang and Wang, Zekun and
Li, Xiao and Cao, Yizhong and Zhuang, Bohan and Men, Rui and Zhang, Jianwei and
Zheng, Bo and Lin, Junyang and Liu, Dayiheng and Zhou, Jingren},
year = {2026},
publisher = {GitHub},
howpublished = {\url{https://github.com/QwenLM/FlashQLA}}
}