FlashQLA: 面向GDN的CP-/Bwd友好型融合线性注意力内核
FlashQLA: CP-/Bwd-Friendly Fused Linear Attention Kernels for GDN
FlashQLA 发布了一组专为梯度下降网络优化的融合线性注意力内核。新内核在设计上对计算模式和后向传播更加友好,旨在提升训练效率。该技术通过优化内核融合策略,改进了注意力机制的计算性能,是提升大规模模型训练速度的关键底层优化。
Qwen 团队发了一篇 fused linear attention 内核的工程论文,目标是把 GDN 架构的推理和反向传播都跑快。做底层优化的工程师值得看一眼,普通开发者可以略过。
Introduction#
Following the release of Qwen3-Next, Gated Delta Network (GDN) has become the workhorse attention layer across the Qwen family — from Qwen3-Next-80B-A3B all the way to the subsequent Qwen3.5 / Qwen3.6 series. As models scale to 397A17B / 122A10B / 35B / 27B and context windows stretch beyond 256K, the overhead of the GDN block in end-to-end training and inference has become non-negligible.
Today we officially open-source FlashQLA: a high-performance linear attention kernel library built on TileLang. FlashQLA applies reasonable operator fusion and performance optimization to the forward and backward passes of GDN Chunked Prefill, achieving 2-3× forward speedup and 2× backward speedup over the FLA Triton kernel across multiple scenarios on NVIDIA Hopper. The efficiency gains are particularly pronounced in pretraining scenarios and edge-side agentic inference.
Key highlights of this release:
Gate-driven automatic intra-card context parallelism. By exploiting the exponential decay property of the GDN gate, FlashQLA automatically enables intra-card CP under TP, long-sequence, and small-head-count settings, improving GPU SM utilization.
Hardware-friendly algebraic reformulation. We reformulate the forward and backward flows of GDN Chunked Prefill to a certain extent, effectively reducing Tensor Core, CUDA Core, and SFU overhead without sacrificing numerical precision.
TileLang fused warp-specialized kernels. Rather than following the step-by-step decomposition into independent kernels, nor fusing the entire computation flow into a single kernel, we take CP and backward requirements into account, use TileLang to build several key fused kernels, and manually implement warpgroup specialization to overlap data movement, Tensor Core computation, and CUDA Core computation.
FlashQLA code and benchmarks are open-sourced at github.com/QwenLM/FlashQLA.
Key Problems in FLA GDN Chunked Prefill#
Let us first review the forward computation flow of GDN Chunked Prefill, taking chunk index $i$i as an example:
- $A_{i} \leftarrow \left(\left(\right. I + S t r i c t L o w e r \left(\right. d i a g \left(\right. \beta_{i} \left.\right) \left(\right. \Gamma_{i} \bigodot K_{i} K_{i}^{\top} \left.\right) \left.\right) \left.\right)\right)^{- 1}$A i←(I+StrictLower(diag(β i)(Γ i⊙K iK i⊺)))−1
- $\left{\right. W_{i} & \leftarrow A_{i} d i a g \left(\right. \beta_{i} \left.\right) d i a g \left(\right. \gamma_{i} \left.\right) K_{i} \ U_{i} & \leftarrow A_{i} d i a g \left(\right. \beta_{i} \left.\right) V_{i}${W iU i←A idiag(β i)diag(γ i)K i←A idiag(β i)V i
- $\left{\right. V_{i} ’ & \leftarrow U_{i} - W_{i} S_{i} \ S_{i + 1} & \leftarrow \gamma_{i , C - 1} S_{i} + K_{i}^{\top} d i a g \left(\right. \frac{\gamma_{i , C - 1}}{\gamma_{i}} \left.\right) V_{i} ’$⎩⎨⎧V i’S i+1←U i−W iS i←γ i,C−1S i+K i⊺diag(γ iγ i,C−1)V i’
- $O_{i} \leftarrow d i a g \left(\right. \gamma \left.\right) Q_{i} S_{i} + \left(\right. L o w e r \left(\right. \Gamma_{i} \left.\right) \bigodot Q_{i} K_{i}^{\top} \left.\right) V_{i}^{'}$O i←diag(γ)Q iS i+(Lower(Γ i)⊙Q iK i⊺)V i′
Ignoring gate preprocessing and CP, each step of this flow corresponds to one kernel in FLA. This flow has two main efficiency problems:
- Most of the above are memory-bound kernels. The flow repeatedly reads $K$K, $V$V and other data, while $W$W, $U$U, $S$S as intermediate variables must be written to HBM and then read by the next kernel, incurring significant memory access overhead.
- The recurrent nature of the SSM state means that the corresponding third step
chunk_gated_delta_rule_fwd_kernelcan only launchbatch_size * num_headsthread blocks simultaneously, resulting in low GPU utilization in small-model, small-batch, or TP scenarios.
The solutions to these two problems are contradictory. For the first problem, the most intuitive solution is to write a fully-fused kernel, where all data is accessed only once and all intermediate variables are kept on-chip. When batch_size * num_heads is large enough, this is certainly optimal. However, such a solution obviously runs into the second problem: for edge-side inference with small models and batch_size=1, or for large-model online deployments with TP where long-sequence inputs from coding agents etc. cannot launch a large enough batch for chunked prefill, the speedup of a fully-fused kernel over the original FLA implementation is limited.
The earliest solution to the second problem comes from how DeltaNet does context parallelism, which splits a long sequence into multiple sub-sequences, uses $S_{0} = 0$S 0=0 to parallelize the recurrence, and then computes an additional $M$M matrix to correct the recurrent results. This scheme was later optimized to insert a step before the recurrence kernel to compute the $S_{0}$S 0 of each sub-sequence, and has now been merged into the FLA repository. For CP rank $j$j, the specific preprocessing flow is:
- $\left{\right. S_{j , i + 1}^{} & \leftarrow \gamma_{j , i , C - 1} S_{j , i}^{} + K_{j , i}^{\top} d i a g \left(\right. \frac{\gamma_{j , i , C - 1}}{\gamma_{j , i}} \left.\right) V ’{j , i} \ M{j , i + 1} & \leftarrow \left(\right. \gamma_{j , i , C - 1} I - K_{j , i}^{\top} d i a g \left(\right. \frac{\gamma_{j , i , C - 1}}{\gamma_{j , i}} \left.\right) W_{j , i} \left.\right) M_{j , i}$⎩⎨⎧S j,i+1∗M j,i+1←γ j,i,C−1S j,i∗+K j,i⊺diag(γ j,iγ j,i,C−1)V’j,i←(γ j,i,C−1I−K j,i⊺diag(γ j,iγ j,i,C−1)W j,i)M j,i
- $S_{j , 0} \leftarrow S_{j , 0}^{*} + M_{j , 0} S_{j - 1 , 0}$S j,0←S j,0∗+M j,0S j−1,0
However, this CP scheme also has its drawbacks: first, it introduces significant extra computation, with the time complexity of recurrently computing the $M$M matrix even exceeding that of the $S$S matrix; second, it does not work well with fully-fused kernels, because matrix inversion and other steps must be performed before the $S_{0}$S 0 of each sub-sequence can be computed.
A Balanced Solution: Fusing Kernels While Enabling Intra-Card CP#
Based on the two problems above, a compromise solution can be derived: split the GDN Chunked Prefill forward computation into two fused kernels, inserting CP-related preprocessing steps between them. After some transformations and simplifications, the following computation flow is obtained:
$A_{i} \leftarrow \left(\left(\right. I + S t r i c t L o w e r \left(\right. d i a g \left(\right. \beta_{i} \left.\right) K_{i} K_{i}^{\top} \left.\right) \left.\right)\right)^{- 1}$A i←(I+StrictLower(diag(β i)K iK i⊺))−1
CP Preprocess * 2.1. $\left{\right. X_{j , i} & \leftarrow - \beta_{j , i} A_{j , i} ’^{\top} K_{j , i} \ Y_{j , i} & \leftarrow \gamma_{j , i , C - 1} K_{j , i} S_{j , i}^{} - d i a g \left(\right. \frac{\gamma_{j , i , C - 1}}{\gamma_{j , i}} \left.\right) V_{j , i} \ Z_{j , i} & \leftarrow K_{j , i} M_{j , i} \ S_{j , i + 1}^{} & \leftarrow \gamma_{j , i , C - 1} S_{j , i}^{} + X_{j , i}^{\top} Y_{j , i} \ M_{j , i + 1} & \leftarrow \gamma_{j , i , C - 1} \left(\right. M_{j , i} + X_{j , i}^{\top} Z_{j , i} \left.\right)$⎩⎨⎧X j,iY j,iZ j,iS j,i+1∗M j,i+1←−β j,iA j,i’⊺K j,i←γ j,i,C−1K j,iS j,i∗−diag(γ j,iγ j,i,C−1)V j,i←K j,iM j,i←γ j,i,C−1S j,i∗+X j,i⊺Y j,i←γ j,i,C−1(M j,i+X j,i⊺Z j,i) * 2.2. $S_{j , 0} \leftarrow S_{j , 0}^{} + M_{j , 0} S_{j - 1 , 0}$S j,0←S j,0∗+M j,0S j−1,0
$\left{\right. V_{i}^{\Delta} & \leftarrow V_{i} - d i a g \left(\right. \gamma_{i} \left.\right) K_{i} S_{i} \ V_{i} ’ & \leftarrow \left(\right. \Gamma_{i} \bigodot A_{i} \left.\right) d i a g \left(\right. \beta_{i} \left.\right) V_{i}^{\Delta} \ S_{i + 1} & \leftarrow \gamma_{i , C - 1} S_{i} + K_{i}^{\top} d i a g \left(\right. \frac{\gamma_{i , C - 1}}{\gamma_{i}} \left.\right) V_{i} ’ \ O_{i} & \leftarrow d i a g \left(\right. \gamma_{i} \left.\right) Q_{i} S_{i} + \left(\right. L o w e r \left(\right. \Gamma_{i} \left.\right) \bigodot Q_{i} K_{i}^{\top} \left.\right) V_{i} ’$⎩⎨⎧V i ΔV i’S i+1O i←V i−diag(γ i)K iS i←(Γ i⊙A i)diag(β i)V i Δ←γ i,C−1S i+K i⊺diag(γ iγ i,C−1)V i’←diag(γ i)Q iS i+(Lower(Γ i)⊙Q iK i⊺)V i’
We also designed a simple mathematical model to automatically determine the degree of parallelism. Let $N$N be the number of chunks in a sequence and $L$L be the number of chunks per CP rank. It is easy to see that the runtime of steps 2.1 and 3 is proportional to $L$L, while the runtime of step 2.2 is proportional to $\frac{N}{L}$L N; therefore we can choose $L = \lambda \sqrt{N}$L=λ N to minimize total time, where $\lambda$λ is a coefficient composed of batch_size, num_heads, and other hyperparameters.
In production, intra-card CP is not always needed. Following the original FLA implementation, step 3 can also increase parallelism by 2-4× via splitting v_head_dim, at the cost of redundant memory access to Q and K. Based on measured data, we enable CP only when batch_size * num_heads <= 40 or batch_size * num_heads <= 56 && seq_len >= 8192.
Further Optimization via Gate Decay#
Revisiting the GDN recurrence:
$$ S_{i + 1} = \alpha_{i} S_{i} \left(\right. I - \beta_{i} k_{i} k_{i}^{\top} \left.\right) + \beta_{i} v_{i} k_{i}^{\top} $$
S i+1=α iS i(I−β ik ik i⊺)+β iv ik i⊺
For $\alpha_{i} \in \left(\right. 0 , 1 \left.\right)$α i∈(0,1), the influence of each $S_{i}$S i on subsequent states decays exponentially, giving it a sliding-window property. For a sufficiently long window size of $W$W, starting computation from $S_{i - W} = 0$S i−W=0 can obtain the accurate $S_{i}$S i, without the need to start from $S_{0}$S 0. We refer to this process as warmup. On real data, we find that $\alpha_{i}$α i is not constantly 1 on 60–80% of linear attention heads, and 6–8 chunks of warmup are sufficient to drive the $S_{i}$S i error below the noise floor.
Therefore, for linear attention heads with the sliding-window property, we can design a lighter CP preprocessing flow that discards the computation of the correction term $M$M and directly obtains an equally accurate sub-sequence $S_{0}$S 0 through warmup:
| 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 denotes warming up with a zero initial state until the gate has decayed sufficiently, then writing out the $S_{0}$S 0 of that CP rank; O denotes subsequent normal recurrent computation. The warmup length for each rank is determined by an independent kernel that collects gate statistics, and the cost of this step is negligible.
TileLang Warp-Specialized Kernel#
We implement FlashQLA in TileLang using a warpgroup-specialization pattern: one producer warpgroup and three consumer warpgroups reside in the same SM, exchange data through shared memory, and synchronize via mbarriers.
Forward#
In the forward pass, the three consumer warp groups compute $V ’$V’, $S$S, and $O$O respectively, overlapping computation and memory traffic through a ping-pong structure.
| WG3 | WG2 | WG1 | WG0 | |
|---|---|---|---|---|
| WG3/0 | WG3/1 | WG3/2 | ||
| BAR 0 | LD$Q$Q | LD$\gamma$γ | ST$O$O | $\gamma , \textrm{ } \gamma_{C - 1} \gamma^{- 1}$γ,γ C−1γ−1 |
| BAR 1 | LD$K$K | LD$\beta$β | ST$S_{i}$S i | TC$U \textrm{ } = \textrm{ } K \textrm{ } S_{i}$U=K S i |
| BAR 2 | LD$V$V | $W \textrm{ } = \textrm{ } \beta \textrm{ } \left(\right. V \textrm{ } - \textrm{ } \gamma \textrm{ } U \left.\right)$W=β(V−γ U) | TC$O \textrm{ } = \textrm{ } Q \textrm{ } S_{i}$O=Q S i | |
| BAR 3 | LD$A$A | TC$V^{\Delta} \textrm{ } = \textrm{ } A_{\gamma} \textrm{ } W$V Δ=A γW | $O \textrm{ } = \textrm{ } s \gamma \textrm{ } O$O=s γ O | |
| BAR 4 | $V ’\textrm{ } = \textrm{ } \gamma_{C - 1} \gamma^{- 1} \textrm{ } V^{\Delta}$V’=γ C−1γ−1 V Δ | TC$O \textrm{ } = \textrm{ } O \textrm{ } + \textrm{ } P_{\gamma} \textrm{ } V^{\Delta}$O=O+P γV Δ | ||
| BAR 5 | TC$S_{i + 1} \textrm{ } = \textrm{ } S_{i + 1} \textrm{ } + \textrm{ } K^{\top} \textrm{ } V^{'}$S i+1=S i+1+K⊺V′ |
Notes:
- $S$S output per chunk is for debugging only; normally only $O$O and the last chunk’s $S$S are output.
CP Preprocessing#
As mentioned earlier, the CP preprocessing splits into two cases: the original approach (computing both $M$M and $S$S) and the sliding-window approach (computing only $S$S). We designed a single fused kernel that handles both:
| WG3 | WG2 | WG1 | WG0 | |
|---|---|---|---|---|
| WG3/0 | WG3/1 | WG3/2 | ||
| BAR 0 | LD$K$K | LD$\gamma$γ | $\gamma_{C - 1} \gamma^{- 1}$γ C−1γ−1 | |
| BAR 1 | LD$V$V | LD$\beta$β | ST$S_{i}$S i | TC$U \textrm{ } = \textrm{ } K \textrm{ } S_{i}$U=K S i $Y \textrm{ } = \textrm{ } - \gamma_{C - 1} \gamma^{- 1} \textrm{ } V \textrm{ } + \textrm{ } \gamma_{C - 1} \textrm{ } U$Y=−γ C−1γ−1 V+γ C−1U |
| BAR 2 | LD$A$A | $\gamma^{\pi} \textrm{ } = \textrm{ } \gamma^{\pi} \textrm{ } \gamma_{C - 1}$γ π=γ π γ C−1 | $\gamma^{\pi} \textrm{ } = \textrm{ } \gamma^{\pi} \textrm{ } \gamma_{C - 1}$γ π=γ π γ C−1 | |
| BAR 3 | TC$Z^{L} \textrm{ } = \textrm{ } K \textrm{ } M^{L}$Z L=K M L TC$M^{L} \textrm{ } = \textrm{ } M^{L} \textrm{ } + \textrm{ } X^{\top} \textrm{ } Z^{L}$M L=M L+X⊺Z L | TC$Z^{R} \textrm{ } = \textrm{ } K \textrm{ } M^{R}$Z R=K M R TC$M^{R} \textrm{ } = \textrm{ } M^{R} \textrm{ } + \textrm{ } X^{\top} \textrm{ } Z^{R}$M R=M R+X⊺Z R |
Notes:
- The last two steps of WG1 and WG2 correspond to the $M$M matrix computation and are triggered only when required.
- $S$S is output per chunk only during backward recomputation.
Backward#
In the backward pass, we reuse the CP preprocessing kernel from the previous section to recompute the $S$S matrix, then fuse bwd_dv, bwd_dhu, bwd_dqkwg, bwd_wy into a single kernel with corresponding algebraic optimizations. Because of on-chip resource constraints, the backward kernel does not use multi-stage pipelining; instead it relies on the long compute chain to hide memory traffic. The full schedule is available in the FlashQLA repo.
| WG3 | WG2 | WG1 | WG0 | |
|---|---|---|---|---|
| BAR 00 | ST$d K$d K | TC$P = Q K^{\top}$P=Q K⊺ | ||
| BAR 01 | TC$d V ’ = K d S_{i + 1}$d V’=K d S i+1 $d V ’ = \gamma_{C - 1} \gamma^{- 1} d V^{'}$d V’=γ C−1γ−1 d V′ | $\Gamma = \gamma \textrm{ } I \textrm{ } \gamma^{- 1}$Γ=γ I γ−1 $P_{\gamma} = s L \left(\right. \Gamma \left.\right) \bigodot \textrm{ } P$P γ=s L(Γ)⊙P | $d S_{i} \textrm{ } = \textrm{ } \gamma_{C - 1} \textrm{ } d S_{i + 1}$d S i=γ C−1d S i+1 | |
| BAR 02 | TC$d V ’ = d V ’ + P_{\gamma}^{\top} \textrm{ } d O$d V’=d V’+P γ⊺d O | $A_{\beta} \textrm{ } = \textrm{ } A \textrm{ } \beta$A β=A β $A_{\gamma} \textrm{ } = \textrm{ } \Gamma \textrm{ } \bigodot \textrm{ } A_{\beta}$A γ=Γ⊙A β | ||
| BAR 03 | TC$U = K S_{i}$U=K S i | |||
| BAR 04 | TC$d V = A_{\gamma}^{\top} \textrm{ } d V ’$d V=A γ⊺d V’ | $W = V - \gamma \textrm{ } U$W=V−γ U | $d \gamma_{C - 1} = \sum \textrm{ } S_{i} \textrm{ } \bigodot \textrm{ } d S_{i + 1}$d γ C−1=∑S i⊙d S i+1 | |
| BAR 05 | ST$d V$d V | LD$V$V | $d V_{\gamma} \textrm{ } = \textrm{ } - \gamma \textrm{ } d V$d V γ=−γ d V $d \gamma \textrm{ } = \textrm{ } \sum_{i} \textrm{ } d V_{\gamma} \textrm{ } \bigodot \textrm{ } U$d γ=∑id V γ⊙U | TC$d A_{\gamma} \textrm{ } = \textrm{ } d V ’ W^{T}$d A γ=d V’W T TC$V ’ = A_{\gamma} \textrm{ } W$V’=A γW |
| BAR 06 | TC$d P_{\gamma} \textrm{ } = \textrm{ } d O \textrm{ } V ’^{\top}$d P γ=d O V’⊺ | |||
| BAR 07 | LD$K$K | TC$d K = V ’ d S_{i + 1}^{\top}$d K=V’d S i+1⊺ | $d A_{\beta} \textrm{ } = \textrm{ } \Gamma \textrm{ } \bigodot \textrm{ } d A_{\gamma}$d A β=Γ⊙d A γ $d \gamma \textrm{ } = \textrm{ } d \gamma \textrm{ } + \textrm{ } \sum_{i} \textrm{ } d P_{\gamma} \textrm{ } \bigodot \textrm{ } L \left(\right. P \left.\right)$d γ=d γ+∑id P γ⊙L(P) $d \gamma \textrm{ } = \textrm{ } d \gamma \textrm{ } - \textrm{ } \sum_{j} \textrm{ } d P_{\gamma} \textrm{ } \bigodot \textrm{ } L \left(\right. P \left.\right)$d γ=d γ−∑jd P γ⊙L(P) $d P \textrm{ } = \textrm{ } s L \left(\right. \Gamma \left.\right) \bigodot \textrm{ } d P_{\gamma}$d P=s L(Γ)⊙d P γ | |
| BAR 08 | $d K = \gamma_{C - 1} \gamma^{- 1} d K$d K=γ C−1γ−1 d K $d \gamma_{C - 1} = \sum \textrm{ } K \textrm{ } \bigodot \textrm{ } d K$d γ C−1=∑K⊙d K $d \gamma \textrm{ } = \textrm{ } - \sum_{i} \textrm{ } K \textrm{ } \bigodot \textrm{ } d K$d γ=−∑iK⊙d K | TC$d Q = d O S_{i}^{T}$d Q=d O S i T | ||
| BAR 09 | LD$Q$Q | TC$d K = d K + d V_{\gamma} \textrm{ } S_{i}^{\top}$d K=d K+d V γS i⊺ | $d Q = s \gamma \textrm{ } d Q$d Q=s γ d Q $d \gamma \textrm{ } = \textrm{ } \sum \textrm{ } Q \textrm{ } \bigodot \textrm{ } d Q$d γ=∑Q⊙d Q | |
| BAR 10 | LD$S$S | TC$d Q = d Q + d P K$d Q=d Q+d P K | ||
| BAR 11 | ST$d Q$d Q | $d \gamma \textrm{ } = \textrm{ } d \gamma \textrm{ } + \textrm{ } \sum_{i} \textrm{ } d A_{\beta} \textrm{ } \bigodot \textrm{ } A \textrm{ } \beta$d γ=d γ+∑id A β⊙A β $d \gamma \textrm{ } = \textrm{ } d \gamma \textrm{ } - \textrm{ } \sum_{j} \textrm{ } d A_{\beta} \textrm{ } \bigodot \textrm{ } A \textrm{ } \beta$d γ=d γ−∑jd A β⊙A β $d \beta \textrm{ } = \textrm{ } \sum_{j} \textrm{ } d A_{\beta} \textrm{ } \bigodot \textrm{ } A$d β=∑jd A β⊙A $d A = d A_{\beta} \textrm{ } \beta$d A=d A ββ | ||
| BAR 12 | TC$d K = d K + d P^{\top} \textrm{ } Q$d K=d K+d P⊺Q | |||
| BAR 13 | TC$d A \textrm{ } = \textrm{ } - A^{\top} \textrm{ } d A \textrm{ } A^{\top}$d A=−A⊺d A A⊺ TC$A_{T} \textrm{ } = \textrm{ } K K^{\top}$A T=K K⊺ | $d O_{\gamma} = s \gamma \textrm{ } d O$d O γ=s γ d O | ||
| BAR 14 | LD$d O$d O LD$A$A | $d \beta \textrm{ } = \textrm{ } d \beta \textrm{ } + \textrm{ } \sum_{i} \textrm{ } d A \textrm{ } \bigodot \textrm{ } A_{T}$d β=d β+∑id A⊙A T $d A_{T} \textrm{ } = \textrm{ } \beta \textrm{ } d A$d A T=β d A $d A_{S} \textrm{ } = \textrm{ } d A_{T} \textrm{ } + \textrm{ } d A_{T}^{\top}$d A S=d A T+d A T⊺ | ||
| BAR 15 | TC$d K = d K + d A_{S} \textrm{ } K$d K=d K+d A SK |
Benchmark#
We benchmarked FlashQLA against the FLA Triton and FlashInfer baseline (FLA 0.5.0, Triton 3.5.1, FlashInfer 0.6.9, TileLang 0.1.8) on the head configurations used by the Qwen3.5 / Qwen3.6 family — $h_{v} \in 64 , 48 , 32 , 24 , 16 , 8$h v∈64,48,32,24,16,8, corresponding to TP1 through TP8.
Specifically, the forward (FWD) benchmarks measure single-kernel latency for different models and TP settings under varying batch lengths, while the backward (BWD) benchmarks examine the relationship between total token count within a batch and latency during a single update step.
Selected H200 single-layer forward results:
| Model / TP | Seqlen | $h_{q k}$h q k | $h_{v}$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× |
The speedup grows with TP degree because FlashQLA improves SM utilization via intra-card AutoCP in the exact regimes — TP sharding and small head number — where the baseline leaves SMs idle.
Usage#
FlashQLA exposes both a high-level API matching FLA’s signature and low-level fwd/bwd entry points:
import torchfrom qla import chunk_gated_delta_ruleo, 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, # optional, [B, H_v, K, V] output_final_state=True, cu_seqlens=cu_seqlens, # optional, varlen support)
Requirements: SM90, CUDA 12.8+, PyTorch 2.8+. Install:
git clone https://github.com/QwenLM/FlashQLA.git
cd FlashQLA && pip install -v .
Acknowledgments#
FlashQLA is inspired by Flash Linear Attention, FlashInfer and TileLang. We thank these communities for the reference implementations.
Citation#
If FlashQLA is useful for your research, please cite:
@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}}}
来源:Qwen:Blog Retrieval(API) · qwen.ai