Kimi DeltaNet 線形注意機構ファミリーガイド
この記事では、ソフトマックス注意、線形注意、DeltaNet、ゲート付きDeltaNet、Kimi Delta Attention(KDA)を含むDeltaNetファミリーの線形注意機構の詳細な導出を提供します。各バリアントがメモリ、選択性、忘却の問題にどのように対処するかを説明します。ブラケット記法を使用し、モデルの背後にある数学的直感を強調します。
この記事では、Kimi Delta Attention(KDA)とその前身を段階的に導出します。まず、標準的なソフトマックス注意から始めます。これは各クエリとすべてのキーとの類似度を計算し、ソフトマックスで正規化して値の重み付き和を出力します。しかし、ソフトマックスの正規化により計算を再帰形式に簡略化できません。
次に、ソフトマックスを取り除き、線形注意を得ます。線形注意は注意計算を外積和に簡略化します。状態St = Σ |vi><ki| で表され、出力はSt |qt>です。これによりO(T)の複雑さを実現しますが、加法的な書き込みにより情報の重複が問題となります。具体的には、書き込みを行った後、同じキーで読み取ると、古い状態の出力に新しい値が加算されるため、記憶が正しい値を返しません。
この問題を解決するために、DeltaNetはデルタルール補正を導入します。現在のキーと値のペアを書き込む前に、古い状態から予測値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に対する1ステップの勾配降下です。DeltaNetの更新はSt = S_{t-1}(I - β_t|kt><kt|) + β_t|vt><kt|とも表現でき、現在のキー方向の古い関連を除去してから新しい関連を追加します。
しかし、DeltaNetは履歴内の古い情報を消去できません。そこで、ゲート付きDeltaNetは忘却ゲートα_tを追加し、各更新前に状態をスケーリングします:S̃_t = α_t S_{t-1}、その後デルタルールを適用します。ゲート機構により、モデルは時代遅れの関連付けを選択的に忘却できます。
最後に、Kimi Delta Attention(KDA)はこれら二つのアイデアを組み合わせます。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ファミリーは段階的な改善を通じて、単純な線形注意から選択的書き込みと忘却能力を持つ複雑なモデルへと発展し、効率的な系列モデリングのための理論的基盤を提供します。