你有沒有想過,當你讓ChatGPT這類大語言模型讀一篇幾萬字的長文檔時,它到底在"卡"哪裡?
答案往往不是它讀不懂,而是它算不過來。模型在讀長文本的第一步,也就是所謂的"預填充"階段,需要讓文本里的每一個詞和其他所有詞都兩兩比對一次,計算它們之間的關聯程度。這個計算量會隨著文本長度的平方增長,文本翻一倍,計算量漲四倍。這就是為什麼處理長文檔時,模型響應會明顯變慢。
這個環節有個專業名字。
> 自注意力機制
:大語言模型理解上下文的核心方法,讓每個詞都能"看到"並衡量與其他所有詞的關聯程度,從而理解語義。
面對這個平方級增長的計算量,過去幾年研究者們主要想了兩條路子。一條是降低數字精度,比如把原本用16位小數表示的數值,壓縮成8位整數來算,這樣速度快、省內存,但精度會打折扣。另一條是減少計算量本身,乾脆不去算某些詞之間的關聯,用稀疏的連接模式代替完整的兩兩比對。
這兩條路子都有效,但都留了一個空白地帶沒人碰過。這篇來自北德克薩斯大學和聖路易斯大學團隊的論文,叫做TileMix
,就是盯著這塊空白下的手。
先搞懂:這塊"空白地帶"到底是什麼
要理解TileMix填的是什麼坑,得先弄明白現在主流的高效注意力計算是怎麼運作的。
現在業界公認最能打的注意力實現方案叫FlashAttention
,它的核心思路是把一個巨大的注意力計算矩陣切成很多小方塊(術語叫"tile",瓦片),一塊一塊地搬進GPU的高速緩存里算,算完就扔,不把整個大矩陣都攤在顯存里占地方。這樣既省顯存又提速。
> FlashAttention:一種通過分塊計算和流式處理注意力矩陣的技術,避免把巨大的中間結果直接攤在顯存里,從而節省內存頻寬並提速。這是當前幾乎所有高效注意力實現的技術基礎。
這個方案有個特點:不管切成多少塊,每一塊內部用什麼精度去算,是固定死的。要麼整個注意力計算過程都用16位浮點數(FP16),要麼統一換成8位整數(INT8),沒有中間地帶。
而稀疏注意力那條路呢,是在決定哪些詞對之間要不要建立連接,跟精度完全是兩碼事。它管的是"連不連",不管"怎麼算"。
TileMix團隊的洞察是:這兩條路子分別管住了"哪裡計算"和"用什麼數值格式計算"這兩個維度里的一個,但從來沒有人把"每一小塊矩陣該用什麼精度"當成一個可以按空間位置單獨決定的執行選項。換句話說,之前的方法要麼是全局統一精度,要麼是全局統一連接方式,但沒人試過讓精度本身像連接方式一樣,按照矩陣里的位置來分區管理。
這就好比一家餐廳的後廚,以前要麼整個後廚統一用高壓鍋(快但可能把食材壓得不夠精細),要麼整個後廚統一用小火慢燉(精細但慢),從沒人想過按菜品所在的灶台分別配置:靠近門口的幾個灶台用高壓鍋處理不太講究的家常菜,最裡面幾個灶台用小火慢燉伺候需要精細火候的招牌菜。如果不這樣區分,你要麼犧牲了所有菜品的口感去換速度,要麼犧牲了所有出餐速度去保口感,沒有折中的空間。
TileMix做的事情,就是把這種"按灶台分配烹飪方式"的思路,搬進了注意力矩陣的計算里。
TileMix到底怎麼設計的
### 第一步:把矩陣切成方塊,再給每塊貼標籤
TileMix的基本操作單元,延續了FlashAttention的思路,把整個注意力矩陣切成一個個硬體友好的小方塊。
> 硬體對齊的計算塊
(tile):為了配合GPU的計算單元(Tensor Core)效率最高的處理粒度而設計的固定尺寸矩陣分塊,通常是幾十到一百多行列的小方陣。
但TileMix多做了一件事:給每一小塊(或者說每一組相鄰的小塊)貼上一個二進制標籤,標籤只有兩種取值,0代表這一塊用FP16去算,1代表用INT8去算。
這個標籤的選取不是隨機的,也不需要額外訓練模型去學習,而是用幾種事先設計好的固定模板來分配。這些模板本身借鑑了之前稀疏注意力研究里總結出的幾種典型注意力分布規律,比如局部窗口模式、全局關鍵位置模式、隨機採樣模式等等,只不過之前這些模式是用來決定"要不要計算",現在被TileMix拿來決定"用什麼精度計算"。
這裡有一個關鍵區別必須說清楚:TileMix從頭到尾都保留了全部的詞對連接,一個都沒有砍掉。它只是在"連接保留"的前提下,決定每個連接用什麼精度去計算。這跟稀疏注意力那種直接切掉部分連接的做法有本質不同。
### 第二步:把決策壓縮進一個64位的數字里
一個長文本可能要切成幾十甚至上百個小方塊,如果給每個塊單獨存一個標記位,管理起來會很麻煩,查找的時候也慢。TileMix的解決辦法是把一整行(對應一個查詢塊)里所有相關方塊的精度決策,打包壓縮進一個64位的整數里,每一個二進制位對應一組方塊的精度選擇。
> 位掩碼
(bitmask):把多個二進制判斷結果壓縮儲存在一個整數的不同二進制位上的技術,查詢時只需要做一次位移和按位與運算就能取出結果,速度極快。
想取出某一組方塊的精度決策時,只需要做一次位移操作加一次按位與運算,這在電腦里幾乎是最快的操作之一,幾乎不占用時間。這就好比你要記住一棟樓里100個房間哪些開燈哪些關燈,與其為每個房間單獨寫一張紙條再逐張翻找,不如把這100個房間的開關狀態編碼成一串二進制數字存在手機備忘錄里,要查第37個房間的狀態,直接算出它在這串數字的第幾位,一眼就能讀出來。如果不這樣壓縮,管理這些海量的精度決策標記會占用大量額外的儲存和查詢開銷,反而抵消了精度優化本身帶來的加速收益。
當文本特別長,方塊特別多,超過64個位不夠用怎麼辦?TileMix設計了一個"分組"機制,讓一個二進制位可以同時管理好幾個相鄰的方塊,就像一個開關同時控制好幾盞燈,這樣即使方塊數量暴漲,64位仍然夠用。
### 第三步:兩條路徑殊途同歸,共享同一套"記賬系統"
這是整篇論文技術含量最高的部分。
在FlashAttention的流式計算里,每處理完一小塊矩陣,都要更新一套共享的狀態變量,包括當前為止見過的最大分數值、歸一化分母、以及累加的輸出結果。這套狀態變量一直被後續的每一塊計算復用,是整個流式計算能夠正確工作的核心機制。
> 在線softmax
:一種邊流式處理矩陣塊邊動態更新歸一化統計量的算法,不需要把整個矩陣都算完才能做歸一化,這也是FlashAttention能夠省內存的關鍵技巧。
問題在於,FP16路徑和INT8路徑產生分數的方式完全不同。INT8那邊要先把浮點數壓縮成整數,用整數乘法算完之後再乘回一個縮放係數還原成浮點數,這個過程會引入捨入誤差和精度損失,跟FP16直接浮點運算的行為不一樣。
TileMix的做法是:不管這一塊是走FP16還是INT8算出來的,在進入共享的狀態更新之前,都要先統一轉換到同一個浮點數值域裡去。這樣兩條路徑產生的分數值,雖然來源不同,但最終都以相同的"語言"匯入同一套記賬系統,不會因為混用兩種精度而搞亂整體的歸一化過程。
這就像一家公司同時收到賬單和美元賬單,財務在做總賬之前,必須先把美元按當天匯率換算成,再放進同一張總賬表里加總。如果不做這一步統一換算,直接把兩種貨幣的數字硬湊到一起相加,得出來的總數毫無意義。TileMix做的正是這件"匯率換算"的工作,只不過換算的對象是數值精度而不是貨幣。
支持的場景:不只是理論上能跑
論文裡特別強調了TileMix在實際部署場景里的幾個支持能力。
它支持分組查詢注意力
,這是現在主流大模型(包括Llama、Qwen這些)普遍採用的一種省顯存的注意力變體,多個查詢頭共享同一組鍵值頭。TileMix讓共享同一個鍵值頭的查詢頭,也共享同一套精度路由決策,不需要重複計算。
它支持變長批處理,也就是同一批次里不同長度的文本可以混著一起處理,不需要都填充到同樣長度浪費計算資源。
它還支持INT8格式的鍵值緩存,這個跟生成階段(而不是本文重點討論的預填充階段)的效率有關,緩存里存的鍵值對用低精度儲存可以省下大量顯存。
這幾項支持能力湊在一起,意味著TileMix不是一個只能在實驗室跑通的原型,而是真的考慮了實際部署會遇到的各種邊角場景。
實驗說了什麼:精度分配確實能省下時間,還能保住準確率
論文在兩個長文本理解任務上做了測試。
第一個任務叫LongEval
,說白了就是讓模型在一堆幾萬字的文本里,精確找出某一行特定編號對應的內容,看它能不能一字不差地把答案摳出來。
第二個任務叫LV-Eval
,覆蓋了11個中英文長文本問答數據集,涉及事實核查、多跳推理、多欄位問答等各種題型,文本長度從16000字到64000字不等。
結果顯示,把注意力全部換成INT8(論文裡叫One配置)之後,模型的表現普遍會明顯下降。以Llama 3.2 3B
模型為例,在64000字長度的多個數據集上,純INT8配置的得分經常只有FP16原始配置的六到七成,掉分幅度相當可觀。
但如果引入TileMix的精度混合路由,情況就不一樣了。比如在16000字長度的"事實召回"這個子任務上,一種叫SpTrans(借鑑了稀疏Transformer設計思路的精度布局)的配置,得分能達到21分左右,而對應的純FP16基線只有6.72分,純INT8配置更是只有4.45分。這個結果在Qwen 2 7B和Qwen 2.5 7B模型上同樣復現出來了,說明這不是偶然的模型特異性現象,而是一種跨模型的規律性發現。
效率方面,在4000字長度的Llama 3.2 3B模型測試里,TileMix的某種混合配置吞吐量達到每秒31.8萬token,而標準FlashAttention只有每秒14.33萬token,純INT8配置是每秒29.8萬token。也就是說,TileMix的混合路由方案不僅沒有比純INT8慢,反而在這個測試點上跑得更快,同時還保留了遠比純INT8豐富的模型質量。
下面這張表格摘錄了論文中一部分關鍵的效率對比數據(單位為每秒千token數):
| 文本長度 | Torch標準實現 | FlashAttention | 純INT8(One) | TileMix混合(75%INT8) |
|---|---|---|---|---|
| 1k | 11.14 | 17.45 | 32.27 | **33.50** |
| 2k | 7.78 | 16.48 | 32.06 | **33.92** |
| 4k | 顯存溢出 | 14.33 | 29.80 | **31.80** |
| 8k | 顯存溢出 | 顯存溢出 | 27.41 | 26.61 |
可以看到,隨著文本變長,標準的Torch實現和普通FlashAttention都相繼出現顯存溢出問題,只有做過精細內存管理的TileMix和純INT8方案還能撐住,而TileMix在多數長度上還能進一步領先純INT8。
數字精度混用到底靠不可靠:論文做的"體檢"
光看任務表現還不夠,研究者還專門做了一輪數值層面的"體檢",檢驗FP16和INT8混合計算之後,模型輸出到底跟純FP16參考結果差多少。
結果顯示,INT8覆蓋比例從0%漲到25%的過程中,輸出數值的平均絕對偏差是逐漸增大的,但這個增大過程是可控、可預測的,覆蓋率越低偏差越小。這說明INT8覆蓋率本身可以當作一個實用的"數值控制旋鈕",需要更保守就調低覆蓋率,能接受更多誤差就調高覆蓋率換取更快速度。
這套體檢還測試了不同模型深度的影響,發現層數越多的模型(比如32層的模型相比12層模型),數值偏差累積得更明顯,這符合直覺,因為誤差會一層一層往下傳遞、放大。
還有一個挺有意思的補充實驗,檢驗的是"靜態路由模板到底有沒有把寶貴的高精度資源用在刀刃上"。研究者專門計算了注意力矩陣里哪些位置的關聯程度對最終結果影響最大(稱為"重要交互"),然後看SpTrans這種精度布局模板有沒有把這些重要交互優先分配給FP16高精度路徑。結果顯示,在名義上25%的方塊被分配給INT8的情況下,實際被劃入INT8的高重要度交互只占到了8.57%,遠低於25%這個整體比例。換句話說,這套靜態模板雖然沒有用到任何實時的重要性檢測機制,卻已經天然地把更多計算資源留給了那些真正舉足輕重的詞對關係。
這就好比一個倉庫管理員,即便完全不知道具體哪些貨物是緊俏商品,只是憑著"靠近出入口的貨架通常放常用品,最裡面角落放冷門貨"這條經驗規則去擺放貨物,結果卻往往歪打正著地把真正的暢銷品擺在了順手可取的位置。這說明經過大量真實場景驗證提煉出的空間分布規律,本身就自帶了某種"重要性感知"的能力,即使沒有專門為此設計檢測機制。
寫在後面
讀到這篇論文時,最讓我意外的不是它提出了一個新點子,而是它指出的那個"空白地帶"竟然一直沒人碰過。低精度量化和稀疏注意力這兩條路線各自發展了好幾年,居然從來沒有人想過把精度當成一個可以按空間位置分配的資源。這有點像兩撥人各自在自己的賽道上狂奔,誰都沒想過賽道之間其實可以架一座橋。
論文裡那個"重要交互暴露率"的實驗,是我覺得最值得單獨拎出來說的細節。研究者沒有滿足於"我的方法效果好",而是反過來去問"為什麼靜態模板會有效",結果發現這些借鑑自稀疏注意力研究的空間分布規律,本身就帶著某種樸素的重要性直覺。這提醒我一件事:很多領域裡積累下來的經驗模式,即便原本是為了解決A問題設計的,挪到B問題上可能依然管用,因為它們背後反映的是數據本身的某種結構性規律,而不是針對具體任務的死板技巧。
這篇論文目前只在英偉達A100這一款GPU上做了驗證,也只測試了FP16和INT8這一對精度組合。如果換成更新的GPU架構,或者嘗試FP8、INT4這些精度格式的組合,這套精度路由的思路是不是依然有效,會不會需要重新設計位掩碼的編碼方式,這是個留白的問題,也是個值得繼續追問下去的方向。
Q&A
Q1:TileMix是什麼,它解決了什麼問題?
A:TileMix是一種針對大語言模型長文本處理的注意力計算加速方法,它把注意力矩陣切成小方塊,讓每個方塊可以單獨選擇用FP16高精度還是INT8低精度計算,同時保留全部詞對連接不做刪減,從而在保持模型理解質量的同時提升長文本預填充階段的計算速度。
Q2:TileMix和普通的INT8量化有什麼區別?
A:普通INT8量化通常是整個計算過程統一換成低精度,而TileMix允許在同一次注意力計算里,不同的矩陣區塊分別選擇FP16或INT8,通過打包的位掩碼實現按區域靈活分配精度,實驗顯示這樣比統一使用INT8能更好地保住長文本理解的準確率。
Q3:使用TileMix之後模型速度能提升多少?
A:論文實驗顯示,在4000字文本長度下,TileMix某些混合精度配置的吞吐量能達到每秒31.8萬token,高於標準FlashAttention的每秒14.33萬token,也略高於純INT8方案的每秒29.8萬token,同時在長文本問答準確率上明顯優於純INT8方案。






