FlashQLA:面向 CP / 反向友好的融合线性注意力内核(用于 GDN)
Qwen 已发布 FlashQLA,这是一款旨在优化门控增量网络(GDN)层的高性能线性注意力内核库。FlashQLA 基于 TileLang 构建,在 NVIDIA Hopper GPU 上相较于 FLA Triton 内核实现了 2-3 倍的前向加速和 2 倍的反向加速,特别有利于预训练和边缘侧的自主推理。
优化 GDN 分块预填充
FlashQLA 解决了原始 FLA 实现中 Gated Delta Network(GDN)分块预填充过程的两个主要效率瓶颈:
- 内存受限内核:标准流程会反复读取和写入中间变量 ($W, U, S$) 到高带宽内存(HBM),导致大量开销。
- GPU 利用率低:状态空间模型(SSM)状态的递归特性将同时运行的线程块数量限制为
batch_size * num_heads。在模型规模小、批次小或张量并行(TP)等场景下,这会导致 GPU 流式多处理器(SM)闲置。
为了解决这些相互冲突的问题,FlashQLA 采用了避免全融合内核(在小批量场景下会失效)的方案,而是将前向计算拆分为两个融合内核,并在它们之间插入上下文并行(CP)预处理步骤。
关键技术创新
基于门控的自动卡内上下文并行(AutoCP)
FlashQLA 在张量并行(TP)、长序列和少头数设置下实现了自动卡内 CP 机制,以提升 SM 利用率。它使用数学模型来确定最优并行度 ($L = \lambda \sqrt{N}$),其中 $N$ 为块数,$L$ 为每个 CP 等级的块数。
为进一步降低开销,FlashQLA 利用 GDN 门的指数衰减特性。对于门 $\alpha_i \in (0,1)$ 的头部,先前状态的影响会指数衰减。FlashQLA 采用“热身”过程(通常为 6–8 块)将状态误差降至噪声底以下,从而可以省去昂贵的校正项 $M$ 矩阵计算,直接获得精确的子序列 $S_0$。
TileLang Warp 专用内核
FlashQLA 使用 TileLang 实现 warpgroup 专用内核。该架构在同一 SM 内使用一个生产者 warpgroup 和三个消费者 warpgroup,通过共享内存交换数据,并使用 mbarriers 同步。
- 前向传播:三个消费者 warpgroup 分别计算 $V'、S、O$,采用 ping-pong 结构实现计算与内存流量的重叠。
- CP 预处理:单个融合内核同时处理原始的 $M$ 与 $S$ 计算以及更轻量的滑动窗口热身方法。
- 反向传播:FlashQLA 将
bwd_dv、bwd_dhu、bwd_dqkwg和bwd_wy融合为单个内核。由于片上资源受限,它依赖长计算链来隐藏内存流量,而非多阶段流水线。
性能基准
在 NVIDIA H200 GPU 上相对于 FLA Triton 和 FlashInfer 基线进行的基准测试显示出显著提升,尤其在张量并行(TP)程度提升时更为明显。
| 模型 / TP | 序列长度 | $h_{qk}$ | $h_v$ | FlashQLA | FlashInfer | FLA vs FLA vs FI |
|---|---|---|---|---|---|---|
| 397B/122B TP8 | 1x32768 | 2 | 8 | 0.310ms | 1.653ms | 2.95× |
| 397B/122B TP4 | 1x32768 | 4 | 16 | 0.486ms | 1.654ms | 2.57× |
| 27B TP2 | 1x32768 | 8 | 24 | 0.659ms | 1.616ms | 2.37× |
| 2B/0.8B TP1 | 1x32768 | 16 | 16 | 0.493ms | 1.640ms | 2.60× |
实现与要求
FlashQLA 提供了兼容 FLA 接口签名的高级 API,以及用于前向和反向传播的低层入口点。
系统要求:
- 硬件:NVIDIA SM90(Hopper)
- 软件:CUDA 12.8+,PyTorch 2.8+
代码和基准测试可在 github.com/QwenLM/FlashQLA 获取。