3 ポイント 投稿者 GN⁺ 3 시간 전 | 1件のコメント | WhatsAppで共有
  • ソフトマックスアテンションから出発し、固定サイズ状態を使う線形アテンション、誤差だけを記録する DeltaNet、状態全体を減衰させる Gated DeltaNet、チャネルごとに減衰させる Kimi Delta Attention(KDA)までを段階的に導出する
  • 基本的な線形アテンションは、過去の key-value 外積の和を状態 (S_t) に保存してシーケンス長に対して線形に動作するが、新しい値を代入せず既存の関連に加える 加算的書き込み干渉 が発生する
  • DeltaNet は、現在の key から予測した値と目標 value の差に (\beta_t) を掛けて記録し、即時再構成条件・オンライン勾配降下・ランク 1 状態更新という 3 つの解釈が同じ式に帰着する
  • Gated DeltaNet はスカラー (\alpha_t) で状態全体を先に減衰させ、KDA はこれを対角行列 (D_t=\operatorname{Diag}(\alpha_t)) に拡張して、key チャネルごとに異なる比率で情報を保持または削除する
  • 同じ KDA 漸化式を、デコード用の 融合再帰 Triton カーネル と学習・長いプレフィル用のチャンク方式で実行し、チャンク方式はトークン内部依存を三角求解で復元して行列積として再構成する

表記法と展開順

  • bra-ket 表記では、(\lvert q\rangle) は列ベクトル、(\langle k\rvert) は行ベクトル、(\langle k\vert q\rangle) はスカラー、(\lvert v\rangle\langle k\rvert) は行列である
  • 1 つの因果的アテンションヘッドと実数ベクトルを用い、DeltaNet の key は正規化 されており、状態は key 空間から value 空間への写像だと仮定する
  • 展開順はソフトマックスアテンション → 線形アテンション → DeltaNetGated DeltaNetKDA で、最後に再帰型およびチャンク型の Triton 実装へつなげる
  • DeltaNet 系の変種のうち 2 つは最新の Qwen と Kimi モデル系 で使われている

二次複雑度アテンションから線形状態へ

  • 一般的な因果的ソフトマックスアテンションは、key と query の類似度を計算し、すべての過去 key に対するスコアを分布として正規化したうえで、value ベクトルの重み付き和を出力する
  • 長さ (T) のシーケンスには (T^2) 個の key-query ペアがある
    • 自己回帰推論では key と value をキャッシュできるが、キャッシュサイズはシーケンスとともに増える
    • 新しい query も過去全体を確認する必要がある
  • ソフトマックスの分母は現在の query とすべての以前の key に共同で依存するため、計算順序を単純に並べ替えるのは難しい
  • ソフトマックスを取り除くと、出力を過去 key-value 外積の和としてまとめられる
    • (S_t=\sum_{i\le t}\lvert v_i\rangle\langle k_i\rvert)
    • (S_t=S_{t-1}+\lvert v_t\rangle\langle k_t\rvert)
    • (\lvert o_t\rangle=S_t\lvert q_t\rangle)
  • 重要な恒等式は ((\lvert v\rangle\langle k\rvert)\lvert q\rangle=\langle k\vert q\rangle\lvert v\rangle) であり、すべての過去 key と value の代わりに、集約された外積を 固定サイズの (d_v\times d_k) 状態 に保存する
  • トークンを 1 回走査するのでシーケンス長に対して線形に動作するが、その代償としてソフトマックスの正規化と選択性を失う
    • より洗練された線形アテンションは特徴マップと正規化項を使う

線形アテンションの加算的書き込み問題

  • 正規化された現在の key に (\lvert v_t\rangle\langle k_t\rvert) を記録した直後、同じ key で読むと (S_t\lvert k_t\rangle=S_{t-1}\lvert k_t\rangle+\lvert v_t\rangle) になる
  • 新しい書き込みは、メモリが (v_t) を返すように代入するのではなく、既存の返り値に (v_t) を += 方式で加える
  • 以前の状態がすでに正しい値を返しているなら同じ value は 2 倍になり、key 同士は互いに直交していないので、各書き込みが既存の書き込みと干渉しうる
  • 線形アテンションは圧縮された連想メモリを提供するが、必要な = に近い更新ではなく加算更新を行う

DeltaNet: 値の代わりに予測誤差を書き込む

  • DeltaNet は、新しい key に対する既存予測 (\widehat v_t=S_{t-1}k_t) を先に読み、value 全体ではなく差分だけを記録する
    • (e_t=\beta_t(v_t-S_{t-1}k_t))
    • (S_t=S_{t-1}+e_tk_t^\mathsf T)
    • 学習された書き込み強度 (\beta_t) は ([0,1]) の範囲にある
  • 同じ key で直ちに再度読むと ((1-\beta_t)S_{t-1}k_t+\beta_tv_t) になる
    • (\beta_t=1) なら正確に (v_t) を返す
    • より小さい値では既存予測を目標方向へ一部だけ移動させる
  • 更新は key 空間で局所的 である
    • 現在の key と直交する query 方向では外積更新が 0 なので応答は変わらない
    • 現在の key 方向の関連だけを選択的に置き換える
  • 再構成損失から導く

    • 状態 (S) を線形写像とみなし、現在の key-value ペアの損失を (\frac12\lVert Sk_t-v_t\rVert_2^2) と置くと、勾配は ((Sk_t-v_t)k_t^\mathsf T) になる
    • (S_{t-1}) から大きさ (\beta_t) の勾配降下を 1 ステップ行うと、DeltaNet の更新式と正確に一致する
    • 同じ更新は 3 通りに解釈できる
      • メモリ演算では (\beta_t) は 既存関連の置換強度 である
      • オンライン学習では (\beta_t) は学習率である
      • 線形代数では予測誤差と key のランク 1 外積である
  • 構造化された状態遷移

    • 更新を展開すると (S_t=S_{t-1}(I-\beta_tk_tk_t^\mathsf T)+\beta_tv_tk_t^\mathsf T) になる
    • 単位 key に対して (I-\beta_tk_tk_t^\mathsf T) は、現在の key 方向で固有値 (1-\beta_t)、すべての直交方向で固有値 1 を持つ
    • 既存の key 方向の関連を先に取り除いてから新しい関連を加えるが、状態全体の寿命管理 はまだ解決していない

Gated DeltaNet: 状態全体を先に忘れる

  • 1 つの行列に過去全体を圧縮すると、すでに状態に統合された個々のトークンだけを選んでスキップすることはできない
  • DeltaNet は現在の key 周辺を補正するが、他の方向の古い情報は残り、将来の読み出しに引き続き寄与しうる
  • Gated DeltaNet は学習されたスカラー保持ゲート (\alpha_t\in[0,1]) を適用する
    1. (\widetilde S_t=\alpha_tS_{t-1}) として忘却する
    2. (\widehat v_t=\widetilde S_tk_t) として予測する
    3. (e_t=\beta_t(v_t-\widehat v_t)) として補正する
    4. (S_t=\widetilde S_t+e_tk_t^\mathsf T) として記録する
  • 忘却 → 予測 → 補正 → 書き込み の順序が重要である
    • 減衰前に予測すると、誤差を計算したメモリと実際に更新するメモリが異なってしまう
  • デルタ則は目標 key に対する置換を、スカラーゲートはグローバル削除を担い、異なる問題を解決する
  • ただし 1 つの (\alpha_t) が行列全体に適用されるため、すべての key チャネルを同じ比率で保持または忘却しなければならない

Kimi Delta Attention: チャネルごとの減衰

  • Kimi Delta Attention はスカラー (\alpha_t) を (d_k) 次元ベクトルに置き換え、(D_t=\operatorname{Diag}(\alpha_t)) を構成する
  • 状態は key 空間から value 空間への写像なので、key チャネルは (S) の列に対応し、右からの積 (S_{t-1}D_t) が各列に異なる保持率を適用する
  • KDA は次の順序で動作する
    1. (\widetilde S_t=S_{t-1}D_t) として key チャネルごとに減衰
    2. (\widehat v_t=\widetilde S_tk_t) として予測
    3. (e_t=\beta_t(v_t-\widehat v_t)) として補正
    4. (S_t=\widetilde S_t+e_tk_t^\mathsf T) として記録
    5. (o_t=S_t(d_k^{-1/2}q_t)) として読み出す
  • Gated DeltaNet から KDA への概念的変化は、(\alpha_t) を (D_t) に昇格させたことだけだが、1 つのチャネルを消しながら別のチャネルは保持できる
  • 対角-低ランク遷移

    • KDA を展開すると (S_t=S_{t-1}A_t+\beta_tv_tk_t^\mathsf T) となり、(A_t=D_t(I-\beta_tk_tk_t^\mathsf T)) である
    • (A_t=D_t-b_ta_t^\mathsf T)、(b_t=D_tk_t)、(a_t^\mathsf T=\beta_tk_t^\mathsf T) と書けるため、対角-低ランク(DPLR) 遷移になる
    • DPLR は key 空間で作用する (d_k\times d_k) 遷移を意味し、メモリ状態自体は依然として (d_v\times d_k) 行列である
    • 系列ごとに次の機能が追加される
      • 線形アテンション: 固定サイズの再帰メモリ
      • DeltaNet: 目標方向の選択的置換
      • Gated DeltaNet: 状態全体の減衰
      • KDA: key チャネルごとの減衰
    • 実装では通常 (g_t=\log\alpha_t\le0) を保存してから、(\exp(g_t)) で保持率を求める
    • 転置された (d_k\times d_v) レイアウトの 5 段階参照実装は naive_recurrent_kda で確認できる

デコード用の融合再帰 Triton カーネル

  • KDA には 2 つの主要な実行方式がある
    • 融合再帰方式: デコード、短いシーケンス、状態保持型サービングに適している
    • チャンク方式: 学習と長いプレフィルに適している
  • fused_recurrent_kda_fwd は、シーケンス・value head・幅 32 の value タイルごとに 1 つの Triton プログラムを実行する
    • BK は一般的なサポート構成では key 次元を覆う
    • 各プログラムは転置状態の [BK, BV] タイルを所有し、トークンを順に走査する
    • 異なる value タイル、head、シーケンスは独立に実行される
  • カーネルは状態減衰、key に対する予測の縮約、residual の計算、外積書き込み、query 読み出しの縮約を漸化式どおりに実行する
  • 一度に新しいトークンが 1 つだけ入るデコードには向いているが、ベクトル演算を Tensor Core に効率的な大きい行列積へ変換できないため、学習と長いプレフィルには不利 である

Chunkwise KDA: 漸化式を行列積へ並べ替える

  • Chunkwise KDA は (C) 個のトークンをまとめて処理しつつ、トークンごとの再帰方式と完全に同じ状態と出力を生成しなければならない
  • 各チャンクは 2 つの結果を計算する
    • 入力状態 (S_c) からチャンク全体を処理した後の (S_{c+1})
    • チャンク内部のすべてのトークンの因果的出力
  • 各トークンのデルタ誤差が同じチャンク内の以前の書き込みに依存することが、核心的な難しさである
  • 累積減衰と一時誤差

    • トークン (i) の対角減衰を (D_i)、チャンク境界からトークン (i) までの累積減衰を (D_{0:i}=D_0D_1\cdots D_i) と置く
    • トークン (j) の書き込みがトークン (i) まで伝播するときは (D_{j+1:i}) が適用され、対角行列なので減衰行列同士は可換である
    • まずチャンク内部の他の書き込みを無視した一時誤差を並列計算する
      • (\bar e_i=\beta_i(v_i-S_cD_{0:i}k_i))
    • 最初のトークンを除く一時誤差は、以前のチャンク内部書き込みの影響を欠いているため、そのままでは使えない
  • 因果依存の復元

    • 以前のトークン (j) が現在のトークン (i) の誤差に与える係数を (\rho_{ij}=\beta_i k_j^\mathsf TD_{j+1:i}k_i) と定義する
    • 実際の誤差は (e_i=\bar e_i-\sum_{j<i}\rho_{ij}e_j) の形で逐次依存を持つ
    • (\rho_{ij}) を厳密な下三角行列 (R_c) に入れると、積み上げた誤差行列は (E_c=\bar E_c(A_c^{kk})^\mathsf T)、(A_c^{kk}=(I+R_c)^{-1}) として計算できる
    • 一般的な密行列の逆行列は不要である
      • (I+R_c) は対角成分が 1 の三角行列である
      • 各 value チャネルに対して 因果的三角求解 を行えばよい
  • チャンク終了状態の計算

    • 入ってくる状態はチャンクのすべての減衰を通過し、各チャンク内部書き込みは自分より後ろにある減衰だけを通過する
    • チャンク終端まで減衰された key を (K_c^{\mathrm{end}}) に行として積むと、状態を次の行列積として整理できる
      • (S_{c+1}=S_cD_{0:C-1}+E_cK_c^{\mathrm{end}})
    • 複数のランク 1 外積書き込みを 1 つの行列積にまとめ、チャンク全体の状態を一度に前進 させる
  • チャンク内部の全出力計算

    • KDA は現在のトークンを書いた後に読むため、トークン (i) の出力には自分自身の書き込みも含まれる
    • 以前の書き込み (j) が query (i) に与える係数を (\chi_{ij}=s,k_j^\mathsf TD_{j+1:i}q_i)、(j\le i) と定義する
    • 係数を下三角の読み出し行列 (A_c^{qk}) に配置する
      • 上三角の 0 は未来トークンからの寄与を遮断する
      • 対角成分は現在のトークンが自身の書き込み後に読む動作を反映する
    • 境界から各 query まで減衰されたベクトルを (Q_c^{\mathrm{boundary}}) に積むと、全出力は次のようになる
      • (O_c=sS_cQ_c^{\mathrm{boundary}}+E_c(A_c^{qk})^\mathsf T)
    • 最初の行列積は減衰したチャンク入力状態を読み出し、2 つ目の行列積はチャンク内部の因果的書き込み寄与を加える

Chunkwise Triton パイプライン

  • チャンク実装は 1 つの巨大カーネルではなく、複数カーネル呼び出しからなる パイプライン である
  • まずチャンク内部の累積ログ減衰を計算する
    • 2 つの prefix sum の差によって、保持ベクトルを長く掛け続けずに (D_{j+1:i}) を表現する
  • 続いて因果的な (A^{qk}) と (A^{kk}) の相互作用行列を作り、(A^{kk}) によってチャンクの補正済み書き込みに対する WY 形式を構成する
  • 状態カーネルは、チャンク間で唯一の逐次走査を行う
    • 各チャンクへ入る状態を生成する
    • チャンクのデルタ誤差を解消する
  • 進入状態が計算された後は、出力カーネルが異なるチャンクとタイルのトークンを並列処理できる
  • 実装ではまず 16 トークンの対角相互作用ブロック を計算し、その後に融合された非対角および三角求解カーネルを実行する
  • chunk_kda_fwd が各段階を調整し、主要エントリポイントは chunk_kda_fwd_intrachunk_gated_delta_rule_fwd_hchunk_gla_fwd_o_gk である
    • コード中の v_new は解消済み誤差である
    • h はチャンク進入状態である
    • kg はチャンク終端まで減衰された key である
  • 再帰方式とチャンク方式は別々のアテンションではなく、同じ KDA 漸化式の 2 つの実行スケジュール である
    • 再帰方式は低遅延デコードのための直列ベクトル演算である
    • チャンク方式は Tensor Core 中心の学習とプレフィルのための行列演算である

1件のコメント

 
GN⁺ 3 시간 전
Hacker News の意見
  • この15年間、機械学習には統一された数学表記が必要だったし、おそらく今後も必要だったはずです。昔は世界各地の研究者の論文ごとに奇想天外な表記が登場していて、さらにひどかったです。
    論文ごとに表記が変わると、理解に摩擦が生じます。少なくともこの記事は最初から表記を明示的に説明していますが、そうする論文は少ないほうです。最初は表記切り替え機能にも気づきませんでしたが、とても便利です。

    • 一文字の記号や明示的なデータ型の代わりに ∣q⟩ のような文字を使う伝統的な数学表記を好む理由が理解できません。簡潔という利点はあるでしょうが、数式を擬似コードや Python のような実際のプログラミング言語で書けば、はるかに理解しやすいと思います。
    • この記事は表記の一側面を説明しているだけで、使っている変数の定義は示していません。kqS が何かは、機械学習を学んでいれば分かるか推測できますが、関連する背景知識がないと記事の大半が不透明になります。
    • 以前は私もそう思っていましたが、コードより数式を眺めている時間のほうがはるかに長いので、記号の意味を知ってしまえば簡潔な表記のほうがずっと読みやすいです。文字で書くと、難しいことで有名な命名まで避けられます。
  • 「自分でも思いつけたかもしれない……」と言いますが、存在しなかった何かを作ったり組み合わせたりするのはものすごく難しいことです。
    誰かが難しい作業をやり遂げて公開すると、すぐに「たいして難しくないね」「自分にもできた」といった反応が出て、すべてが単純に見え始めます。開発中に新しいものを発明したと思ったら、実はすでに1970年代に作られ広く使われていたものだと後で気づくこともよくあります。ただ自分の経路と交わらず、存在を知らなかっただけです。

  • 私にとってはブラケット記法が、すべてをシンプルで直感的にしてくれます。ベクトル表記ではどちらが横でどちらが縦なのか混乱し、塊だけを追って集中力を失いがちでしたが、ブラケットだと全体が非常に直感的でした。
    見逃している良い記事がたくさんありそうなので、他の記事もこの記法に変換してみようと思います。ちなみに物理学の博士号を持っていて、軽いディスレクシアがあります。

  • 「外積は行列で、内積は数値だ。過去のすべてのキーと値を保存する代わりに、固定サイズの状態 S_t に外積の和を保存する」といった文体を見ると、LLM が書いた記事だと確信してしまいます。

    • たぶん、バズワード入りのタイトルを依頼するところから始めたのでしょう。
    • Claude にダッシュ()を使わないようプロンプトすると、こういう結果になります。
  • 可視化されたチュートリアルもあります: https://snowchord.com/blog/linear-attention-visualized/

  • こういう記事やタイトルを見るたびに、自分よりはるかに賢い大勢の人たちに深い感謝と謙虚な気持ちを覚えます。高校や学部ではとても賢い人として通っていましたし、平均よりは利口ですが、私を青二才に見せる人も何百万人も確実にいます。
    ここで賢いというのは、巨大で複雑な概念やシステムを頭の中に保持して推論する能力を指していて、数学者にとって特に重要な才能に見えます。

    • AI ツールが作業をますます加速させるとしても、新しいアイデアの大半の源泉は今後も人間だと思います。
      友人と酒を飲みながらやってみた思考実験は、子どもたちを画面やアルゴリズムが供給する大衆的コンテンツから隔離し、最先端モデルを訓練するようにメディアや資料の質を厳格に管理する、学習に適した環境で育てるというものでした。子どものための修道院のようにして、数学、工学、コンピュータサイエンス、ディープラーニングなどを通じて現実に関する最新知識を教えるやり方です。
      結局、先端的な AI ツールを活用して知識の境界を広げるには、依然として非常に賢く、思考が大きく汚染されていない人が必要です。AI が人間を完全に置き換えるという考えは、方向を誤っています。
  • 参考までに、ブラケット記法という名前は実際に括弧(bracket)に由来します。
    https://en.wikipedia.org/wiki/Bra-ket_notation

  • 最初はためらいましたが、ケット記法のおかげで演算がずっと明確になり、気に入りました。ただ、二次アテンションの d_k のように、一部の変数について簡単なおさらいもあればよかったと思います。

  • 最初はこの解法を思いつけなかったことに落ち込みましたが、JavaScript で二分探索を自分で書くのにも苦労することに気づいて、すぐに気が楽になりました。Kimi Delta Attention を自分が思いついた可能性はまったくありません。

    • 線形代数のコードは、意外と書きやすい面があります。一般的なコンピュータサイエンスのコードのように再帰が複雑に絡み合うことはなく、すべての変数の間に数学的関係があり、よくある数学概念はすでによく実装されたライブラリを利用できます。
      ループも二、三段階以上深くなることはめったになく、それより複雑なら、どうせライブラリに任せたほうがよいです。