博客

TLX Block Attention:一种用于定长块稀疏自注意力的 Warp-Specialized Blackwell 内核

代码获取地址:https://github.com/facebookresearch/ads_model_kernel_library 

在这篇文章中,我们介绍了 TLX Block Attention 的设计——这是一个针对 NVIDIA Blackwell GPU 的 Triton 内核。它利用编译时已知的分块对角(block-diagonal)注意力模式,消除了通用注意力实现中存在的各类算法开销。在 NVIDIA B200 GPU 上,该内核的前向传播速度比 Flash Attention v2 快约 1.85 倍,反向传播速度快约 2.50 倍;当将旋转嵌入(rotary embeddings)融合进注意力尾声(epilogue)时,合并后的注意力与旋转反向传播任务可实现约 3.5 倍的加速。

这项工作建立在 TLX(Triton Language Extensions)之上——这是一套对 Triton 编译器的底层扩展,旨在 NVIDIA Blackwell GPU 上实现对 warp 专用化(warp specialization)、异步 Tensor Core 操作和内存层级管理的硬件级控制。TLX 弥合了 Triton 的高级 Python 生产力与传统上需要原生 CUDA 或 CUTLASS 才能实现的精细化硬件控制之间的鸿沟。有关 TLX 的更多信息,请参阅 triton-ext 仓库

───────────────────────────────────────

1. 介绍

自注意力(Self-attention)是一种让模型权衡序列中每个元素对其他元素相关性的机制——本质上是在问“输入中的哪些部分应该影响我对其他部分的理解?”它是 Transformer 架构的核心构建模块,使这些模型能够捕捉数据中丰富且依赖上下文的关系。一个直观的理解是:一个人过去的决定如何影响现在和未来的决定?

块对角自注意力(Block-diagonal self-attention)——即序列被划分成固定大小的组,且仅在组内进行注意力计算——是推荐系统和特征交互模型中广泛使用的模式(BlockBERT, Qiu et al., EMNLP 2020)[1]。在我们的广告排名堆栈中,生产工作负载通常运行 1152 的批大小,序列长度高达约 4k token,头维度为 64 或 128,并且随着序列长度的增加,注意力结构中的稀疏度约为 70%。随着这些模型变得越来越深、越来越宽,注意力的开销成为了主要的瓶颈。

如今这些工作负载运行在带有分块掩码或滑动窗口的 Flash Attention v2 等通用内核上。FlexAttention (FA4) [7] 虽然支持块稀疏模式,但其最小分块大小为 256,这与这些模型所需的 64-token 分块不兼容。在 64-token 的块大小下,带有分块掩码的 Flash Attention v2 仍然是目前可用的最强基准,但它仍有巨大的性能提升空间。Flash Attention 的分块迭代、在线 softmax 校正、logsumexp 记录以及辅助内核启动对于任意长度的因果注意力至关重要——但当模式为分块对角且在编译时已知时,这些纯属开销。

本工作的核心论点:当你在编译时已知注意力模式,你就可以构建出快得多的东西。 我们利用了每个 Q 块精确地对一个 K/V 块进行注意力的固定约束,将这一知识贯穿整个算法,从而将多次迭代的累加器折叠为单一 GEMM,消除校正阶段,并移除辅助内核启动。

───────────────────────────────────────

2. 为什么选择块注意力?

2.1 固定块约束及其带来的简化级联

标准 Flash Attention [2] 通过在多个 K/V 块上迭代 Q 块来处理任意长度的序列,维护运行统计信息(行最大值和 log-sum-exp),并在每一步应用校正因子以保持数值稳定性。

列表 1:标准 Flash Attention 内循环,显示多块迭代和在线 softmax 校正。

# Flash Attention inner loop (standard)
for k_tile in K_tiles:
    S = Q @ k_tile.T                   # partial scores
    m_new = max(m_old, rowmax(S))
    alpha = exp(m_old - m_new)         # correction factor
    O = alpha * O + exp(S - m_new) @ v_tile
    l = alpha * l + rowsum(exp(S - m_new))
O = O / l                              # final normalization
# Store L = m + log(l) to HBM for backward

对于任意序列,这既正确又优雅。但对于具有固定 64-token 块大小的分块对角注意力,整个“Q 块在 K 块上迭代”的循环减少为单次迭代。每个 Q 块及其对应的 K/V 块都是同一个块。这一约束在算法中产生了级联效应:

  1. 无多块迭代。 得分矩阵 S = Q · Kᵀ ∈ ℝ^{64×64} 在一次 GEMM 后即完成计算。没有循环需要维护跨块状态。
  2. 无在线 softmax 校正。 由于只有一个块,在 S 上计算的行最大值和总和立即全局正确。校正因子 α = exp(m_old − m_new) 恒等于 1,可直接舍弃。
  3. 无 logsumexp (L) 存储。 Flash Attention 将每行的 log-sum-exp L 存储到 HBM,以便反向传播重新计算 softmax。在单块情况下,反向传播可以直接从 Q、K、V 重新计算 P = softmax(S),无需任何辅助张量——消除了每对前向/反向传播中一次完整的 HBM 读写。
  4. 无 Di 预处理内核。 标准 Flash Attention 反向传播在主反向阶段之前启动单独的内核来计算 Di = rowsum(dO ⊙ O)。在 TLX Block Attention 中,Di 在 dP/dS 反向阶段内联计算,消除了内核启动及其相关的内存流量。
  5. 无带重缩放的输出累加。 在单块情况下,输出 O = P · V 是单次 GEMM 的直接结果,而不是多个缩放后的部分结果的累加。这使得所有 async_dot 调用都可以使用 use_acc=False,告知 Tensor Core 硬件无需在调用间保留 TMEM 累加器,从而可以自由重用。

列表 2:use_acc=False 向硬件发出信号,表示不需要跨块累加,从而实现 TMEM 重用。

# From the kernel: use_acc=False signals no accumulation needed
tlx.async_dot(
    q_tile[buff_idx],
    k_tile_T,
    TMEMqk[tmem_idx],
    use_acc=False,           # Fresh result — no accumulation
    mBarriers=[qk_SMEM_free[buff_idx], qk_TMEM_full[tmem_idx]],
)

2.2 与标准 Flash Attention 的对比

下表总结了算法上的差异:

维度 标准 Flash Attention TLX Block Attention
每个 Q 块的 K 块数量 多个(全序列) 恰好 1 个(同一块)
得分矩阵 多块累加 单个 [64, 64] — 完成
Logsumexp L 张量 存入 HBM 用于反向传播 不需要
运行中的最大值/和 跨块维护 计算一次,在寄存器中消费
校正因子 α 每次迭代都需要 不需要(已剔除)
输出累加 带有重缩放的增量式 单次 P·V GEMM
use_acc 模式 True(跨块累加) False(即时结果)
Di 预处理 单独的内核启动 内联计算

表 1:标准 Flash Attention 与 TLX Block Attention 之间的算法差异。

这些不仅仅是微优化——它们代表了整个算法阶段的消除。反向传播受益尤为显著:无需存储 L 张量消除了每批次 × 头数 × 序列的 HBM 往返,而内联 Di 计算则消除了一个带有相关驱动开销和内存带宽的内核启动。

───────────────────────────────────────

3. 内核架构:Warp 专用化流水线

3.1 TLX

我们选择 Triton 作为创作框架,因为它提供了原生 Python、面向分块的编程模型,能够自然地映射到下述 Warp 专用化流水线结构,同时避免了原生 CUDA 或 CUTLASS 的样板代码,并保持了跨编译器演进的可移植性。Triton 的 TLX (Triton Language Extensions) 进一步以一种平衡硬件控制与开发者生产力的抽象级别,暴露了 Blackwell 特有的原语(如 async_dot、local_trans 以及显式 TMEM/SMEM 屏障管理)。根据我们的经验,TLX 的性能与更底层的替代方案相当(且通常更高),同时得益于其 Python 原生的简洁性,开发迭代速度显著加快。

具体而言,该内核依赖于几个超出基础 Triton 的 TLX 原语:用于发出带有显式累加器控制的 Warp 专用化 tcgen05 MMA 操作的 tlx.async_dot;用于 TMA 驱动的 SMEM 填充的 tlx.async_descriptor_load;用于 TMEM 到寄存器传输的 tlx.local_trans;以及协调 Warp 组间生产者-消费者流水线的 mBarrier 同步模型。这些扩展可在 triton-ext 仓库 中获取。

3.2 Warp 专用化

TLX Block Attention 使用了 Warp 专用化[8] —— 同一个 CTA 中的不同 Warp 被永久分配给不同的硬件单元,并在内核的整个生命周期中执行不同的代码路径。这与传统的 CUDA 模型形成了鲜明对比,在传统模型中,所有 Warp 执行相同的代码,仅通过条件判断产生分叉。

阶段 Warp 寄存器 硬件单元 角色
加载 1 48 TMA 引擎 async_descriptor_load 用于 Q, K, V
QK MMA 1 48 tcgen05 Tensor Core async_dot(Q, Kᵀ) → TMEMqk
Softmax 4 120 CUDA 核心 + SFU 掩码 / 缩放 / exp2 / 归一化 → P 写入 SMEM
PV MMA 1 48 tcgen05 Tensor Core async_dot(P, V) → TMEMpv
尾声 8 200 CUDA 核心 + L2 + TMA 引擎 TMEM → 寄存器 → BF16 → SMEM → TMA 存储
总计 15 每个 CTA 480 个线程

表 2:前向流水线阶段配置。寄存器分配是有意非对称的——硬件加速阶段分配最少寄存器;CUDA 核心阶段分配最多。

Fig. 1 — Forward pipeline warp timeline (conceptual, one iteration):

Time →
Load     [─ TMA Q,K ─][─ TMA V ─]
QK MMA         [── async_dot Q·Kᵀ ──]
Softmax                  [── exp2/normalize → P ──]
PV MMA                            [── async_dot P·V ──]
Epilogue                                   [── local_load → BF16 → store ──]

每个阶段的输出都会触发一个屏障,解除下一个阶段的阻塞,从而在硬件单元之间创建生产者-消费者流水线。当尾声(Epilogue)Warp 将分块 i 写入全局内存时,MMA Warp 正在计算分块 i+1,加载(Load)Warp 正通过 TMA 获取分块 i+2 —— 同时有三个块在流水线中。

3.3 屋顶线(Roofline)背景

在 BLOCK_D=64,HEAD_DIM=128 时,算术强度约为 33 FLOP/byte —— 远低于 B200 的拐点 ~281 FLOP/byte [4]。该内核的设计意图是受内存带宽限制。这就是为什么通过 TMA 隐藏延迟以及最小化不必要的内存流量(消除 L 张量、融合旋转操作)是主要的优化杠杆。

3.4 缓冲管理

为了使硬件单元持续繁忙,内核使用了三缓冲 SMEM(3 个槽)和双缓冲 TMEM(2 个槽),消耗了 256 KB SMEM 总预算中的约 169 KB。通过三个 SMEM 槽,加载 Warp 可以在 MMA Warp 处理分块 i+1、尾声 Warp 排空分块 i 的同时预取分块 i+2。反向内核降级为双缓冲 SMEM(约 162 KB),以便在相同的 256 KB 预算内容纳更多的梯度块。

───────────────────────────────────────

4. 反向传播:无需 Logsumexp 张量的梯度计算

在标准 Flash Attention 中,反向传播要求前向传播将 logsumexp 张量 (L) 保存到高带宽内存(HBM)。此张量对于在反向传播期间重建注意力概率 (P) 是必要的。此外,标准注意力需要一个单独的预处理内核来计算 Δᵢ(dO ⊙ O 的行和)。

由于分块对角注意力在单个分块中计算整个 64×64 得分矩阵,我们可以完全绕过这两个要求。反向内核既不读取任何 logsumexp 张量,也不需要单独的预处理步骤。相反,它会在内联中完全重新计算 S = Q · Kᵀ 和 P = softmax(S) —— 当分块可以在单次通过中完成时,这是一项低成本操作。

这种简化级联使我们能够构建一个完全融合的、7 阶段的 Warp 专用化反向流水线:

阶段 Warp 寄存器 硬件单元 角色
加载 1 48 TMA 引擎 加载 Q, K, V, dO(+ 旋转 sin/cos)
QK MMA 1 48 tcgen05 Tensor Core 重新计算 S = Q · Kᵀ
Softmax/P 4 120 CUDA 核心 + SFU 重新计算 P = softmax(S)
dV MMA 1 48 tcgen05 Tensor Core dV = Pᵀ · dO
dP/dS 4 120 TC + CUDA 核心 dP = dO · Vᵀ, Δᵢ, dS
dQ/dK MMA 1 48 tcgen05 Tensor Core dQ = dS · K, dK = dSᵀ · Q
尾声 8 200 CUDA 核心 + L2 + TMA 引擎 存储 dQ, dK, dV(+ 融合旋转操作)
总计 20 每个 CTA 640 个线程

表 4:7 阶段反向流水线配置。

反向传播本质上比前向传播更复杂。它需要 20 个 Warp(每个 CTA 640 个线程)来平衡高强度的计算需求。最值得注意的是,它完全饱和了 SM 上的 256 KB 张量内存(TMEM)。五个不同的 TMEM 缓冲区 —— TMEMqk, TMEMdv, TMEMdp, TMEMdq, 和 TMEMdk —— 总计达到 100% 的 TMEM 利用率。为了适应这一点,反向内核将前向传播中的三缓冲 SMEM 降级为双缓冲 SMEM(约 162 KB / 256 KB,占 63%),同时保持双缓冲 TMEM。

───────────────────────────────────────

5. 变长序列调度

现实世界中的推荐和特征交互模型并不处理整齐统一的序列长度。相反,流量主要由打包在单个扁平缓冲区中的锯齿状、变长序列组成。天真地为每个序列映射一个 CTA 会导致 SM 空闲——当短序列先结束而其他序列仍在处理长序列时——这会造成严重的工作负载失衡。

为了最大化 SM 占用率,内核启动了 min(NUM_SMS, total_blocks) 个持久化程序 —— 每个 SM 恰好有一个持久化线程块。工作负载通过两个预计算数组进行平衡:

  1. BLOCK_PER_BATCH:每个序列的 64-token 分块数量的前缀和。
  2. BLOCK_PER_PROGRAM:分配给每个 SM 的平衡分块范围 —— 使用闭式 divmod 算术而非累积和计算得出。

为了消除 GPU 同步开销,当 CPU 端偏移张量可用(cpu_offsets)时,所有标量调度算术(分块计数、divmod、前缀和)都在内核启动前在 CPU 上计算完成 —— 实现零 GPU 同步点。

在内核内部,每个 SM 必须确定给定的全局分块索引属于哪个序列(批次索引)。这使用了一种无分支的二分查找,该查找恰好执行 32 次迭代(对于任何合理的批大小都足够),且零线程同步。

───────────────────────────────────────

6. 融合旋转反向传播:以更高速度实现更高精度

对于自注意力层,自注意力之前是投影 + 正弦波 [6]。在反向传播中,这变成了注意力反向传播 -> 正弦波反向传播,通常需要两个不同的内核启动。

6.1 基准:双内核反向传播

传统的反向传播需要两个单独的内核启动:

  1. 注意力反向传播内核 —— 通过 Tensor Core 以 FP32 累加 dQ、dK、dV,然后在存储到全局内存时截断为 BF16。
  2. 旋转反向传播内核 —— 从全局内存重新加载 BF16 梯度,应用旋转共轭 R(−θ),并存储最终的 BF16 结果。

这种分离有三个成本:

问题 影响
精度损失 FP32 梯度在旋转变换之前即被截断为 BF16 —— 然后在最终存储时再次截断。两个量化点,每个引入约 0.4% 的相对误差(BF16 仅有 7 位尾数)。下游投影 GEMM 会放大累积的误差。
内存带宽浪费 dQ, dK, dV 被写入后立即被重新读取 —— 在 [total_seq_len, 1152] 张量(head_dim=128, 3 个 KV 头)上产生完整的往返。随着序列长度达到数百万,这些流量相当可观。
内核启动开销 明明一次即可完成,却进行了两次单独调度。

6.2 融合方案

注意力反向传播内核已经分配了一个专门的 Warp 组用于梯度存储尾声。我们利用这一点,将旋转共轭注入该尾声中,此时梯度仍处于 FP32 寄存器中:

  1. Tensor Core 以 FP32(TMEM)存储 dQ, dK, dV。
  2. 将 FP32 值加载到寄存器中。
  3. 以完整 FP32 精度应用 R(−θ) —— 轻量级的 sin/cos 加载 + 逐元素乘法。
  4. 转换为 BF16 并发出单次全局存储。

各步骤对比:

维度 基准(分离) 融合内核
注意力反向传播计算 FP32 FP32
中间存储 BF16 → 全局内存 FP32 寄存器
旋转 sin/cos 操作 BF16 FP32
BF16 量化点 2 1(仅最终存储)
全局内存往返 2 0
内核启动 2 1

在反向尾声中融合旋转共轭。交错操作在 FP32 中对配对的 [cos, sin] 分量应用 R(−θ)。

# Apply rotary conjugate to dV (neg_sin handles the conjugate)
dv0, dv1 = dvLocal.reshape(BLOCK_D, HALF_DIM, 2).split()
dvLocal = tl.interleave(
    dv0 * cos_local - dv1 * neg_sin,
    dv1 * cos_local + dv0 * neg_sin,
)

───────────────────────────────────────

7. 性能结果

所有基准测试均在 NVIDIA B200 GPU(x86 CPU)上以 BF16 精度进行。主要配置使用 B=1152 序列,HEAD_DIM=128,H=4 个头,max_seq_len=2000,且稀疏度=0.7 —— 离散均匀分布(代表生产流量分布)。

7.1 内核级加速

传播 带块注意力的 Flash Attention v2 (ms) TLX Block Attention (ms) 加速比
前向 1.81 0.98 1.85×
反向 5.89 2.36 2.50×
总计 7.70 3.33 2.31×

表 5:内核级性能对比(B=1152, D=128, H=4, BF16, B200, max_seq_len=2000, 稀疏度=0.7)。

反向传播的加速(2.50×)大于前向加速(1.85×),主要是因为反向传播受益于两项独立的简化:(1)消除了 logsumexp 存储和 Di 预处理,以及(2)内联 P 重计算,避免了标准 Flash Attention 反向传播所必需的 L 张量 HBM 往返。

7.2 跨工作负载的扩展性

表 6:跨序列长度和稀疏比率的扩展性能。无论分布形状如何,加速比都是一致的(batch=1152,对于 >7000 batch=768)。相比 Flash Attention v2 (jfa) 的内核加速。

7.3 融合旋转反向传播

将旋转反向传播融合进注意力尾声的影响尤为显著:

配置 时间 (ms)
注意力反向传播(独立) 1.556
旋转反向传播(独立) 4.880
未融合总计 6.436
融合注意力_旋转反向传播 1.819
加速比 3.54×

表 7:融合与非融合旋转反向传播的时间明细。独立的旋转内核主导了未融合的总时间。seq_len=1735537, heads=3, head_dim=128, batch=1152。

独立的旋转反向传播比注意力反向传播本身的成本高出 3 倍以上 —— 它纯粹受限于内存带宽,读写 [M, D] 张量而没有任何有意义的计算。将其融合进注意力尾声将此带宽成本分摊到现有的 TMEM → 寄存器流水线上,使组合操作从 6.436 ms 缩减至 1.819 ms。

端到端来看,将此内核集成到自注意力层中,在这些层上实现了 +30.6% 的模型 FLOPs 利用率(MFU)提升

7.4 数值精度

将旋转反向传播融合进 FP32 尾声还带来了可衡量的精度提升。与高精度 PyTorch 参考实现相比,TLX Block Attention 将查询梯度 (dQ) 中的最大梯度误差降低了 2 倍以上:

指标 Flash Attention v2 TLX Block Attention 更准确
最大 dQ 差异 0.2559 0.1201 TLX
最大 dK 差异 0.1689 0.1689 平手
最大 dV 差异 0.0112 0.0112 平手
平均 dQ 差异 0.000309 0.000220 TLX

表 8:相对于 PyTorch 参考实现的梯度数值精度。由于单量化点融合旋转路径,TLX Block Attention 将最大 dQ 误差降低了 53%。

dQ 受益最大,因为查询梯度 (dQ = dS · K) 通过了带有 1 个量化点而非 2 个量化点的融合旋转共轭。dK 也通过了旋转共轭(RoPE 同时旋转 Q 和 K),但其最大绝对误差恰好由 MMA 累加本身而非旋转内存往返所主导,因此消除中间 BF16 转换带来的逐元素改进在最大值上并未体现。

───────────────────────────────────────

8. 适用性

如果你的模型使用分块对角注意力 —— 即每个 token 仅对固定局部组内的其他 token 进行注意力计算 —— 那么此内核非常适用。

  • 在 NVIDIA Blackwell GPU 上进行训练。 该内核使用 tcgen05 MMA 指令、TMEM 分配和 Blackwell 时代的 TMA 描述符 —— 这些在 Ampere 或 Hopper 上均不存在。async_dot / local_trans / tlx API 专门针对 Blackwell 架构(sm_100+)。
  • HEAD_DIM ∈ {64, 128}。 这是支持的头维度;其他数值需要重新编译并可能需要重新计算 SMEM/TMEM 预算。

───────────────────────────────────────

9. 结论

TLX Block Attention 展示了单一架构约束所带来的复合威力。通过识别出广泛的特征交互和序列模型仅需要严格的分块对角注意力,一系列简化成为可能。

消除块间注意力意味着没有多块累加。没有多块累加意味着没有在线 softmax 校正因子。没有在线 softmax 校正意味着 logsumexp 张量可以在反向传播中完全丢弃。没有单独的 logsumexp 张量腾出了足够的寄存器和内存带宽预算,可以完全将旋转嵌入直接融合进反向传播尾声,这独立地提升了速度和数值精度。

结果是一个完美适配 Blackwell 架构 TMA 和 TMEM 硬件原语的 Warp 专用化内核:前向传播 15 个 Warp,反向传播 20 个,每个 Warp 组被永久分配给与其瓶颈相匹配的硬件单元。该设计在 Flash Attention v2 基础上实现了 2.3 倍的内核级加速,当融合旋转操作时反向传播加速比达到 3.5 倍,并在生产自注意力层上实现了 +30.6% 的 MFU 增益。

该内核已在 github.com/facebookresearch/ads_model_kernel_library 开源 —— 请在您自己的块稀疏注意力工作负载上尝试,并让我们知道您的发现。

───────────────────────────────────────

致谢

作者感谢 Triton [5] 和 PyTorch 团队持续开发 tlx Blackwell 扩展,使此内核成为可能。特别感谢更广泛的 GPU 内核研究社区,他们在 Flash Attention、Warp 专用化流水线和持久化内核调度方面的工作为这些优化奠定了基础。

───────────────────────────────────────

参考文献

  1. Qiu, J., Ma, H., Levy, O., Yih, S. W., Wang, S., & Tang, J. (2020). BlockBERT: Efficient Attention Using Block Structures. EMNLP Findings 2020. https://arxiv.org/abs/1911.02972
  2. Dao, T., Fu, D. Y., Ermon, S., Rudra, A., & Ré, C. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS 2022. https://arxiv.org/abs/2205.14135
  3. Dao, T. (2024). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. ICLR 2024. https://arxiv.org/abs/2307.08691
  4. NVIDIA Corporation. (2024). NVIDIA Blackwell Architecture Technical Brief. https://resources.nvidia.com/en-us-blackwell-architecture
  5. Tillet, P., Kung, H. T., & Cox, D. (2019). Triton: An Intermediate Language and Compiler for Tiled Neural Network Computations. MAPL 2019. https://www.eecs.harvard.edu/~htk/publication/2019-mapl-tillet-kung-cox.pdf
  6. Su, J., Lu, Y., Pan, S., Murtadha, A., Wen, B., & Liu, Y. (2021). RoFormer: Enhanced Transformer with Rotary Position Embedding. https://arxiv.org/abs/2104.09864
  7. He, H. & Guessous, D. (2024). FlexAttention: The Flexibility of PyTorch with the Performance of FlashAttention. PyTorch Blog. https://pytorch.ac.cn/blog/flexattention/
  8. Yu, H., Ren, M., Maher, B., Nay, S., Zhu, G., & Jiang, S. (2024). Enabling Advanced GPU Features in PyTorch – Warp Specialization. PyTorch Blog. https://pytorch.ac.cn/blog/warp-specialization/