AI・機械学習
PSSA: Rustでゼロから書かれた非Transformer言語モデル
PSSA: A non-transformer language model written from scratch in Rust (github.com)
要約
PSSAは、Transformerアーキテクチャに依存しない、Rustでゼロから構築された新しい言語モデルです。逐次的な状態空間層とエピソード記憶バンクを利用し、実行中に自身の重みを書き換える「プラスチック」な性質を持ちます。Transformerと比較して、同等のパラメータ数でより高速な学習と生成速度を実現し、特にCPUでのパフォーマンスが12倍優れていると報告されています。このモデルは、Transformerの計算コストの課題を解決する可能性を示唆しています。
全文翻訳
PSSA:プラスチック状態空間アーキテクチャ
PSSAは、Transformerではない小さな言語モデルです。再帰的な状態空間層を通じてテキストを一度に1トークンずつ読み込み、エピソード記憶のバンクにルックアップを行い、実行中に自身の重みの一部を書き換えます。これはRustでゼロから書かれており、PyTorch、TensorFlow、またはその他のMLフレームワークは一切使用していません。同等のパラメータ数で同じコーパスを使用した場合、Transformerよりも速く学習し、同じCPUで約12倍速くテキストを生成します。
なぜRustなのか、そしてそれがポイントではない理由
速度のためでもなく、言語がアーキテクチャをより良くするためでもありません。PSSAは、トークンごとの重み更新、フォワードパス中に書き込まれるメモリバンク、およびすべてのバッチカーネルが微分可能なスカラー参照パスを必要としました。これを自動微分フレームワーク内で表現することは、あらゆるステップでフレームワークと戦うことを意味したため、線形代数は直接記述されました。これにより、プラスチック部分は単純になり、勾配は約3e-8までチェック可能になりました。アーキテクチャが主張するところです。実装言語は詳細であり、Pythonポートも歓迎します。
Transformerとの違い
Transformerはコンテキスト内のトークンのペアすべてにスコアリングするため、ステップあたりのコストはシーケンス長の二乗に比例して増加し、すべてのコンテキストがステップごとに再読み込みされます。PSSAは、シーケンス全体で固定サイズの状態を1つの左から右へのパスで運び、コンテキストを再読み込みする代わりにメモリバンクにルックアップを行うため、コストは長さに比例して増加します。
モデル
各トークンは1つのPSSAレイヤーを通過します:選択的状態空間再帰、エピソード記憶バンクからの双曲線空間での制限付き読み取り、その読み取りが残差ストリームにどれだけ到達するかを決定する学習ゲート、およびSiLU MLP。デフォルトはd_m = 256チャネル、チャネルあたりd_s = 16状態、ランク16アダプターです。
再帰
レイヤー正規化されたトークン埋め込みをxとする。トークン自体から3つの射影が読み取られ、これが再帰を選択的にするのではなく固定するものになります:
delta = softplus(W_delta x) チャネルごとのステップサイズ、delta in R^d_m
B = W_B x 入力マップ、B in R^d_s
C = W_C x 出力マップ、C in R^d_s
遷移は対角的で、(チャネル、状態)ペアごとに1つのレートがあり、構成によって負に保たれるため、再帰が爆発することはありません:
A = -softplus(A_raw) A in R^(d_m x d_s)
この連続システムをステップdeltaで離散化すると、トークンごとの更新が得られます。hはトークン間およびトレーニング中のチャンク境界を越えて伝達されます:
Abar_ij = exp(delta_i * A_ij)
Bbar_ij = delta_i * B_j
h_ij <- Abar_ij * h_ij + Bbar_ij * x_i
y_i = sum_j C_j * h_ij
A_rawは、HiPPO初期化の精神で、各チャネルの16のレートが1.5から200トークンの対数間隔のタイムスケールに配置されるように初期化されます。したがって、単一のチャネルは、最後の2つのトークンと最後の200個のトークンを同時に保持することから始まり、トレーニングはこれらを最初から発見するのではなく、それらのホライゾンを移動させます。レイヤーのこの半分は選択的対角SSMであり、新規性を主張しません。これはS4およびMambaと同じファミリーであり、スカラーファーストで記述されているため、バックワードパスを項ごとにチェックできます。
メモリ読み取り
PSSAに固有の部分はyに起こることです。クエリは現在のトークンと現在の状態の両方から形成されるため、取得は再帰がどこまで進んだかに条件付けられ、手元のトークンだけでなく、その状態にも依存します:
q = W_qx x + W_qh y
qh = proj(q) Poincareボールへの同相写像、|qh| < 1
読み取りは4つのスロットに制限され、温度tau_memでの双曲線距離のソフトマックスによって重み付けされます:
w = softmax(-d_H(qh, k_s) / tau_mem) (4つの最も近いスロットに対して)
m = sum_k w_k * v_k
双曲線距離はボールの境界に向かって増加するため、一般的なコンテキストを保持するスロットと特定の単一エピソードを保持するスロットは、読み取りを広げることなく分離可能です。4つのスロットは、バンクがどれだけ保持していても、トークンあたりの固定コストです。
ゲート、アダプター、MLP
読み取りは無条件に残りのストリームに参加しません。学習されたチャネルごとのゲートが、低ランクのSiLUアダプターとともに、どれだけそれが着地するかを決定します。これはターゲット付き更新を運びます:
g = sigmoid(W_gate x)
z = s * y + g (要素ごと)
W_proj m + adapter(x)
u = W_2 silu(W_1 z)
z_out = z + u
書き込みパス
書き込みは、アーキテクチャがプラスチックと呼ばれる理由です。入力状態がバンクがすでに保持しているものに対して新規であるときにスロットが挿入され、各スロットは上書きされる頻度をレート制限する不応性カウンターを保持し、高速なプラスチック更新は外部ストアに永遠に残るのではなく、閉形式のリッジ回帰によってベース遷移行列に折りたたまれます:
A_base <- A_base + (H^T H + lambda I)^-1 H^T dH
不応性カウンターは、矛盾する更新のストリームが、繰り返し証拠によってすでに安定化されたスロットを消去するのを防ぎ、統合は、バンクが長距離構造が保存される唯一の場所になるのを防ぎます。
ここでの新しさとは何か、そうでないものとは何か
再帰は標準的な選択的SSMメカニズムです。主張は、再帰状態に条件付けられた双曲線制限付き読み取り、書き込みの新規性と不応性ルール、および高速重みから遷移行列へのリッジ統合ステップです。すべてがスカラー参照パスに対して実装されており、バッチ処理および並列実装は、コミットごとに微分され、現在最大勾配誤差約3e-8まで一致しています(cargo run --release --example twin_check)。
結果
2つのモデル、同じコーパス、同じトークナイザー、同じオプティマイザースケジュール、同じシード、同じパラメータ数。1つはPSSA、もう1つは標準的なTransformerです。クリーニングされたWikiText-103の1270万トークン以上で:PSSAは3.98のトレーニング交差エントロピーで完了し、Transformerは4.43でした。これは0.45natsの差、53.7対83.7のパープレキシティです。Transformerは、PSSAが約200万トークンで既に通過した損失に到達するために、その全1270万トークン予算を費やしました。2つの曲線は交差せず、接触しません。これは、チェーン全体でログ記録されたすべての更新にわたるPSSA実行です:29,243のログ記録された更新、7.63から3.98まで、生のティックの上に41ポイントの移動平均が描かれています。
どちらのモデルも見たことのないテキストでも機能します
トレーニング損失だけでは、モデルが供給されたストリームに適合したと言うだけです。そのため、両方のチェックポイントは、どちらの実行も決して触れなかったコーパスの一部から切り取られた198,939トークンのスライスでスコアリングされました:両方の実行の各チェックポイント、64のPSSAリンクと43のTransformerリンクが、その未見のスライスの制限された9,934トークンウィンドウでスコアリングされました。曲線は交差しません:PSSAは最初のリンクから先行し、0.51nats低い値で終了します。以下の表は、完全なスライスに対する各実行の最終チェックポイントです。
保持されたスライス、198,939の未見トークン
PSSA Transformer
交差エントロピー 3.997 4.429
パープレキシティ 54.4 83.8
次のトークン精度 24.1% 18.0%
保持されたギャップ、0.43natsは、実質的にトレーニングギャップです。PSSAはよりハードに記憶しているのではなく、より良く一般化しています。そして、実行がはるかに高速です。
同じ2 vCPUマシンで固定作業、512トークン/更新で199,059トークン:415に対して1,716トークン/秒、つまり4.13倍。両方のモデルはCPUでタイミングされました。両方のアーキテクチャは、同じ学習率0.003で最適値を出したため、どちらの実行もチューニングの利点で勝っているわけではありません。スイープは120,000トークンセグメントでの短いプローブであり、最終的な数値ではなく設定チェックです。
同じCPU、同じプロンプト、同じサンプラーで200トークンを生成:
PSSA Transformer
200トークン 226 ms 2,735 ms
相対的 12倍高速
ベースライン
再帰モデルは固定サイズの状態を保持するため、各新しいトークンのコストは、それ以前の長さに関係なく増加しません。Transformerはステップごとにコンテキスト全体を再読み込みします。
実際の違いは何ですか?
再帰的な状態空間コア。学習された連続状態行列は、完全なコンテキストウィンドウでのアテンションではなく、固定サイズの状態で情報を前方に運びます。
エピソード記憶バンク。双曲線(Poincareスタイル)検索と制限付きトップ4検索を備えた512スロット。実行中に読み書きされます。
プラスチック重み。高速更新は、