FlashMorph:讓 Transformer 高效「變形」成混合注意力模型
Morphing into Hybrid Attention Models
背景
大型語言模型要處理長文本(例如 128K、256K 甚至 512K token 的上下文)時,標準 Transformer 的全注意力(full attention)機制會是效能瓶頸,因為它的計算與記憶體成本會隨序列長度呈平方成長。近年來業界提出「混合注意力模型」(hybrid attention model)作為折衷方案:只保留一部分層維持全注意力,其餘層則替換成計算量接近線性成長的線性注意力(linear attention),藉此在長文本推論時大幅降低延遲與記憶體開銷,同時盡量維持模型的長文本檢索與推理能力。
這種做法的關鍵問題在於——「該保留哪些層的全注意力?」不同層對模型整體能力的貢獻並不相同,選錯層可能導致長文本檢索(long-context recall)能力明顯下滑。過去的做法大致分兩種:一種是固定的排列模式(例如每隔幾層保留一層全注意力),另一種是逐層打分(layerwise scoring),即單獨評估每一層的重要性後再挑選分數最高的層組合。這兩種方法都隱含一個假設——層與層之間的重要性是「獨立」的,可以分開評估。但論文指出,這個假設其實有問題:在一個全局的混合架構配置下,各層之間會互相影響、互相依賴,單獨評估某一層是否重要,並不能反映它在整個混合配置中真正的作用。此外,像 KL-guided scoring(KL-LS)這類需要對大量數據做評估的方法,計算成本也相當高昂,難以規模化應用到更大的模型或更多層數選擇的搜尋空間中。
因此,本篇論文由 Disen Lan、Jianbin Zheng、Yuxi Ren 等作者提出,將「混合層選擇」重新定義為一個「有預算限制的子集優化問題」(budget-constrained subset optimization problem),並提出了名為 FlashMorph(Fast LAyer Selection for Hybrid MORPHing)的新方法,試圖同時解決效果不佳與計算成本過高兩個痛點。
方法
FlashMorph 的核心思路是不再對每一層「單獨打分」,而是讓所有層的重要性在一個統一的、可微分的框架下「一起被學習」,從而捕捉層與層之間的交互效應。整體流程分為四個階段:
第一步:構建可變形模型(morphable model)。 對原始 Transformer 中每一個全注意力層,額外接上一個由該層轉換而來的線性注意力分支(linear-attention branch),讓每一層同時擁有全注意力與線性注意力兩種「模式」可供選擇,並透過隱藏狀態對齊蒸餾(hidden-state alignment distillation)確保轉換後的線性分支能盡量逼近原本全注意力層的行為。
第二步:聯合優化層級門控(gate)。 為每一層引入一個可學習的標量門控值 α^(l) ∈ [0,1],代表該層傾向使用全注意力還是線性注意力的程度。此時凍結所有模型權重,只在合成的長文本檢索資料(synthetic long-context retrieval data,例如 passkey 檢索任務)上聯合優化這些門控值。優化過程中加入「線性化正則項」(linearization regularization),鼓勵模型盡可能多依賴線性注意力以換取效率,只有在真正必要時才保留全注意力。這一步是整套方法的關鍵——因為所有層的門控是「同時」被優化的,能夠反映出全局配置下的層間依賴關係,而不是像過去方法那樣孤立評估。
第三步:離散化與架構實例化。 在優化收斂後,依照預設的「全注意力層預算」(例如保留 1/4 或 1/3 的層維持全注意力),把連續的門控值離散化,挑出 top-K 個門控值最高的層作為最終混合架構中保留全注意力的層,其餘層則正式替換為線性注意力。
第四步:恢復訓練(recovery)。 對這個新的混合架構,進行標準的 logits 蒸餾(以原始 Transformer 的輸出分佈作為監督訊號)以及長文本微調(long-context finetuning),讓模型的整體能力和長文本表現恢復到接近原模型的水準。
實驗中,FlashMorph 在 Qwen3 系列模型(0.6B、1.7B、8B,以及 MoE 架構的 30B-A3B)上進行測試,並搭配三種主流線性注意力變體:Lightning Attention、Gated Linear Attention(GLA)以及 Gated DeltaNet(GDN),驗證方法的通用性。
實驗結果
在效果面,FlashMorph 找到的混合層配置在多項基準測試上都優於既有的層選擇方法,包括固定排列(Uniform Interleaving)、基於超網路搜尋的 PostNAS、KL 引導打分的 KL-LS,以及逐層替換評估的 HALO。
在長文本檢索能力方面,使用 Needle-in-a-Haystack(NIAH)測試、上下文長度從 32K 延伸到 256K token:Qwen3-1.7B 搭配 GDN 變體時,在 32K 長度的 NIAH-Single-1 任務上達到 100% 準確率,即使拉長到 256K,仍維持 88.2% 的準確率,且在多個變體組合下都優於 HALO 與 KL-LS。
在一般能力與知識密集任務上,以 0.6B 模型為例:常識推理類基準(ARC-Easy 63.1%、PIQA 67.8%、HellaSwag 47.5%,平均 62.1%)以及檢索密集型任務(SQuAD 38.4%、FDA 71.3%、SWDE 76.7%,平均 62.1%)均維持在有競爭力的水準,顯示轉換為混合架構並未明顯犧牲通用能力。
而在效率這個核心賣點上,以 Qwen3-1.7B 為例的層選擇成本對比極為顯著:FlashMorph 只需 2000 萬 token 的訓練量與 2.1 GPU 小時即可完成層選擇;相較之下,HALO 需要 2.34 億 token、15.4 GPU 小時;KL-LS 更需要高達 200 億 token、1071.8 GPU 小時。換算下來,FlashMorph 比 HALO 快約 7.3 倍,比 KL-LS 快約 510 倍,大幅降低了搜尋混合配置所需的算力投入。
推論速度方面,在全注意力與線性注意力比例為 3:1 的混合架構下,128K 上下文的 prefill 階段可達 2.24 倍加速,256K 時提升到 2.81 倍;decode 階段在 256K 時加速 1.56 倍,512K 時加速 2.07 倍,證明轉換後的模型在真正長文本場景下能帶來實質的推論效能提升。
意義
FlashMorph 的貢獻可以從兩個層面來看。第一,它重新定義了「混合層選擇」這個問題的本質——不是逐層獨立評分,而是全局聯合優化的子集選擇問題,這個視角轉變讓模型能夠真正捕捉層與層之間的協同效應,找到品質更好的混合配置,而不只是把每一層的「局部最優」簡單拼在一起。
第二,也是更具實用價值的一點,是它把原本動輒需要數百 GPU 小時甚至上千 GPU 小時的搜尋成本,壓縮到僅需 2 GPU 小時左右,且訓練資料量減少了三個數量級(從 200 億 token 降到 2000 萬 token)。這意味著即便是資源有限的團隊或研究者,也能對大型模型(甚至像 30B-A3B 這種 MoE 架構)進行 Transformer 到混合注意力架構的轉換嘗試,而不必依賴大型實驗室級別的算力預算。
對於正在部署長文本應用(如長文檔問答、程式碼庫級別的程式輔助、多輪長對話代理)的公司而言,這類「morphing」式的轉換方法提供了一條務實的路徑:不需要從零重新訓練一個全新的線性注意力或混合架構模型,而是可以直接把現有訓練好的 Transformer(如 Qwen3 系列)以相對低廉的成本改造成兼顧效率與長文本能力的混合模型。這種「舊模型再利用」的思路,對於降低長文本推論的服務成本、加快混合架構的普及,具有相當直接的實務意義。