TileMix:以Tile为中心的混合精度注意力用于大模型推理加速 TileMix: Tile-Centric Mixed-Precision Attention for LLM Inference Acceleration
在融合注意力内核内按分数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、即插即用。
实验结果
质量上,混合路由普遍收复均匀 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 长上下文问答(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 系统与内核工程经验”的中等难度。
论文图表
对比四类注意力加速策略的执行示意图:量化方法降低权重/激活精度但注意力 softmax 与累加常保持高精度;稀疏/块注意力只执行被选中的 token 交互;TileMix 保留全部合法交互,并在单个融合内核内把分数 tile 组路由到 FP16 或 INT8 路径。蓝色为高精度、灰色为低精度、白色为被删除的交互。
一图看懂本文与量化、稀疏两条主流路线的本质区别——“不删交互,只换精度”,是理解全文定位的起点。
LLaMA 3.2 3B 在 Multi-News 数据上的跨层、跨头注意力图可视化,显示注意力值在 query-key 平面上非均匀分布(不同层/头的模式差异明显),为结构化 tile 组精度分配提供定性动机。
回答“为什么精度值得按空间结构分配”:注意力质量天然非均匀,均匀降精度必然浪费预算,这是方法动机的实证支撑。
LLaMA 3.2 3B 的完整 LV-Eval 结果矩阵:FP16/One 与 6 种布局 × 3 种覆盖率的全部数字,含数据集 3 16k 处 SpTrans 21.04/20.53/21.65(FP16 仅 6.72)的异常强势区间。
主表的完整版,是核查布局-覆盖率-任务交互细节(尤其 factrecall 异常)的一手数据。
Qwen 2 7B 的完整 LV-Eval 矩阵:数据集 3 16k 处 SpTrans 达 50.6/49.7/52.0(FP16 16.39),而多数布局在其他设置下居于 One 与 FP16 之间。
跨模型验证混合路由质量的最强单点证据(比 FP16 高约 3 倍),也是分析布局-任务交互的关键数据。
LLaMA 3.2 3B-Instruct 上 Torch/Flash/One 与全部布局×覆盖率的吞吐与 TOPS 全表:覆盖率越高吞吐越高,混合配置多与 One 相当或略低,1k 处 Band75 达 34.56 K tokens/s。
效率主表的扩展版,揭示覆盖率-吞吐关系与布局间吞吐差异的完整图景。
Qwen 2 7B 的吞吐/TOPS 全表:One 约 152-158 TOPS,各布局在 4k 处 13.5-19.2 K tokens/s,FlashAttention 仅 7.09。
补充第二个 7B 模型的效率证据,显示 TileMix 相对 FlashAttention 的加速在 7B 上同样接近或超过 2 倍。
32 层模型的最大绝对偏差:Flash 约 0.32-0.38,100% INT8 约 0.36-0.45,25% INT8 约 $5.0\text{-}5.9\times10^{-3}$。
深度扩展实验的终点,确认覆盖率而非深度是偏差的主导因素。
静态 SpTrans 布局下高重要性交互被路由到 INT8 的加权暴露率:SpTrans25 在 Top5/10/20/30 阈值下仅 8.48-8.63%,SpTrans50 约 17.8-19.2%,SpTrans75 约 21.2-21.7%,均低于名义覆盖率。
解释结构化布局质量优势的机理:空间模板无需在线选择就天然把大量高重要性注意力质量留在 FP16。