間序列建模的選擇性狀態(tài)空間|從高效序列架構(gòu)演進(jìn)視角)
摘要本文解讀 COLM 2024 杰出論文《Mamba: Linear-Time Sequence Modeling with Selective State Spaces》。該論文提出選擇性狀態(tài)空間模型S6通過融合選擇機(jī)制、硬件感知并行掃描與簡(jiǎn)化同質(zhì)架構(gòu)讓沒有注意力機(jī)制的遞歸模型在語言建模上首次達(dá)到 Transformer 級(jí)質(zhì)量其特別之處在于把 SSM 參數(shù)變成輸入的函數(shù)、按內(nèi)容選擇性傳播與遺忘信息。實(shí)驗(yàn)表明Mamba-3B 匹敵 2 倍尺寸 Transformer、推理吞吐 5 倍、訓(xùn)練線性擴(kuò)展在語言、音頻、基因組三模態(tài)均取得 SOTA并支持百萬級(jí)上下文為高效序列架構(gòu)提供了重要借鑒。視頻講解點(diǎn)擊觀看 B 站視頻摘要論文基本信息背景與動(dòng)機(jī)為什么高效序列模型長(zhǎng)期打不過 Transformer研究主線從問題到結(jié)論基準(zhǔn)/方法設(shè)計(jì)分類全景方法細(xì)節(jié)從 S4 到 S6一個(gè)參數(shù)化的改變門控定理統(tǒng)一 SSM 與 RNN并行掃描與硬件實(shí)現(xiàn)實(shí)驗(yàn)設(shè)計(jì)與結(jié)果評(píng)測(cè)協(xié)議合成任務(wù)選擇性的直接證據(jù)語言建模匹敵兩倍尺寸的 Transformer效率訓(xùn)練與推理雙贏DNA 與音頻百萬級(jí)長(zhǎng)上下文消融$\Delta$ 最重要$B,C$ 協(xié)同結(jié)果對(duì)比總結(jié)關(guān)鍵發(fā)現(xiàn)局限性常見問題FAQMamba 和 Transformer 的核心區(qū)別是什么為什么說 Mamba 的訓(xùn)練是線性時(shí)間的選擇機(jī)制到底選擇什么Mamba 在哪些任務(wù)上不如 TransformerMamba-2 與 Mamba 是什么關(guān)系為什么歸納頭任務(wù)如此重要參考鏈接論文基本信息項(xiàng)目?jī)?nèi)容標(biāo)題英文Mamba: Linear-Time Sequence Modeling with Selective State Spaces標(biāo)題中文Mamba線性時(shí)間序列建模的選擇性狀態(tài)空間作者Albert Gu, Tri Dao機(jī)構(gòu)卡內(nèi)基梅隆大學(xué)機(jī)器學(xué)習(xí)系 · 普林斯頓大學(xué)計(jì)算機(jī)科學(xué)系會(huì)議COLM 2024Outstanding Paper AwardarXivhttps://arxiv.org/abs/2312.00752項(xiàng)目網(wǎng)站https://github.com/state-spaces/mamba背景與動(dòng)機(jī)為什么高效序列模型長(zhǎng)期打不過 Transformer現(xiàn)代基礎(chǔ)模型幾乎全部建立在 Transformer 的自注意力之上。注意力的核心能力是稠密路由信息每個(gè) token 與上下文窗口內(nèi)所有 token 交互這讓它擅長(zhǎng)建模復(fù)雜數(shù)據(jù)但帶來兩個(gè)根本缺陷——無法利用有限窗口之外的信息以及訓(xùn)練復(fù)雜度隨窗口長(zhǎng)度二次增長(zhǎng)推理時(shí)還要維護(hù)隨上下文線性增長(zhǎng)的 KV 緩存。為克服這些缺陷學(xué)界提出了大量次二次復(fù)雜度架構(gòu)線性注意力Linear Attention、門控卷積、遞歸模型以及結(jié)構(gòu)化狀態(tài)空間模型SSM。SSM 家族從 S4 出發(fā)演化出 H3SSM 層兩側(cè)夾門控連接、HyenaMLP 參數(shù)化全局卷積、RetNet并行注意力路徑與 RWKVLTI 遞推WKV 可看作兩個(gè) SSM 之比。這些模型訓(xùn)練線性、推理恒定時(shí)間但在語言這類信息稠密、離散的模態(tài)上始終不如注意力——沒有一個(gè)被證明能在規(guī)模上跨域有效。論文指出共同病根這些模型都是時(shí)不變LTI的參數(shù)不隨輸入變化因而無法做內(nèi)容感知推理——不能根據(jù)當(dāng)前 token 決定記住什么、忘記什么。更根本地序列建模的本質(zhì)是把上下文壓縮進(jìn)有限狀態(tài)注意力完全不壓縮所以慢遞歸模型狀態(tài)有限所以快但效果受限兩者之間的橋梁就是選擇性讓壓縮按內(nèi)容進(jìn)行。研究主線從問題到結(jié)論圖 9Mamba 論文的研究主線——從注意力的效率缺陷出發(fā)定位 LTI 病根用選擇機(jī)制與硬件感知掃描完成從問題到結(jié)論的閉環(huán)。基準(zhǔn)/方法設(shè)計(jì)Mamba 的核心設(shè)計(jì)圍繞三個(gè)支柱展開與既有 SSM 架構(gòu)S4/H3/Hyena形成鮮明對(duì)比。圖 1選擇機(jī)制總覽——先前的 SSM 因時(shí)不變而可避免物化大狀態(tài)選擇性 SSM 把輸入相關(guān)動(dòng)態(tài)放回模型靠硬件感知算法控制內(nèi)存。選擇機(jī)制S6讓 $\Delta$、$B$、$C$ 三個(gè)參數(shù)成為輸入的函數(shù)$B s_B(x)$、$C s_C(x)$、$\Delta \tau_\Delta(\mathrm{Linear}1(x))$$\tau\Delta$ 為 softplus。參數(shù)沿序列長(zhǎng)度維展開模型由時(shí)不變變?yōu)闀r(shí)變。硬件感知掃描參數(shù)直接從 HBM 載入 SRAM離散化與遞推在 SRAM 內(nèi)融合完成用 work-efficient 并行前綴掃描并行化非線性遞推反傳不存中間狀態(tài)、重計(jì)算——內(nèi)存占用與 FlashAttention 級(jí)優(yōu)化 Transformer 相當(dāng)。簡(jiǎn)化架構(gòu)把 H3 塊與 MLP 塊合并為單一 Mamba 塊同質(zhì)堆疊擴(kuò)展因子固定 $E2$使用 SiLU 激活。模型既無注意力也無獨(dú)立 MLP 塊大部分參數(shù)$3ED^2$在線性投影上。分類全景圖 10Mamba 在次二次序列架構(gòu)譜系中的位置——從 S4 到 H3/Hyena/RetNet/RWKV再到引入選擇機(jī)制的 Mamba。方法細(xì)節(jié)從 S4 到 S6一個(gè)參數(shù)化的改變S4 定義連續(xù)系統(tǒng) $h(t) Ah(t) Bx(t)$、$y(t) Ch(t)$經(jīng)零階保持離散化$\bar{A} \exp(\Delta A)$得到遞推 $h_t \bar{A}h_{t-1} \bar{B}x_t$。訓(xùn)練用全局卷積可并行推理切回遞推恒定時(shí)間/步。S6 的關(guān)鍵改動(dòng)$B$、$C$ 變?yōu)檩斎牒瘮?shù)$\Delta$ 由輸入的線性投影經(jīng) softplus 得到。$A$ 可以保持靜態(tài)——因?yàn)?$\bar{A} \exp(\Delta A)$$\Delta$ 的選擇性會(huì)自動(dòng)傳導(dǎo)到離散參數(shù)。這打破了卷積等價(jià)性時(shí)變但換來了按內(nèi)容決定記憶/遺忘的能力。門控定理統(tǒng)一 SSM 與 RNN論文證明了選擇機(jī)制與經(jīng)典門控的聯(lián)系當(dāng) $N1$、$A-1$、$B1$ 時(shí)選擇性 SSM 遞推精確退化為門控 RNN$g_t \sigma(\mathrm{Linear}(x_t))$$h_t (1-g_t)h_{t-1} g_t x_t$由此得到 $\Delta$ 的機(jī)理解釋大 $\Delta$ 重置狀態(tài)、聚焦當(dāng)前輸入小 $\Delta$ 保持歷史、忽略當(dāng)前輸入。這一視角也解釋了為何 $s_\Delta$ 投影到 1 維即可——輸入 $x_t$ 該被忽略時(shí)所有通道應(yīng)一致忽略它。選擇機(jī)制由此帶來三類能力過濾變間距噪聲如語言中的um、過濾無關(guān)上下文性能隨上下文單調(diào)提升、在文檔/回合邊界重置狀態(tài)。并行掃描與硬件實(shí)現(xiàn)遞推模式的 FLOPs 為 $O(BLDN)$低于卷積模式的 $O(BLD\log L)$ 常數(shù)因子。但時(shí)變遞推無法卷積化必須處理兩個(gè)挑戰(zhàn)遞推的串行性與狀態(tài)物化。解法是內(nèi)核融合離散化遞推在 SRAM 內(nèi)完成HBM 只讀寫 $B\times L\times D$ 的輸入輸出 并行前綴掃描 反向重計(jì)算。掃描受內(nèi)存帶寬限制融合是提速關(guān)鍵A100 上比此前 SSM 實(shí)現(xiàn)快 3 倍比樸素掃描快 40 倍。圖 2Mamba 塊結(jié)構(gòu)——兩堆塊對(duì)應(yīng) Transformer 交錯(cuò)的注意力與 MLP 塊的 $12D^2$ 參數(shù)內(nèi)部 SSM 貢獻(xiàn)的參數(shù)很少。實(shí)驗(yàn)設(shè)計(jì)與結(jié)果評(píng)測(cè)協(xié)議四個(gè)設(shè)定合成任務(wù)選擇性復(fù)制、歸納頭檢驗(yàn)內(nèi)容感知能力語言建模用 The Pile 300B tokens覆蓋 125M–1.3B 參數(shù)縮放律Chinchilla 協(xié)議與零樣本下游評(píng)測(cè)基因組用 HG38 預(yù)訓(xùn)練 大猿物種分類微調(diào)上下文 $2^{10}\to2^{20}$音頻用 YouTubeMix 波形預(yù)訓(xùn)練BPB SC09 語音生成NLL/FID/IS。基線包括 GPT3 配方 Transformer、LLaMa 配方 Transformer、H3、Hyena、RetNet、RWKV、SaShiMi。合成任務(wù)選擇性的直接證據(jù)架構(gòu)內(nèi)部層Selective Copying 準(zhǔn)確率S4S4LTI18.3%H3S457.0%H3Hyena30.1%MambaS456.4%MambaS6選擇性99.8%圖 3合成任務(wù)——選擇性復(fù)制與歸納頭直接檢驗(yàn)內(nèi)容感知能力。歸納頭任務(wù)中模型在長(zhǎng)度 $2^8256$ 上訓(xùn)練可外推到 $2^{20}1048576$4000 倍保持高準(zhǔn)確率其他方法最多外推 2 倍——選擇機(jī)制是唯一能外推的關(guān)鍵。語言建模匹敵兩倍尺寸的 Transformer模型Pile ppl ↓LAMBADA acc ↑HellaSwag ↑平均 acc ↑Mamba-130M10.5644.335.344.7Pythia-160M29.6433.030.240.6Mamba-370M8.2855.646.550.0Pythia-410M9.9551.440.648.2Mamba-1.4B6.8064.959.159.7Pythia-1.4B7.5161.752.155.2Mamba 在每個(gè)尺寸檔全面勝出1.4B 平均準(zhǔn)確率 59.7 甚至超過同 tokenizer 的 Pythia-2.8B59.1。Mamba-3B 在常識(shí)推理上比 Pythia-3B 高 4 分匹敵 2 倍尺寸 Transformer——這是第一個(gè)匹配 LLaMa 配方 Transformer 的無注意力模型。圖 4Pile 2K 上下文縮放律——首個(gè)匹配 Transformer 的無注意力模型。效率訓(xùn)練與推理雙贏融合掃描比樸素實(shí)現(xiàn)快40 倍推理時(shí)作為遞歸模型每步恒定時(shí)間、無需 KV 緩存吞吐量達(dá)同尺寸 Transformer 的5 倍。圖 5訓(xùn)練與推理效率基準(zhǔn)。DNA 與音頻百萬級(jí)長(zhǎng)上下文基因組上固定模型大小時(shí)性能隨上下文單調(diào)提升至 $2^{20}$1M基線持平甚至下降大猿物種分類上下文 1M中 Mamba 準(zhǔn)確率領(lǐng)先。音頻上6.1M 參數(shù)的 Mamba 在 SC09 上 FID 0.94對(duì)比 SaShiMi 的 1.99降幅超過一半24.3M 版本 FID 0.67、mIS 144.9超越 WaveGAN/DiffWave 等 GAN 與擴(kuò)散基線附錄 G 詳表。圖 8長(zhǎng)上下文 DNA 分類——選擇機(jī)制過濾無關(guān)上下文能力的直接驗(yàn)證。消融$\Delta$ 最重要$B,C$ 協(xié)同架構(gòu)內(nèi)部層PPL ↓H3S4real10.34H3S68.95MambaS4real10.56MambaS68.69選擇性 $\Delta$選擇性 $B$選擇性 $C$PPL ↓???10.93???9.81???8.71$\Delta$ 是最重要的選擇性參數(shù)門控連接狀態(tài)維數(shù) $N$ 從 1 增到 16 困惑度下降超 1.0、僅增 1% 參數(shù)但只有 $B,C$ 也選擇性時(shí)才有效附錄 E。圖 6歸納頭外推曲線——選擇性機(jī)制帶來可無限外推的內(nèi)容感知能力。圖 7融合掃描內(nèi)核的訓(xùn)練效率。結(jié)果對(duì)比總結(jié)圖 11結(jié)果對(duì)比總結(jié)——質(zhì)量匹敵兩倍尺寸 Transformer效率線性擴(kuò)展長(zhǎng)上下文能力為三模態(tài)通用骨干奠定基礎(chǔ)。關(guān)鍵發(fā)現(xiàn)選擇機(jī)制是性能分水嶺Selective Copying 從 S4 的 18.3% 躍升到 S6 的 99.8%歸納頭可外推 4000 倍$2^8\to2^{20}$LTI 模型完全做不到。首次匹配 Transformer125M–1.3B 縮放律上 Mamba 是第一個(gè)無注意力模型匹配 LLaMa 配方且序列越長(zhǎng)優(yōu)勢(shì)越明顯。匹敵兩倍尺寸Mamba-1.4B 零樣本平均 59.7 超 Pythia-2.8B 的 59.1Mamba-3B 常識(shí)推理比 Pythia-3B 高 4 分。效率數(shù)量級(jí)提升推理吞吐 5 倍于 Transformer融合掃描比樸素實(shí)現(xiàn)快 40 倍訓(xùn)練內(nèi)存與 FlashAttention 同級(jí)。三模態(tài) SOTA音頻 FID 從 1.99 降至 0.94減半以上基因組與音頻性能隨上下文單調(diào)提升至 1M。$\Delta$ 門控理論$N1$ 時(shí)選擇性 SSM 精確退化為門控 RNN統(tǒng)一了 SSM 離散化與 RNN 門控兩套理論。局限性規(guī)模有限實(shí)證僅到約 3B 參數(shù)低于 Llama/RWKV/RetNet 的 7B更大規(guī)模下是否保持優(yōu)勢(shì)未知規(guī)模化需額外工程。連續(xù)-離散免費(fèi)午餐選擇機(jī)制犧牲了 LTI 在連續(xù)信號(hào)音頻/視頻上的強(qiáng)歸納偏置音頻實(shí)驗(yàn)需切回復(fù)數(shù)參數(shù)化附錄 G。生態(tài)欠賬微調(diào)、prompting、ICL、指令微調(diào)、RLHF、量化等 Transformer 生態(tài)的成熟適配機(jī)制尚未在 Mamba 上建立。硬件依賴性能依賴定制融合內(nèi)核selective scan新硬件需重新工程化。常見問題FAQMamba 和 Transformer 的核心區(qū)別是什么Transformer 用自注意力在窗口內(nèi)稠密路由信息訓(xùn)練二次、推理需 KV 緩存Mamba 用選擇性狀態(tài)空間遞推訓(xùn)練線性、推理恒定時(shí)間靠輸入相關(guān)的參數(shù)決定記憶與遺忘首次在不犧牲質(zhì)量的前提下實(shí)現(xiàn)線性復(fù)雜度。為什么說 Mamba 的訓(xùn)練是線性時(shí)間的Mamba 的時(shí)變遞推雖不能卷積化但硬件感知的并行前綴掃描把 $O(BLDN)$ 的 FLOPs 并行化且內(nèi)存帶寬受限的操作通過內(nèi)核融合SRAM 內(nèi)完成離散化與遞推保持高效因此訓(xùn)練隨序列長(zhǎng)度線性擴(kuò)展。選擇機(jī)制到底選擇什么選擇的對(duì)象是信息的流入與流出$\Delta$ 決定當(dāng)前輸入被聚焦還是被忽略大 $\Delta$ 重置狀態(tài)、小 $\Delta$ 保持歷史$B,C$ 分別控制輸入進(jìn)狀態(tài)、狀態(tài)出輸出的細(xì)粒度門控——本質(zhì)上讓固定容量的狀態(tài)按內(nèi)容做最優(yōu)壓縮。Mamba 在哪些任務(wù)上不如 Transformer在連續(xù)信號(hào)模態(tài)如音頻、視頻上Mamba 的時(shí)變選擇機(jī)制弱于 LTI SSM 的強(qiáng)歸納偏置需要復(fù)數(shù)參數(shù)化彌補(bǔ)此外大規(guī)模7B驗(yàn)證、生態(tài)工具鏈微調(diào)/量化/RLHF也落后于 Transformer。Mamba-2 與 Mamba 是什么關(guān)系Mamba-2Dao Gu, ICML 2024通過狀態(tài)空間對(duì)偶SSD統(tǒng)一了 SSD 與注意力把選擇性掃描進(jìn)一步映射到類注意力結(jié)構(gòu)硬件效率再提升約 8 倍同時(shí)保持了 Mamba 的選擇性核心。為什么歸納頭任務(wù)如此重要?dú)w納頭induction heads被廣泛認(rèn)為是 LLM 上下文學(xué)習(xí)能力的關(guān)鍵機(jī)制。Mamba 在長(zhǎng)度 256 訓(xùn)練后外推到 1M 仍保持高準(zhǔn)確率直接證明選擇性機(jī)制具備可無限外推的內(nèi)容感知能力而所有 LTI 對(duì)比方法最多外推 2 倍。參考鏈接本文Gu Dao, Mamba: Linear-Time Sequence Modeling with Selective State Spaces (arXiv:2312.00752), COLM 2024開源代碼github.com/state-spaces/mambaS4Gu, Goel, Ré, Efficiently Modeling Long Sequences with Structured State Spaces, ICLR 2022HiPPOGu et al., HiPPO: Recurrent Memory with Optimal Polynomial Projections, NeurIPS 2020H3Dao et al., Hungry Hungry Hippos, ICLR 2023FlashAttentionDao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, NeurIPS 2022RWKVPeng et al., RWKV: Reinventing RNNs for the Transformer Era, Findings of EMNLP 2023給大家推薦一款自用寫文獻(xiàn)綜述、無虛構(gòu)文獻(xiàn)的 AI復(fù)旦大學(xué) FudanNLP 團(tuán)隊(duì)自研 切問學(xué)術(shù)官網(wǎng)qiewenpaper.com覆蓋3.6 億篇可溯源真實(shí)中英文文獻(xiàn)能自動(dòng)整合文獻(xiàn)觀點(diǎn)生成規(guī)范綜述還能挖掘研究創(chuàng)新點(diǎn)、復(fù)現(xiàn)實(shí)驗(yàn)配合視頻教學(xué)新手快速上手文獻(xiàn)綜述寫作后記博客的關(guān)鍵詞集中在編程、算法、機(jī)器人、人工智能、數(shù)學(xué)等等持續(xù)高質(zhì)量輸出中。討論QQ群白拾的小屋 (750365700)?B站賬號(hào)白拾的物理AI組會(huì)活躍于知識(shí)區(qū)和動(dòng)畫區(qū)?GitHub主頁YhbCode000工程文件