← 返回 2026-08-25

TileMix:以Tile为中心的混合精度注意力用于大模型推理加速 TileMix: Tile-Centric Mixed-Precision Attention for LLM Inference Acceleration

Hanzhi Zhang, Qiao Zhang, Qinglei Cao, Heng Fan, Yan Huang, Kewei Sha, Yunhe Feng 📅 2026-08-18 👍 5 2026-08-30 18:30
GPU算子融合 LLM推理加速 Triton 注意力内核 混合精度量化 长上下文

在融合注意力内核内按分数Tile组路由FP16/INT8精度,稠密全连接下加速长上下文预填充

前置知识

FlashAttention 与在线 softmax(online softmax)

一种 IO 感知的精确注意力实现:把 Q/K/V 切成 SRAM 能容纳的 tile,外层遍历 query tile、内层流式读入 K/V tile,在片上用在线 softmax 递推维护行最大值 $\tilde m$、归一化子 $\tilde z$ 与未归一化输出累加器 $\tilde O$,每来一个新 tile 先按新旧最大值之比重缩放旧状态再并入新分数,最后除以 $\tilde z$ 得到输出。它避免在 HBM 中物化 $L\times L$ 的分数与概率矩阵,把中间显存从 $O(L^2)$ 降到 $O(L)$。

TileMix 完全建立在 FlashAttention 式流式内核之上,其全部创新(tile 组精度分派、异构路径共享 softmax 状态)都发生在在线 softmax 递推内部,不懂这个骨架就看不懂方法部分和 Figure 2。

Prefill(预填充)与 KV 缓存

LLM 推理分两阶段:prefill 一次性处理整段提示词,注意力要对全部 $L^2$ 个 query-key 对计算分数,是计算密集型;decode 逐 token 生成,需缓存历史 K/V(KV cache),是访存密集型。长上下文场景下 prefill 的二次方打分是主要时延来源,也是本文的优化对象。

论文所有吞吐数据(如 4k 序列 31.80 K tokens/s)都是端到端 prefill 吞吐,另单独提供 decode 用的 INT8 KV cache 接口;理解两阶段的计算特性差异,才能明白为什么它主打吞吐而非单 token 延迟。

INT8 块量化与 Tensor Core

把浮点矩阵按块(本文 $\text{BLK}_Q=128$、$\text{BLK}_K=64$)线性映射到 8 位整数并记录每块 scale(本文用 absmax:$\delta=\max|x|/127$,最近邻取整),用 NVIDIA Tensor Core 的 INT8 MMA 指令做矩阵乘、在 INT32 寄存器中累加,之后乘回 $\delta_A\delta_B$ 还原数值。INT8 吞吐约为 FP16 的两倍,但引入舍入与累加误差。

TileMix 的 INT8 路径正是 $Q_8K_8^\top$ 的 INT8 MMA + INT32 累加 + scale 还原,它与 FP16 路径在舍入、累加、rescale 上的行为差异,正是共享在线 softmax 状态这一核心设计要解决的问题。

分组查询注意力(GQA)

让 $H_q$ 个 query 头共享 $H_k$ 个 KV 头(映射关系 $h_k=\lfloor h_q H_k/H_q\rfloor$),大幅缩小 KV cache 显存占用,是 LLaMA 2/3、Qwen 等主流开源模型的标准配置。

TileMix 的路由图以 KV 头为索引,并在映射到同一 KV 头的 query 头之间共享同一个 64 位路由字;GQA 支持是其宣称的系统贡献之一,也直接决定路由元数据的形状与开销。

硬件对齐 tile 与 GPU 内存层级

GPU 显存分 HBM(大容量、相对慢)与 SRAM/寄存器(小、极快)。融合内核按 Tensor Core 友好的规则 tile 尺寸(如 128×64)切分计算,tile 同时决定数据搬运粒度与并行工作划分;不规整的访存模式或内层分支会破坏流水线利用率。

TileMix 刻意让路由决策不破坏硬件对齐的 compute tile(分组因子 $g$ 只让一个比特管理 $g$ 个相邻 key tile),并精确核算每个 tile 事件的 HBM/L2 读(16.25 KB)与片上驻留(96 KB),这是它相对稀疏方法的工程优势所在。

研究动机

大语言模型处理长文档摘要、多页问答、检索增强生成等任务时,prefill 阶段的稠密自注意力要对全部 query-key token 两两计算分数,计算量为 $O(L^2)$,是长上下文推理的主要瓶颈。现有三条加速路线各有短板:低精度量化(如 INT8)通常在张量、算子或注意力阶段级别统一指定精度,注意力内部仍是一条算术路径,均匀 INT8 过于粗糙——本文实验中 LLaMA 3.2 3B 在 LV-Eval 数据集 9 的 16k 上下文上 FP16 为 45.68,全 INT8(One)跌到 34.70;稀疏注意力(MInference、FlexPrefill)通过删减 token 交互降算,同设置下分别只有 38.16 和 35.24,且改变了注意力连接性、有漏掉关键交互的风险;FlashAttention 等 IO 感知融合内核虽以 tile 组织数据搬运,但打分精度依然统一。换言之,“哪些交互被执行”与“用什么精度执行”这两个维度,此前从未在融合内核的流式循环内部被联合调度。

本文的目标是本文目标是设计一个无需训练的长上下文 prefill 注意力内核,把数值精度从“算子级配置”下沉为“融合内核内可执行的空间决策”:将注意力矩阵切成硬件对齐的分数 tile,按 tile 组分别分派到 FP16 或 INT8 Tensor Core 打分路径,同时保留全部合法 token 交互(稠密连接不变),让两条异构路径共享同一个在线 softmax 状态。工程目标还包括原生支持 GQA、变长 batch 与 INT8 KV cache,并在 FP16 与均匀 INT8 之间给出一个可用 INT8 覆盖率旋钮连续调节的精度-效率前沿,使预填充吞吐超过 FlashAttention 的同时,长上下文质量显著优于均匀 INT8、接近甚至局部超过 FP16。

与已有工作不同的是,本文的独特切入是把已有两股工作“正交拼接”并塞进内核内层循环。以往混合精度量化把格式分配在张量、算子、阶段或量化块级别,精度路由发生在流式循环之外;稀疏方法用空间结构决定交互“执行与否”。TileMix 反其道而行:保留稀疏注意力研究中的空间模板(band、global、BigBird、SpTrans 等),但把它们从“删除交互的掩码”重新解释为“分配精度的路由图”,在 FlashAttention 式内核的内层循环里以常数时间完成 tile 组精度分派,且因果掩码等合法性判定与精度决策彻底解耦——没有任何交互被删除,只是换了一条算术路径。“稠密连接 + 空间异构精度 + 共享 softmax 状态 + 位图元数据”这一组合,是此前注意力内核工作所没有的。

核心方法

直觉上,注意力矩阵中不同区域对量化的敏感度不同(对角线附近、少数重头交互更敏感),与其全局降精度或删除交互,不如“该细的地方细、能粗的地方粗”。技术上,TileMix 沿用 FlashAttention 的两级循环:外层遍历 query tile $\{Q_m\}$(共 $T_m=\lceil L_q/\text{BLOCK}_M\rceil$ 个),内层把 K/V tile 从 HBM 流式读入 SRAM 和寄存器,维护在线 softmax 状态 $(\tilde m, \tilde z, \tilde O)$——行最大值、归一化子与未归一化累加器。每个(KV 头 $h_k$,query tile 行 $m$)对应一个 64 位路由字,内层循环对每个 key-tile 组 $g_j$ 用一次移位-掩码 $(b_{h_k,m}\gg g_j)\,\&\,1$ 取出精度决策,把分数 tile $S_{m,n}=Q_mK_n^\top/\sqrt d$ 分派到 FP16 matmul 或 INT8 MMA(INT32 累加后按块 scale 还原),两条路径统一到同一浮点分数域后再更新共享的 FP16 在线 softmax 状态,PV 计算始终为 FP16。

核心创新是让“以 tile 组为单位的精度”成为融合稠密注意力内的一等执行抽象。与 SageAttention、FlashINT8 等按调用或阶段统一低精度的量化注意力内核不同,TileMix 的精度选择发生在内层循环的每个 key-tile 组上,用打包位图表达、常数时间移位掩码查找、仅 $O(H_kT_m)$ 元数据;与 MInference、FlexPrefill 等稀疏方法不同,它不删除任何合法交互,完整保留注意力图的连接性与 softmax 归一化分母。另一个关键设计是共享状态异构执行:FP16 与 INT8 路径的舍入、累加、rescale 行为不同,TileMix 让两者在指数化之前进入同一浮点分数域,再共同更新同一个 $(\tilde m,\tilde z,\tilde O)$ 状态,从而把异构算术安全地缝进一个流式 softmax 递推,这是内核层面此前未有的组合。

方法步骤详情

方法分五步。第 1 步,构造路由模板 $R_{h_k,m,g_j}\in\{0,1\}$(1=INT8,0=FP16),布局可选 One/Zero/Band(对角带 FP16)/Global(全局位 FP16)/RowRand/AlignedSparse/BigBird/SpTrans(步进+条带尾 FP16),按 tile 组粒度匹配目标 INT8 覆盖率 $\rho_{\text{INT8}}$(评估 25/50/75%),模板跨层、跨 batch、跨 KV 头广播复用,不用任何评测数据。第 2 步,位图打包:分组因子 $g$ 使 $\text{BLOCK}_{\text{mask}}^N=g\cdot\text{BLOCK}_N$,每行至多 $T_{\text{mask}}\le 64$ 组,打包为 $b_{h_k,m}=\sum_j R_{h_k,m,g_j}2^{g_j}$,key tile $n$ 的组号为 $g_j=\lfloor n/g\rfloor$。第 3 步,量化准备:quantize-once 模式在启动前按 $\delta=\max|x|/127$ 的 absmax 块量化($\text{BLK}_Q=128$、$\text{BLK}_K=64$,每块每头一个 scale)生成 $Q_8,K_8$ 及 scale 供所有 INT8 组复用;另有 fused on-the-fly 模式片上现算以省显存。第 4 步,融合执行:外层每个 Triton 程序持一行 query tile,内层逐 K/V tile 流式计算,按路由位分派 $Q_{8,m}K_{8,n}^\top$(INT8 MMA+INT32 累加,乘 $\delta_Q\delta_K/\sqrt d$)或 $Q_mK_n^\top/\sqrt d$(FP16),两路径同域后更新共享 FP16 在线 softmax,causal/边界掩码独立于路由生效。第 5 步,系统接口:GQA 下 $h_k=\lfloor h_qH_k/H_q\rfloor$ 共享路由字,变长 batch 用 cu_seqlens 前缀和做 padding-free 路由,decode 提供 INT8 KV cache 接口($V_8$ 片上转 FP16 并折入 scale)。

技术新颖性

技术新颖性有四点。其一,粒度:现有量化注意力(SageAttention、FlashINT8、TurboAttention 等)在张量、头、阶段或量化块级定精度,TileMix 把精度决策首次推进到融合内核内层循环的二维分数 tile 组级,且与硬件对齐的 compute tile 结构保持一致。其二,表达与开销:64 位打包位图加分组因子 $g$ 让一个路由字覆盖任意长序列($T_{\text{mask}}\le 64$),查找为常数时间移位掩码,元数据仅 $O(H_kT_m)$,对规整长上下文执行几乎零干扰。其三,数值结构:异构算术路径共享在线 softmax 状态是新的内核组合,附录实验进一步表明 FP16/FP32 累加差异远小于路由布局与覆盖率的影响,说明共享状态设计是稳健的。其四,概念复用:把稀疏注意力的空间先验(band/global/random/BigBird/SpTrans)无损转译成精度预算模板,“覆盖率”取代“稀疏率”成为新的连续旋钮,全程训练-free、即插即用。

TileMix fused attention with tile-group precision routing
Figure 2: TileMix fused attention with tile-group precision routing
Key-tile grouping and bitmask encoding for one query-tile row m
Figure 3: Key-tile grouping and bitmask encoding for one query-tile row m
Constant-time shift-and-mask routing lookup
Figure 4: Constant-time shift-and-mask routing lookup
Storage and execution layout of the two TileMix operand-preparation modes for the primary prefill path
Figure 7: Storage and execution layout of the two TileMix operand-preparation modes for the primary prefill path

实验结果

质量上,混合路由普遍收复均匀 INT8 的损失并常追平甚至超过 FP16。LV-Eval(LLaMA 3.2 3B,11 个中英数据集,16k-64k)上 One 几乎全面落后 FP16(如数据集 9 的 16k:FP16 45.68 vs One 34.70),SpTrans25 达 45.77,并优于稀疏基线 MInference 38.16、FlexPrefill 35.24 与 INT8 内核 SageAttention 42.12;最戏剧性的是 factrecall_en 16k:SpTrans 三档在 LLaMA 3.2 3B 得 21.04/20.53/21.65(FP16 仅 6.72),Qwen2-7B 上 50.6/49.7/52.0(FP16 16.39),Qwen2.5-7B 上 31.4/30.7/32.2(FP16 10.22),跨模型复现。LongEval 行检索(3.1k-38.7k)显示质量恢复取决于 FP16 放哪里(布局)而非只放多少(覆盖率)。效率上,A100 40GB、batch 8 端到端 prefill:4k 时 SpTrans75 达 31.80 K tokens/s,约为 FlashAttention 14.33 的 2.2 倍、略超 One 的 29.80;8k 时 Torch 与 FlashAttention 均 OOM,TileMix 仍达 26.6 K tokens/s(One 136.76 TOPS)。数值上偏差随覆盖率单调可控:1k 序列 0%/5%/10%/25% 覆盖对应 $7.27\times10^{-5}/7.47\times10^{-4}/1.19\times10^{-3}/2.03\times10^{-3}$;SpTrans25 仅把约 8.5% 的高重要性注意力质量暴露给 INT8,低于名义 25% 覆盖率,解释了结构化布局的质量优势。

LV-Eval long-context question answering accuracy for LLaMA 3.2 3B
Table 1: LV-Eval long-context question answering accuracy for LLaMA 3.2 3B
Implementation-level prefill throughput (Thpt, K tokens/s) and TOPS for LLaMA 3.2 3B-Instruct
Table 2: Implementation-level prefill throughput (Thpt, K tokens/s) and TOPS for LLaMA 3.2 3B-Instruct
Mean absolute output deviation from the fixed Torch FP16 reference under different score-tile-group INT8 coverage ratios
Table 3: Mean absolute output deviation from the fixed Torch FP16 reference under different score-tile-group INT8 coverage ratios
Notation for attention, tiling, online softmax, quantization, head mapping, and routing
Table 4: Notation for attention, tiling, online softmax, quantization, head mapping, and routing
Memory and data-movement accounting for FlashAttention and TileMix
Table 5: Memory and data-movement accounting for FlashAttention and TileMix
Line-level retrieval accuracy on the LongEval benchmark for Vicuna 7B
Table 6: Line-level retrieval accuracy on the LongEval benchmark for Vicuna 7B
LV-Eval long-context question answering results on Qwen2-7B
Table 7: LV-Eval long-context question answering results on Qwen2-7B
LV-Eval long-context question answering for Qwen 2.5 7B
Table 10: LV-Eval long-context question answering for Qwen 2.5 7B
Throughput (Thpt, K tokens/s) and TOPS on Qwen 2.5 7B across sequence lengths
Table 12: Throughput (Thpt, K tokens/s) and TOPS on Qwen 2.5 7B across sequence lengths
Throughput (Thpt, K tokens/s) and TOPS on Vicuna 7B across sequence lengths
Table 14: Throughput (Thpt, K tokens/s) and TOPS on Vicuna 7B across sequence lengths
Single-layer model on random inputs
Table 15: Single-layer model on random inputs
Numerical behavior of a 12-layer attention model on random inputs
Table 16: Numerical behavior of a 12-layer attention model on random inputs
Direct numerical difference between TileMix and FlashAttention under different precision layouts
Table 18: Direct numerical difference between TileMix and FlashAttention under different precision layouts
Numerical differences compared to full FP16 with FP32 accumulation under different precision layouts and mixing ratios
Table 19: Numerical differences compared to full FP16 with FP32 accumulation under different precision layouts and mixing ratios
Direct comparison between FP16 and FP32 accumulation under different precision layouts and mixing ratios
Table 20: Direct comparison between FP16 and FP32 accumulation under different precision layouts and mixing ratios
Layer-wise numerical differences between FP16 and FP32 accumulation on LLaMA 3.1 8B
Table 21: Layer-wise numerical differences between FP16 and FP32 accumulation on LLaMA 3.1 8B
Layer-wise numerical differences between FP16 and FP32 accumulation on Qwen 2.5 14B
Table 22: Layer-wise numerical differences between FP16 and FP32 accumulation on Qwen 2.5 14B
Line-level retrieval accuracy on the LongEval benchmark for LLaMA 3.2 3B under different tile-group routing layouts
Figure 5: Line-level retrieval accuracy on the LongEval benchmark for LLaMA 3.2 3B under different tile-group routing layouts
Line-level retrieval accuracy on the LongEval benchmark for Qwen 2.5 7B under different tile-group routing layouts
Figure 8: Line-level retrieval accuracy on the LongEval benchmark for Qwen 2.5 7B under different tile-group routing layouts
Line-level retrieval accuracy on the LongEval benchmark for Qwen 2 7B under different precision policy layouts
Figure 9: Line-level retrieval accuracy on the LongEval benchmark for Qwen 2 7B under different precision policy layouts
查看结构化数据
任务指标本文基线提升
LV-Eval 长上下文问答(LLaMA 3.2 3B,数据集 9 loogle_SD,16k) 准确率 SpTrans25 = 45.77 One(全INT8)34.70;FP16 45.68 较 One +11.07 分并反超 FP16,同时高于 MInference 38.16 与 SageAttention 42.12
LV-Eval factrecall_en 16k(LLaMA 3.2 3B) 准确率 SpTrans 21.04/20.53/21.65(25/50/75% INT8) FP16 6.72;SageAttention 5.88;MInference 5.41 约为 FP16 的 3.1 倍,所有稀疏与 INT8 基线的 3.5 倍以上
LV-Eval factrecall_en 16k(Qwen2-7B / Qwen2.5-7B) 准确率 SpTrans 50.6-52.0(Qwen2);31.4-32.2(Qwen2.5) FP16 16.39 / 10.22;One 10.7 / 6.7 约 3.1 倍于 FP16,跨模型复现的布局-任务交互
预填充吞吐(LLaMA 3.2 3B-Instruct,4k 序列,batch 8,A100 40GB) K tokens/s SpTrans75 = 31.80 FlashAttention 14.33;One 29.80 +122% vs FlashAttention,+6.7% vs One
预填充吞吐与算力(8k 序列,同配置) K tokens/s / TOPS SpTrans75 26.61 K tokens/s;One 136.76 TOPS Torch 与 FlashAttention 均 OOM FlashAttention 无法运行的设置下仍保持约 2 倍于其 4k 吞吐的速度
LongEval 行级检索(Vicuna 7B,500 行约 11.7k tokens) 精确匹配率 混合布局约 0.80-0.83 FP16 0.85;One(INT8)0.76 回收约一半的均匀 INT8 精度损失(700 行处 0.45-0.50 vs One 0.44/FP16 0.52)
随机输入数值偏差(单层注意力,1k 序列,对照 Torch FP16) 平均绝对偏差 25% INT8 覆盖 $1.95\times10^{-3}$(Table 15) 100% INT8 $5.40\times10^{-2}$ 偏差降低约 27 倍,覆盖率构成实用的数值控制旋钮

局限与改进

作者承认的局限:方法只针对 long-context prefill 的前向推理,decode 仅以独立的 INT8 KV cache 接口简要描述;当前实现绑定 NVIDIA A100 的 FP16/INT8 Tensor Core 路径,其他数值格式需要格式特定的 scale 处理与内核调度;路由是静态、无数据的结构化模板,虽然内核接口可以消费自适应策略,但论文未实现。我的补充观察:其一,布局-任务交互强且部分反直觉——factrecall_en 16k 上混合精度反而比 FP16 高 2-3 倍,作者仅归为 layout-task interaction 而未给出机理,提示结果对基准的敏感性;其二,25/50/75% 指 tile 组覆盖率而非 FLOP 占比,实际加速与覆盖率呈非线性(Band75、Global75 等高覆盖布局常劣于保守覆盖);其三,量化用最简单的逐块 absmax 对称整型,没有离群值校准(对比 SmoothQuant 类方法),在激活更重尾的更大模型上稳健性未知;其四,评测集中在 3B/7B 模型与检索/QA 类任务,缺少生成质量(困惑度、摘要)与 H100/FP8 上的验证;其五,quantize-once 的算子级 HBM 驻留为 136.01 MiB,比 FlashAttention 的 96 MB 多约 40 MiB,显存紧张场景是隐性成本。

独立分析的弱点

第一,静态模板与输入无关:同一张路由图被所有请求、所有层共享,无法针对具体输入的重头交互做在线保护——Table 23 显示 SpTrans 碰巧只把 8.57% 的高重要性质量暴露给 INT8,这更像结构性幸运而非设计保证;改进方向是把轻量统计(如每头累积注意力质量分布)离线校准进模板,或真正接入内容感知策略。第二,路由按 KV 头 × query 行共享且跨层广播,粒度仍偏粗,而数值分析表明偏差随深度累积(Qwen2.5-14B 逐层均值从 $10^{-5}$ 升至约 $1.4\times10^{-4}$),可为不同层配不同覆盖率形成敏感度驱动的分配。第三,INT8 路径无离群值处理,直接 absmax 量化 K 在存在激活离群值的模型上可能失效,可引入 per-channel scale 或 Hadamard 旋转不变变换。第四,效率叙事部分依赖 FlashAttention 在 8k/batch 8 时 OOM,应在 24GB 消费卡、batch 1 及 H100 FP8 上补齐对比以增强说服力。第五,factrecall 上的异常增益缺乏解释,应补充统计显著性检验、多随机种子以及代码、摘要等更广任务验证,避免对检索类基准的隐性过拟合。

未来方向

作者指出的方向:内核接口已能消费任意静态或自适应路由策略,下一步自然是数据驱动/内容感知的精度路由——类似 MInference 的动态模式识别,但作用于精度而非连接性;扩展到 FP8/INT4 与 Hopper/Blackwell 架构需要格式特定的 scale 与调度设计;精度路由与稀疏化正交,可叠加成“精度路由 + 交互剪枝”的复合加速;INT8 KV cache 的 decode 路径值得与 prefill 统一为全流程混合精度推理。基于本文成果可延伸的研究:把覆盖率当超参数做逐层/逐头自动搜索(把 ResQ 等低秩残差混合精度思想搬到 tile 粒度);研究布局-任务交互的机理,例如用注意力熵或重头集中度预测最优布局;把 TileMix 接入 vLLM/SGLang 等服务框架做端到端吞吐、成本与能耗评估;探索训练时感知 tile 级量化噪声的量化感知训练,进一步提高可容忍的 INT8 覆盖率上限。

复现评估

复现条件较好。代码已在 GitHub 开源,依赖公开基准(LongEval、LV-Eval)与开源权重模型(LLaMA 3.2 3B、Qwen 2/2.5 7B、Vicuna 7B),方法训练-free,推理侧即插即用。硬件门槛为一块 NVIDIA A100 40GB:论文所有主结果都在该配置下测得(batch 8、3 次预热、5 次计时,计时含量化、scale 还原、路由、内存搬运与内核调度的端到端 prefill 开销),单卡即可复现主要数字。数值实验(单层/12 层/32 层注意力 + 随机输入)成本极低,适合快速验证。难点主要在工程侧:内核为 Triton 实现并针对 A100 离线 autotune,换 GPU(尤其非 A100 或消费级卡)需重新调优;位图打包、GQA 头映射、变长前缀和元数据等细节需仔细对齐才能复现精确吞吐。质量实验矩阵庞大(8 布局 × 3 覆盖 × 4 模型 × 多长度),论文以完整表格披露(Vicuna 因 16k 上限改用表格),可核对性好。总体属“可复现但需要 GPU 系统与内核工程经验”的中等难度。