注意力按硬件tile混精度,A100上4k prefill吞吐约为FlashAttention的2.2倍

TileMix: Tile-Centric Mixed-Precision Attention for LLM Inference Acceleration

Hanzhi Zhang, Qiao Zhang, Qinglei Cao, Heng Fan, Yan Huang, Kewei Sha, Yunhe Feng

cs.AI

2026-08-18

TileMix在融合注意力内核里按硬件对齐的score tile把QK分到FP16或INT8,不砍任何合法连接。LLaMA 3.2 3B在A100上4k prefill吞吐31.80K token/s,约为FlashAttention的2.2倍,长上下文质量接近全FP16。

这篇在解决什么

长上下文 prefill 的瓶颈是 dense self-attention:序列长度 L 时,QK 分数是 O(L²)。现有加速分三条路。量化把权重和激活压到 INT8,一次 kernel 调用通常只走一条精度路径。稀疏注意力靠丢掉一部分 token 交互换速度,改的是连接图。FlashAttention 把 score、softmax、PV 融进 SRAM tile,解决的是 IO,精度仍然整核统一。

缺一块空间决策:合法连接全部保留,但每个硬件对齐的 score tile 自己选 FP16 还是 INT8。North Texas LLaVi Lab 和 Saint Louis University 的 TileMix 做的就是这件事。

方法

精度在这里变成 fused dense attention 里可执行的空间决策,免训练。

注意力矩阵切成 BLOCKM × BLOCKN 的硬件 tile,对应 FlashAttention 的 query/key 块。沿 key 维把相邻 compute tile 合成 routing group,每个 group 一个 bit:1 走 INT8,0 走 FP16。每个 KV head、每行 query tile 压进一个 64-bit 字。分组因子 g 让 key 变长时 routing 位数不超过 64,查找是移位加掩码,元数据规模是 O(Hk × Tm)。

可以把它想成把棋盘按格子上色:有的格子用细笔(FP16),有的用粗笔(INT8),棋子一个都不拿掉。

INT8 路径只量化 Q 和 K。query 块 128×d,key 块 64×d,对称 absmax 除以 127,INT8 Tensor Core 做 MMA、INT32 累加,再用块 scale 和 1/√d 还原。V 和 PV 全程 FP16。两条路径还原后进入同一浮点 score 域,更新共享的 FP16 online-softmax:行最大值、归一化项、输出累加器。因果掩码和边界合法性独立于路由图,不合法的交互本来就不会算。

评测不用内容自适应,也不在线找 heavy hitter。路由是静态、无数据的模板,把稀疏注意力里的空间结构借来当精度布局:Band 保对角局部,Global 保指定位置,Row-Random 每行随机留 FP16,Aligned Sparse 右对齐扩展,BigBird 把前三者并起来,SpTrans 用 stride 加尾部。报告的 25%、50%、75% 是合法 tile group 的 INT8 占比,不是 FLOPs 占比。同一模板跨层、跨 batch、跨 KV head 复用。支持 grouped-query attention、变长 batch,以及 decode 侧的 INT8 KV cache 接口。实现是 Triton,A100 上 autotune 一次后固定。

结果

主模型 LLaMA 3.2 3B,附录扩到 Vicuna-7B、Qwen-2-7B、Qwen-2.5-7B。硬件是 NVIDIA A100 40GB。对照包括全 FP16、同一 kernel 上全 INT8(叫 One)、FlashAttention、MInference、FlexPrefill、SageAttention。吞吐是端到端 prefill,含量化、scale 还原、路由和调度,batch 8。

LV-Eval 长上下文问答,LLaMA 3.2 3B 一组 16k/32k/64k 数字:

方法16k32k64k
FP1632.0415.087.75
One(全 INT8)28.7811.625.42
SpTrans2531.7515.568.01
SageAttention29.7913.505.77
MInference26.9311.555.31
FlexPrefill27.1111.745.39

全 INT8 明显掉点。混精度把质量捞回 FP16 附近,部分长度略高。稀疏基线砍连接,多数子集低于 dense 混精度。LongEval 逐行检索上,布局比名义覆盖率更要紧:rowrand 和 sptrans 在 INT8 比例升高时更稳,alignsparse、band、global 更适合保守覆盖。

论文还写 SpTrans 在 16k factual recall 上跨 LLaMA 和 Qwen 都远超 FP16,例如 LLaMA 该子集 FP16 是 6.72、SpTrans25 是 21.04。这个跳变不像普通量化误差。

Prefill 吞吐(LLaMA 3.2 3B-Instruct,K token/s):

长度FlashAttentionOneSpTrans75SageAttention
1k17.4532.2733.5019.91
4k14.3329.8031.8019.91
8kOOM27.4126.6118.79

4k 上 SpTrans75 是 31.80,FlashAttention 是 14.33,约 2.22 倍,接近全 INT8。8k 时 PyTorch 和 FlashAttention 在该协议下 OOM,TileMix 还能跑。相对 Torch FP16 的平均绝对偏差:0% INT8 约 10⁻⁵;8k、10% 覆盖升到 6.32×10⁻³,25% 为 6.84×10⁻³。附录里 SpTrans25 只把约 8.5% 的高重要性 attention mass 送到 INT8,低于名义 25% tile 覆盖。

为什么重要

给推理系统一个新旋钮:精度可以按硬件 tile 空间分配,不必在全高精度、全 INT8 和砍连接之间单选。免训练,能接 GQA 和变长 batch,部署门槛低。

适合长文档 prefill 已经是瓶颈、机器有 INT8 Tensor Core、又能接受静态布局的场景。这是渐进的系统工作,不是新注意力算法。布局和覆盖率要按模型和任务选,没有万能模板。吞吐含完整流水线,高覆盖混配有时比 One 还快,dispatch 和访存布局会吃掉或送出一点速度。

局限与存疑

作者写了三点。只评了 prefill。kernel 落在 A100 的 FP16/INT8 Tensor Core,换 FP8 或 INT4 要重做 scale 和调度。评测路由是静态模板,接口能吃自适应策略,但没拿来比。

另外几处要自己掂量。质量主表是 3B,7B 在附录,没有 70B 级。吞吐只报到 8k,质量报到 64k,两端对不齐。覆盖率按 tile group 计,不是 FLOPs。SpTrans 在部分 16k factual recall 子集上大幅超过 FP16,论文当成布局-任务交互,幅度过大,更像评测与模板的耦合,需要独立复核。静态路由不在线检测 heavy hitter,换域或换模型可能失效。decode 有 INT8 KV cache 接口,正文数字几乎全是 prefill。

术语

原文与代码

相关论文

全部论文解读