返回首頁

推論系統

TileMix 在注意力矩陣內混用 FP16/INT8,A100 的 4K prefill 吞吐達 FlashAttention 2.22 倍

TileMix 不刪除 token 連線,而是在同一個融合 kernel 內,按注意力分數 tile 選擇 FP16 或 INT8 Tensor Core 路徑。LLaMA 3.2 3B 的 4K 測試達 31.80K token/s,但結果目前限於 A100、靜態路由及 prefill 階段。

https://wellcomeimages.org/indexplus/obf_images/5b/25/577186793973149fae8b5205fc4e.jpg Gallery: https://wellcomeimages.org/indexplus/image/L0005890.html Wellcome Collection gallery (2018-04-01): https://wellcomecollection.org/works/xqav7ue… · CC BY 4.0 · Image source
zh-Hant

TileMix 把混合精度由整個張量或整次 kernel 呼叫,細化成注意力分數矩陣內的硬體對齊 tile group。每個合法 tile group 由一個路由位元決定採用 FP16,或以 INT8 Tensor Core 計算 QK 分數並以 INT32 累加;兩條路徑經尺度還原後,再共同更新同一組 FP16 online-softmax 最大值、正規化因子與輸出累加器。它因此保留完整的稠密注意力連線,不像稀疏注意力直接移除部分 token 互動,也不需要重新訓練模型。

為避免路由判斷拖慢 FlashAttention 類型的內層迴圈,系統把同一 query tile row 的決策壓入 64 位元 bitmask,透過移位與遮罩常數時間查詢。較長序列可讓一個位元控制相鄰多個 key tile,路由中繼資料因而只隨 KV head 數與 query tile row 數成長。公開的 Triton 實作另支援 grouped-query attention、無 padding 的變長批次、INT8 KV cache 介面,以及 LLaMA、Qwen 模型注入程式。

作者在單張 NVIDIA A100 40GB、batch size 8 上,把量化、尺度還原、路由、資料搬移及排程均計入端到端 prefill 時間。LLaMA 3.2 3B-Instruct 的 4K 輸入中,將 75% tile group 路由至 INT8 的 SpTrans 設定達 31.80K token/s;同一封裝下 FlashAttention 為 14.33K token/s,相當於 2.22 倍,而全 INT8 路徑為 29.80K token/s。LongEval 與含中英文資料的 LV-Eval 顯示,FP16 tile 的空間位置與 INT8 覆蓋率都會影響品質;混合設定常能收回全 INT8 遺失的準確度,但高覆蓋率並非在所有資料集都等同 FP16。

這項工作的實際價值是提供可調的精度—吞吐介面,而非宣稱一個普遍無損的加速比。下一步要觀察自適應路由是否值得其額外成本,以及 FP8、Hopper/Blackwell GPU、decode 階段與更大型模型上的表現。現有證據只涵蓋 A100 的 FP16/INT8 路徑;程式庫亦只有兩次提交、沒有正式 release,因此仍需獨立重現與整合測試。

來源

  1. TileMix: Tile-Centric Mixed-Precision Attention for LLM Inference Acceleration
  2. HanzhiZhang-Ulrica/TileMix