@PyTorch: PyTorch 成员 Meta 刚刚开源了一个 GPU 内核,使注意力在 NVIDIA Blackwell 上加速 2.3 倍。TLX Block Atte…

X AI KOLs Following 工具

摘要

Meta 开源了 TLX Block Attention,这是一个 warp 特化的 Triton 内核,在 NVIDIA Blackwell GPU 上为块对角自注意力实现了 2.3 倍的加速,与旋转嵌入融合时加速可达 3.5 倍。

PyTorch 成员 Meta 刚刚开源了一个 GPU 内核,使注意力在 NVIDIA Blackwell 上加速 2.3 倍。TLX Block Attention 是一个 warp 特化的 Triton 内核,专为块对角自注意力而构建——这是一种广泛应用于推荐和特征交互模型的模式。 通过利用注意力结构的编译时知识,Flash Attention 算法的整个阶段被消除:没有多 tile 循环,没有校正因子,没有辅助张量。 结果:内核加速 2.3 倍,与旋转嵌入融合时加速 3.5 倍,生产层 MFU 提升 30.6%。 了解更多:https://bit.ly/4fOoh9E 代码:https://bit.ly/4e6SPlK
查看原文
查看缓存全文

缓存时间: 2026/05/26 16:56

PyTorch 成员 Meta 刚刚开源了一个 GPU 内核,使注意力在 NVIDIA Blackwell 上提速 2.3 倍。TLX Block Attention 是一个专为块对角自注意力(推荐系统和特征交互模型中广泛使用的模式)构建的 warp 专用 Triton 内核。

通过利用编译时已知的注意力结构,Flash Attention 算法的整个阶段都被消除了:无需多 tile 循环、无需校正因子、无需辅助张量。

结果:内核加速 2.3 倍,与旋转嵌入融合时加速 3.5 倍,在生产层上 MFU 提升 30.6%。

在此了解更多信息 https://bit.ly/4fOoh9E 代码:https://bit.ly/4e6SPlK


面向固定块稀疏自注意力的 Warp 专用 Blackwell 内核 – PyTorch

来源:https://pytorch.org/blog/tlx-block-attention-a-warp-specialized-blackwell-kernel-for-fixed-block-sparse-self-attention/ 代码地址:https://github.com/facebookresearch/ads_model_kernel_library

在这篇文章中,我们介绍了 TLX Block Attention 的设计——这是一个面向 NVIDIA Blackwell GPU 的 Triton 内核,利用编译时已知的块对角注意力模式,消除了通用注意力实现中存在的整类算法开销。在 NVIDIA B200 GPU 上,该内核的前向速度比 Flash Attention v2 快约 1.85 倍,反向速度快约 2.50 倍;当旋转嵌入融合到注意力后处理中时,组合注意力及旋转反向传播速度提升约 3.5 倍。

这项工作基于 TLX(Triton 语言扩展)——一组对 Triton 编译器的低级扩展,在 NVIDIA Blackwell GPU 上暴露了对 warp 专用化、异步张量核心操作以及内存层次管理的硬件原生控制。TLX 弥合了 Triton 的高层 Python 生产力与通常需要原始 CUDA 或 CUTLASS 的细粒度硬件控制之间的差距。更多关于 TLX 的信息,请参阅 triton-ext 仓库(https://l.facebook.com/l.php?u=https%3A%2F%2Fgithub.com%2Ftriton-lang%2Ftriton-ext&h=AUAOq8QYbTX6IusNi8rHABdpQH9B8xew6KYsjPTIj7t7ziwHNEIGjtsBFtKp-gtiLGDsDq92cpWTOX1ErSS05AY_GKdORcPejIQxMNN3_VT564XkpaSCnYMv-54v3iXi3URUVIBfceAS0uRmvw)

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

1. 引言

自注意力是一种机制,让模型衡量序列中每个元素与其他所有元素的相关程度——本质上是问“输入的哪些部分应该告知我对其他部分的理解?”这是 Transformer 架构的核心构建块,使得这些模型能够捕捉数据中丰富的、依赖于上下文的关系。一个好的直觉可能是:一个人过去的决策如何影响现在和未来的决策?

块对角自注意力——将序列划分为固定大小的组,每组只关注自身内部——是推荐系统和特征交互模型中广泛使用的模式(BlockBERT, Qiu et al., EMNLP 2020 (https://arxiv.org/abs/1911.02972))[1]。在我们的广告排序栈中,生产工作负载通常运行批大小为 1152、序列长度最多约 4000 个 token、头维度为 64 或 128,并且随着序列长度增加,注意力结构中约 70% 的稀疏性。随着这些模型变得更深更宽,注意力成本成为主要瓶颈。

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

这项工作的核心论点:当你在编译时知道注意力模式时,你可以构建更快的东西。 我们利用每个 Q tile 恰好对应一个 K/V tile 的固定约束,将此知识传播到整个算法中,从而将多次迭代的累加器简化为单个 GEMM,消除校正阶段,并移除辅助内核启动。

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

2. 为什么是块注意力?

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

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

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

``

Flash Attention 内循环(标准)

for k_tile in K_tiles: S = Q @ k_tile.T # 部分分数 m_new = max(m_old, rowmax(S)) alpha = exp(m_old - m_new) # 校正因子 O = alpha * O + exp(S - m_new) @ v_tile l = alpha * l + rowsum(exp(S - m_new)) O = O / l # 最终归一化

将 L = m + log(l) 存储到 HBM 以供反向使用

``

这对于任意序列是正确的且优雅的。但对于块大小为 64 token 的块对角注意力,整个 Q-tile-over-K-tiles 循环减少为单次迭代。每个 Q tile 及其对应的 K/V tile 是同一个 tile。这一单一约束级联地影响整个算法:

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

列表 2:use_acc=False 向硬件指示无需跨 tile 累加,从而启用 TMEM 重用。

``

来自内核:use_acc=False 指示无需累加

tlx.async_dot( q_tile[buff_idx], k_tile_T, TMEMqk[tmem_idx], use_acc=False, # 新结果——无需累加 mBarriers=[qk_SMEM_free[buff_idx], qk_TMEM_full[tmem_idx]], ) ``

2.2 与标准 Flash Attention 的比较

下表总结了算法差异:

方面标准 Flash AttentionTLX Block Attention
每个 Q tile 的 K tile 数多个(整个序列)恰好 1 个(同一块)
分数矩阵多个 tile 累加单个 [64, 64] — 完整
Logsumexp L 张量存储到 HBM 供反向使用不需要
运行中的最大值/求和跨 tile 维护一次性计算,寄存器内使用
校正因子 α每次迭代需要不需要(省略)
输出累加增量式带缩放单个 P·V GEMM
use_acc 模式True(跨 tile 累加)False(新结果)
Di 预处理单独内核启动内联计算

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

这些不是微优化——它们代表了整个算法阶段的消除。反向传播尤其受益:缺少存储的 L 张量消除了每个 batch × head × sequence 的一次 HBM 往返,而内联 Di 计算则消除了内核启动及其相关的驱动开销和内存带宽。

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

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

3.1 TLX

我们选择 Triton 作为编写框架,因为它提供了 Python 原生、面向 tile 的编程模型,自然地映射到下面描述的 warp 专用化流水线结构——同时避免了原始 CUDA 或 CUTLASS 的样板代码,并且在不同编译器版本间保持可移植性。Triton 的 TLX(Triton 语言扩展)进一步暴露了 Blackwell 特定的原语,如 async_dot、local_trans 和显式的 TMEM/SMEM 屏障管理,其抽象级别在硬件控制与开发者生产力之间取得了平衡。根据我们的经验,TLX 提供的性能可媲美(且经常超越)更低级的替代方案,同时由于其 Python 原生的简洁性,迭代速度显著更快。

具体来说,此内核依赖于几个超出基本 Triton 的 TLX 原语:tlx.async_dot 用于发出 warp 专用化的 tcgen05 MMA 操作并带有显式累加器控制;tlx.async_descriptor_load 用于 TMA 驱动的 SMEM 填充;tlx.local_trans 用于 TMEM 到寄存器的传输;以及 mBarrier 同步模型,用于协调跨 warp 组的生成者-消费者流水线。这些扩展可在 triton-ext 仓库(https://l.facebook.com/l.php?u=https%3A%2F%2Fgithub.com%2Ftriton-lang%2Ftriton-ext&h=AUD8yMu0haMw6d272mfnKLyPQGAzLzIYglBbh1LZ4klMCKgd0std74g4lH8r6sst8DVeP6KNWx3T3ZRv0OyeTM7u-Jr2t3Av244-fCsvpiVCJOrDVpPuf3PpZTCJNEkBtTTT-2P7XghPEt25DA) 中找到。

3.2 Warp 专用化

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

阶段Warp 数寄存器数硬件单元角色
加载148TMA 引擎为 Q、K、V 执行 async_descriptor_load
QK MMA148tcgen05 张量核心async_dot(Q, KT) → TMEMqk
Softmax4120CUDA 核心 + SFU掩码 / 缩放 / exp2 / 归一化 → P 到 SMEM
PV MMA148tcgen05 张量核心async_dot(P, V) → TMEMpv
后处理8200CUDA 核心 + L2 + TMA 引擎TMEM → 寄存器 → BF16 → SMEM → TMA 存储
总计15——每个 CTA 480 线程

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

`` 图 1 — 前向流水线 warp 时间线(概念性,一次迭代):

时间 → 加载 [─ TMA Q,K ─][─ TMA V ─] QK MMA [── async_dot Q·KT ──] Softmax [── exp2/normalize → P ──] PV MMA [── async_dot P·V ──] 后处理 [── local_load → BF16 → store ──] ``

每个阶段的输出发出一个屏障信号,解锁下一阶段,从而在硬件单元之间创建生成者-消费者流水线。当后处理 warp 将 tile i 写入全局内存时,MMA warp 正在计算 tile i+1,而加载 warp 正在通过 TMA 获取 tile i+2——同时有三个 tile 在飞行中。

3.3 Roofline 上下文

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

3.4 缓冲区管理

为了保持硬件单元持续忙碌,该内核使用三重缓冲 SMEM(3 个槽位)和双重缓冲 TMEM(2 个槽位),消耗约 169 KB 的 256 KB SMEM 预算。通过三个 SMEM 槽位,加载 warp 可以预取 tile i+2,而 MMA warp 处理 tile i+1,后处理 warp 清空 tile i。反向内核降级为双重缓冲 SMEM(约 162 KB),以便在相同的 256 KB 预算内容纳额外的梯度 tile。

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

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

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

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

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

阶段Warp 数寄存器数硬件单元角色
加载148TMA 引擎加载 Q、K、V、dO(+ sin/cos 用于旋转)
QK MMA148tcgen05 张量核心重新计算 S = Q · KT
Softmax/P4120CUDA 核心 + SFU重新计算 P = softmax(S)
dV MMA148tcgen05 张量核心dV = PT · dO
dP/dS4120TC + CUDA 核心dP = dO · VT, Δi, dS
dQ/dK MMA148tcgen05 张量核心dQ = dS · K, dK = dST · Q
后处理8200CUDA 核心 + L2 + TMA 引擎存储 dQ、dK、dV(+ 融合旋转)
总计20——每个 CTA 640 线程

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

反向传播本质上前向更复杂。它需要 20 个 warp(每个 CTA 640 线程)来平衡密集的计算需求。最值得注意的是,它完全占满了 SM 上的 256 KB 张量内存。五个不同的 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 tile 数量的前缀和。
  2. BLOCK_PER_PROGRAM:分配给每个 SM 的平衡 tile 范围——使用封闭形式的 divmod 算术而非累积和来计算。

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

在内核内部,每个 SM 必须确定某个全局 tile 索引属于哪个序列(批索引)。这使用一个无分支的二分查找,恰好执行 3

相似文章