AI News HubLIVE
站內改寫2 分鐘閱讀

Kimi DeltaNet 線性注意力系列指南

本文詳細推導了DeltaNet系列線性注意力機制,包括softmax注意力、線性注意力、DeltaNet、門控DeltaNet和Kimi Delta注意力(KDA)。解釋了每個變體如何通過解決記憶、選擇性和遺忘問題來改進前一個。文章使用bra-ket表示法,並強調模型背後的數學直覺。

來源Hacker News AI作者: kkm

在這篇文章中,我們將逐步推導Kimi Delta Attention(KDA)及其前身,從標準的softmax注意力開始,最終到達KDA這一高效的線性注意力變體。文章採用bra-ket表示法,因為這種表示法能清晰地展示推導中的形狀。在bra-ket模式中,∣q⟩表示列向量,⟨k∣表示行向量,⟨k∣q⟩是一個數,而∣v⟩⟨k∣是一個矩陣。我們假設只有一個因果注意力頭,向量為實數,DeltaNet的鍵已歸一化,狀態從鍵空間映射到值空間。

首先,我們從標準的softmax注意力開始。對於第t個token的查詢,注意力權重計算為ati = exp(s⟨ki∣qt⟩) / ∑_{j≤t} exp(s⟨kj∣qt⟩),其中s = d_k^{-1/2}。輸出是所有值的加權和。該計算需要計算所有T^2個鍵-查詢對,在自迴歸推理中,緩存會隨着序列增長而增加,每個新查詢仍需檢查整個歷史。

移除softmax後,我們得到線性注意力。將尺度s吸收到查詢中,輸出變為∣ot⟩ = ∑_{i≤t} ⟨ki∣qt⟩∣vi⟩ = (∑_{i≤t} ∣vi⟩⟨ki∣)∣qt⟩。定義狀態St = ∑_{i≤t} ∣vi⟩⟨ki∣,則注意力變為遞歸寫入後讀取:St = S_{t-1} + ∣vt⟩⟨kt∣,∣ot⟩ = St∣qt⟩。這實現了O(T)的複雜度,但加法寫入存在信息重疊問題:寫入並不使記憶返回∣vt⟩,而是將∣vt⟩加到已有值上,當舊狀態已產生正確值時,加法寫入會使新狀態產生兩倍值。

為了解決這個問題,DeltaNet引入delta規則修正。在寫入當前鍵值對之前,先從舊狀態中預測v̂_t = S_{t-1}∣kt⟩,然後計算誤差∣et⟩ = β_t (∣vt⟩ - v̂_t),最後寫入誤差:St = S_{t-1} + ∣et⟩⟨kt∣。當β_t=1時,立即讀取同一鍵能得到精確值∣vt⟩。該修正具有局部性:對於與當前鍵正交的查詢,寫入不影響其輸出。DeltaNet也可以從在線學習角度理解:它是對重構損失L_t(S) = 1/2 ||S∣kt⟩ - ∣vt⟩||^2的一步梯度下降,步長為β_t。DeltaNet的更新等價於St = S_{t-1}(I - β_t∣kt⟩⟨kt∣) + β_t∣vt⟩⟨kt∣,其中I - β_t∣kt⟩⟨kt∣在當前鍵方向上的特徵值為1-β_t,正交方向上為1,從而在添加新關聯前移除舊關聯。

然而,DeltaNet無法清除歷史中的陳舊信息。為了管理記憶生命週期,Gated DeltaNet添加了一個遺忘門α_t,在每次更新前將狀態縮放:S̃_t = α_t S_{t-1},然後應用delta規則。這允許模型遺忘過時的關聯。Gated DeltaNet的完整更新包括遺忘、預測、修正和寫入四步。

最後,Kimi Delta Attention(KDA)結合了門控和delta規則。KDA的更新方程如下:首先遺忘:S̃_t = S_{t-1} Diag(α_t),其中α_t是向量,控制每個維度的遺忘;然後預測:v̂_t = S̃_t∣kt⟩;接着計算誤差:∣et⟩ = β_t (∣vt⟩ - v̂_t);最後寫入:St = S̃_t + ∣et⟩⟨kt∣;輸出:∣ot⟩ = St (d_k^{-1/2}∣qt⟩)。KDA的推導源自對記憶操作的直觀要求:寫操作應能覆蓋舊值,遺忘機制應能控制歷史的持久性。

在推導完KDA後,文章還討論瞭如何通過Triton實現循環和分塊計算,以提高效率。總之,DeltaNet家族通過逐步改進,從簡單的線性注意力發展到具有選擇性寫入和遺忘能力的複雜模型,為高效序列建模提供了理論基礎。