ThriftAttention: 长上下文FP4注意力的选择性混合精度

arXiv cs.LG 论文

摘要

ThriftAttention提出了一种选择性混合精度注意力方法,该方法仅对一小部分查询-键块使用FP16计算,其余使用FP4,从而在长上下文推理中实现接近FP16的质量和FP4的效率。

arXiv:2605.23081v1 公告类型: 新 摘要: 高效的注意力算法对于缓解长上下文工作负载中注意力的二次成本至关重要。先前的工作利用Blackwell GPU上的块缩放量化技术,将注意力计算迁移至4位精度以加速推理。然而,这些技术在长上下文设置中会导致显著的质量下降。我们表明,量化误差的输出影响高度不均匀,并且随着每次查询-键交互的重要性而增加,将功能相关的误差集中在包含最重要令牌的少量注意力块中。我们提出了ThriftAttention,一种低比特注意力变体,以FP4推理效率提供接近FP16的长上下文质量。该方法分两个阶段进行。首先,一种启发式方法快速选择少量重要的查询-键块对进行FP16精度计算。其次,所选块以FP16计算,其余块以FP4计算,两个路径通过在线softmax合并为单个输出。我们在长上下文基准测试和模型家族中证明,仅计算5%的查询-键块为FP16,ThriftAttention平均恢复了FP4到FP16性能差距的89.1%。我们展示了ThriftAttention的优势随着序列长度增长而增加,缓解了在较长上下文中观察到的系统性FP4质量下降。代码可在https://github.com/joesharratt1229/ThriftAttention获取。
查看原文
查看缓存全文

缓存时间: 2026/05/25 09:00

# ThriftAttention:面向长上下文FP4注意力机制的混合精度选择策略
来源:https://arxiv.org/html/2605.23081

###### 摘要

高效的注意力算法对于缓解长上下文任务中注意力的二次成本至关重要。先前的工作利用Blackwell GPU上的块级量化技术,将注意力计算迁移至4位精度以加速推理。然而,这些技术在长上下文场景下会导致显著的质量下降。我们证明,量化误差对输出的影响具有高度非均匀性,且随每个查询-键交互的重要性增加而增大,将功能相关的误差集中在一小部分包含最重要token的注意力块中。我们提出**ThriftAttention**,一种低比特注意力变体,能够以FP4推理效率实现接近FP16的长上下文质量。该方法分两个阶段进行:(1) 一种启发式方法快速选择少量重要的查询-键块对进行FP16精度计算;(2) 选中的块以FP16计算,其余块以FP4计算,两者通过在线softmax合并为单一输出。我们在长上下文基准测试和多个模型家族中证明,仅需将5%的查询-键块计算提升至FP16,ThriftAttention即可平均恢复FP4→FP16性能差距的89.1%。我们展示了ThriftAttention的优势随序列长度增加而增长,有效缓解了在较长上下文中观察到的系统性FP4质量下降。代码开源在:https://github.com/joesharratt1229/ThriftAttention。

---

## 1 引言

高效推理对于大型语言模型的部署至关重要 (Pope et al., 2023; Wan et al., 2024)。注意力机制 (Vaswani et al., 2017; Zhang et al., 2025b) 是长上下文任务中的关键瓶颈,其二次成本和KV缓存内存流量主导了执行时间 (Kwon et al., 2023; Patel et al., 2024)。NVIDIA的Blackwell架构 (NVIDIA Corporation, 2024) 引入了原生FP4张量核心,其算术吞吐量是Blackwell GPU上等效FP16指令的4倍 (NVIDIA Corporation, 2026),同时将KV缓存内存流量减少相似幅度。近期工作,包括SageAttention3 (Zhang et al., 2025c),利用此硬件加速注意力计算。这引入了推理效率与输出质量之间的根本性矛盾。

![图1](Caption: ThriftAttention在保持接近FP16质量的同时,实现接近FP4的延迟。131k上下文长度下(Qwen3-8B)负对数似然(NLL)恢复与推理效率的帕累托前沿。性能恢复衡量为恢复的FP4到FP16 NLL差距的百分比。)

两方面的先前工作分别解决了部分问题,但均未完全解决。FP4注意力方法将质量下降视为提高吞吐量的代价。稀疏性方法 (Tang et al., 2024; Zhang et al., 2025e, 2023) 采取识别重要查询-键交互并仅计算这些交互的方式。然而,在推理的生成阶段,稀疏性方法必须丢弃至少75%的KV块才能匹配FP4延迟。如此激进的稀疏率在纯推理稀疏性方法中可能成为性能下降的主要来源,因为完全遗漏块的误差是不可恢复的。

在这项工作中,我们开发了**ThriftAttention**,一种免训练的混合精度注意力机制,能够以FP4推理效率提供接近FP16的长上下文质量。这解决了低比特注意力变体在长上下文下的退化问题。我们展示了功能相关的量化误差并非均匀分布在查询-键交互中,而是集中在少数注意力分数幅度大且对最终输出分布最重要的交互上。

这一发现启发了一个简单的两阶段方法。首先,一个轻量级启发式方法对每个查询-键块对进行评分:$S_{ij} = \bar{q}_i \cdot \bar{k}_j$,其中$\bar{q}_i$和$\bar{k}_j$分别是查询块$q_i$和键块$k_j$的token均值。得分最高的top-$k$块被选为FP16精度,其余查询-键块计算分配为FP4。对两组块分别计算注意力,然后通过在线softmax合并为单一输出。

#### 结果。

ThriftAttention 恢复了 FP4 和 FP16 注意力之间大部分的质量差距,同时保留了 FP4 推理的效率优势。在 5% 的 FP16 块预算下,ThriftAttention 平均恢复 FP4→FP16 性能差距的 89.1%。当预算提高到 10% 和 25% 时,恢复率分别增加到 91.8% 和 92.4%。在端到端生成中,ThriftAttention 在长上下文下将推理延迟最多降低 2 倍。序列长度分析表明,ThriftAttention 的优势随上下文长度增加而增长,而均匀 FP4 注意力退化最为严重。

#### 贡献。

我们的工作做出以下贡献:

- • 我们引入了 ThriftAttention,一种免训练的注意力方法,将最重要的块交互以 FP16 计算,其余以 FP4 计算。据我们所知,这是首次在注意力计算中以这种混合精度方式使用亚字节格式。
- • 我们在 LongBench-v1、HELMET、RULER 和 PG-19 上,跨越 Llama、Qwen 和 Ministral 模型家族评估了 ThriftAttention,表明少量 FP16 预算可以恢复大部分 FP16 质量,同时保留低比特推理效率。

![图2](Caption: ThriftAttention 概述)

---

## 2 相关工作

**I/O 高效注意力。** FlashAttention (Dao et al., 2022) 引入了分块以减少 GPU 内存 I/O,后续版本 (Dao, 2024; Shah et al., 2024; Zadouri et al., 2026) 提高了并行性并增加了硬件特定优化。

**量化注意力。** 虽然后训练量化在线性层中已经非常成熟 (Dettmers et al., 2022; Frantar et al., 2022; Lin et al., 2024; Dettmers et al., 2023; Ashkboos et al., 2024; Liu et al., 2025),但其在注意力上的扩展仍然有限。SageAttention (Zhang et al., 2025d, a) 通过 INT8/FP8 量化和异常值平滑加速注意力,SageAttention3 (Zhang et al., 2025c) 在此基础上使用两级微尺度缩放,将 FP4 扩展至 Blackwell。其他工作针对 KV 缓存压缩 (Liu et al., 2024; Hooper et al., 2024; Lin et al., 2025) 或将量化矩阵乘法与稀疏性结合 (Kang et al., 2024)。

**稀疏注意力。** Quest (Tang et al., 2024) 使用逐坐标最小-最大边界进行查询感知的 KV 块选择。Token 级别的驱逐和选择策略 (Zhang et al., 2023; Xiao et al., 2024; Li et al., 2024; Jiang et al., 2024) 减少了活跃 KV 集,而 NSA (Yuan et al., 2025) 和 SLA (Zhang et al., 2026) 在训练期间学习稀疏结构。SpargeAttn (Zhang et al., 2025e) 将块稀疏性预测与量化注意力结合,在计算剩余部分之前跳过接近零的块,并以 INT8/FP8 计算余下部分。这与我们的工作最接近,但它将稀疏性作为主要加速机制,且未使用子8位数字格式。其他方法包括线性注意力 (Katharopoulos et al., 2020; Choromanski et al., 2021; Wang et al., 2020; Qin et al., 2024; Yang et al., 2025b),在注意力分布集中时效果不佳。

**定位。** ThriftAttention 将全精度分配给最重要的块,而不是施加均匀量化。对于给定块,误差上限为 FP4 量化噪声,而非如稀疏方法中跳过的或近似的注意力分数大小。

---

## 3 方法

![图3](Caption: Qwen3-8B 中跨层和跨头的典型 FP16→FP4 注意力量化误差 $e = ||P_{FP16} - P_{FP4}||$,按查询/键块划分,序列长度=4096)

### 3.1 ThriftAttention 的动机

考虑单个查询 token 关注 N 个键。注意力输出为

$$ o = \sum_{j=1}^N p_j v_j, \quad p_j = \frac{\exp(s_j)}{\sum_k \exp(s_k)}, \quad s_j = q \cdot k_j / \sqrt{d}. $$ (1)

FP4 量化将每个分数扰动 $\epsilon_j$,得到 $\tilde{s}_j = s_j + \epsilon_j$。一阶输出扰动为

$$ \delta o = \tilde{o} - o \approx \sum_j \frac{\partial o}{\partial s_j} \epsilon_j. $$ (2)

利用 softmax 雅可比矩阵 ($\partial p_j / \partial s_j = p_j(1-p_j)$, 对于 $k \neq j$ 有 $\partial p_k / \partial s_j = -p_k p_j $):

$$ \frac{\partial o}{\partial s_j} = p_j(v_j - o). $$ (3)

代入并取范数:

$$ \|\delta o\| \leq \sum_j |\epsilon_j| \cdot p_j \cdot \|v_j - o\|. $$ (4)

每个键的误差贡献是三项的乘积:$|\epsilon_j|$(分数量化误差)、$p_j$(注意力权重)和 $\|v_j - o\|$(值相对于输出的偏差)。$p_j$ 因子使得这一贡献非均匀:具有大预 softmax 分数的 token 通过 softmax 指数产生大的 $p_j$,从而放大其自身的量化误差。相反,低注意力分数会削弱键 token 的量化误差对 $\tilde{o}$ 的影响。

图3 展示了这一结构。可视化 Qwen3-8B 中代表性层和头的查询-键块对上的 $e = ||P_{\text{FP16}} - P_{\text{FP4}}||$,误差集中在每个查询的一小部分块中。这些通常是近对角块和非初始注意力汇,正是注意力分数最大的地方。

这种集中性表明,均匀 FP4 注意力的大部分质量损失可以通过选择性地将这些高误差块提升为 FP16 并保持其余块为 FP4 来恢复。这在长上下文中最为关键,因为每个 token 的误差在更多位置上累积。

### 3.2 ThriftAttention 算法

#### FP4 量化。

令 $X \in \mathbb{R}^{N \times d}$。我们将 $X$ 量化为 FP4 张量 $X^q \in \mathbb{R}^{N \times d}$ 以及一个 FP8 微尺度张量 $S_X \in \mathbb{R}^{N \times d/16}$。我们使用 Blackwell GPU 支持的 NVFP4 微尺度格式 (NVIDIA Corporation, 2024; Rouhani et al., 2023),其中 $X^q$ 的每个元素以 E2M1 格式存储,$S_X$ 中每个按组尺度以 E4M3 格式存储。此量化独立应用于 $Q$、$K$ 和 $V$,其中 $Q, K, V \in \mathbb{R}^{N \times d}$。

#### 块重要性评分。

我们将 $Q$ 划分为 $T_q = N / B_q$ 个块,将 $K$、$V$ 划分为 $T_k = N / B_k$ 个块:

$$ Q = [Q_1; \dots; Q_{T_q}], \quad K = [K_1; \dots; K_{T_k}], \quad V = [V_1; \dots; V_{T_k}], $$

其中 $Q_i \in \mathbb{R}^{B_q \times d}$ 且 $K_j, V_j \in \mathbb{R}^{B_k \times d}$。令 $\mathcal{B}_i^Q$ 和 $\mathcal{B}_j^K$ 分别表示查询块 $i$ 和键/值块 $j$ 的 token 索引集。我们计算每个块的 token 均值:

$$ \bar{Q}_i = \frac{1}{B_q} \sum_{t \in \mathcal{B}_i^Q} Q_t, \quad \bar{K}_j = \frac{1}{B_k} \sum_{t \in \mathcal{B}_j^K} K_t. $$

块对 $(i,j)$ 的重要性分数为

$$ \hat{S}_{ij} = \bar{Q}_i \bar{K}_j^\top. $$

#### 混合精度注意力计算。

对于每个查询块 $i$,我们选择 top-$k$ 个键块 $\mathcal{T}_i = \operatorname{TopK}( \{\hat{S}_{ij}\}_{j=1}^{T_k}, k )$。查询-键块对路由到 FP4 或 FP16 路径:

$$ S_{ij} = \begin{cases} Q_i K_j^\top / \sqrt{d}, & j \in \mathcal{T}_i, \\[3.0pt] \operatorname{Matmul}_{\mathrm{FP4}}(Q_i^q, K_j^q, S_{Q,i}, S_{K,j}) / \sqrt{d}, & j \notin \mathcal{T}_i. \end{cases} $$

$$ \widetilde{P}_{ij} = \operatorname{OnlineSoftmax}(S_{ij}). $$

输出累积遵循两条路径。对于 $j \in \mathcal{T}_i$,我们遵循标准 FlashAttention-2 在线 softmax (Dao, 2024) 过程。对于 $j \notin \mathcal{T}_i$,概率块通过 SageAttention3 (Zhang et al., 2025c) 的两级方案量化:

$$ (\widehat{P}_{ij}, S_{P,ij}^{(2)}) = \phi(\widetilde{P}_{ij} / s_{P,ij}^{(1)}), \quad O_i \mathrel{+}= s_{P,ij}^{(1)} \cdot \operatorname{Matmul}_{\mathrm{FP4}}(\widehat{P}_{ij}, V_j^q, S_{P,ij}^{(2)}, S_{V,j}), $$

其中 $s_{P,ij}^{(1)} = \operatorname{rowmax}(\widetilde{P}_{ij}) / (448 \times 6)$。FP16 和 FP4 输出更新在线合并。完整过程在算法1中给出,展示的是非因果版本。对于因果 LLM 注意力,我们应用标准因果掩码并限制块选择。

相似文章

FP8注意力中的P-Cast精度:凹陷引发的崩溃与S=2^8的最优性

arXiv cs.AI

本文分析了在将softmax输出转换为FP8(E4M3)时,由于注意力凹陷现象导致的FP8注意力精度损失。它表明正向KV迭代会导致非凹陷注意力值下溢,并提出反向迭代和静态缩放因子S=256来消除下溢,实现了3-10倍的MSE改进。

@VukRosic99: 长上下文Transformer面临两大瓶颈:二次注意力计算和KV缓存(在1M tokens时可达数百GB)…

X AI KOLs Timeline

MiniCPM-SALA是一款9B参数的混合注意力模型,通过在稀疏注意力和线性注意力之间交替插入(每3个线性层插入1个稀疏层)来克服长上下文Transformer的二次计算和KV缓存瓶颈。在256K tokens下,其推理速度比Qwen3-8B快3.5倍,并能在消费级GPU上支持高达1M tokens。该模型采用经济高效的持续训练方法,训练成本降低约75%。