FlashQLA:面向 CP / 反向友好的融合线性注意力内核(用于 GDN)

Qwen 已发布 FlashQLA,这是一款旨在优化门控增量网络(GDN)层的高性能线性注意力内核库。FlashQLA 基于 TileLang 构建,在 NVIDIA Hopper GPU 上相较于 FLA Triton 内核实现了 2-3 倍的前向加速和 2 倍的反向加速,特别有利于预训练和边缘侧的自主推理。

优化 GDN 分块预填充

FlashQLA 解决了原始 FLA 实现中 Gated Delta Network(GDN)分块预填充过程的两个主要效率瓶颈:

  1. 内存受限内核:标准流程会反复读取和写入中间变量 ($W, U, S$) 到高带宽内存(HBM),导致大量开销。
  2. 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_dvbwd_dhubwd_dqkwgbwd_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 获取。

Sources