AI・機械学習
上級

RetNet(レットネット)

Retentive Networkの略で、Microsoftが2023年に提案したTransformer代替アーキテクチャ。並列学習・リカレント推論・チャンク並列推論の3つの計算モードを切り替え可能で、学習はTransformer同等の並列性、推論はRNN同等の効率性を実現する。

0 回閲覧
0 いいね

RetNetとは

RetNet(Retentive Network)は、Microsoftの研究チームが2023年7月に発表した「Retentive Network: A Successor to Transformer for Large Language Models」論文で提案されたアーキテクチャである。名称の「Retentive」は、情報を「保持」する仕組みであるRetention(保持)メカニズムに由来する。TransformerのAttentionに代わるRetentionメカニズムにより、学習・推論・チャンク処理の3つの等価な計算モードを提供する。

Retentionメカニズム

RetNetの核心であるRetentionは、以下のように定義される:

Retention(X) = (QK^T ⊙ D) V

ここでDは因果的な減衰マスク:D_{nm} = γ^{n-m}(n ≥ m)、0(n < m)。γは減衰率(0 < γ < 1)で、ヘッドごとに異なる値が設定される。

この定式化の重要な点は、同じ計算を3つの異なるモードで等価に実行できることである。

3つの計算モード

1. 並列モード(学習時):

Retention_parallel = (QK^T ⊙ D) V  [O(n²)だがGPUで高速]

Transformerと同様に全トークンを一括処理。GPUのテンソルコアを最大活用し、学習スループットを最大化する。

2. リカレントモード(推論時):

s_n = γ s_{n-1} + K_n^T V_n  [状態更新]
Retention_n = Q_n s_n           [出力計算]

RNNのように状態ベクトルsを逐次更新。O(1)の計算量とメモリで各トークンを生成でき、自己回帰推論に最適。

3. チャンクワイズモード(長文推論時):

チャンク内: 並列モードで計算
チャンク間: リカレントモードで状態伝播

入力をチャンク(例: 512トークン)に分割し、チャンク内は並列、チャンク間は再帰。長文の推論効率と品質を両立する。

ベンチマーク比較

モデルパラメータPPL (WikiText-103)学習速度推論メモリ推論速度
Transformer1.3B15.21.0xO(n) KV1.0x
RetNet1.3B15.11.0xO(1) 状態3.5x
Transformer2.7B13.81.0xO(n) KV1.0x
RetNet2.7B13.71.0xO(1) 状態3.4x
Transformer6.7B12.51.0xO(n) KV1.0x
RetNet6.7B12.40.97xO(1) 状態3.2x

Multi-Scale Retention

RetNetではマルチヘッドAttentionに相当するMulti-Scale Retentionを採用する:

  • 各ヘッドに異なる減衰率γを割り当てる(例: γ = 1 - 2^{-5}, 1 - 2^{-6}, ..., 1 - 2^{-8})
  • 低い減衰率のヘッドは近傍トークンに、高い減衰率のヘッドは遠方トークンに注目
  • GroupNormで各ヘッドの出力を正規化し、スケール差を吸収
  • これにより、異なる時間スケールの依存関係を同時にモデリング

Transformerとの比較

特性TransformerRetNet
学習並列性完全並列(O(n²))完全並列(O(n²))
推論計算量O(n)(トークンあたり)O(1)(トークンあたり)
推論メモリO(n)(KVキャッシュ)O(1)(状態ベクトル)
位置エンコーディングRoPE/ALiBi指数減衰(γ)
In-context learning優秀良好
長文処理KVキャッシュが制約チャンクワイズで効率的
学習安定性確立済みGroupNormが重要

実装上の特徴

RetNetの実装にはいくつかの工夫が必要である:

  • 複素数値演算: Retentionの減衰マスクは複素指数関数e^{iθn}を用いて位置情報をエンコード。RoPEの回転行列に相当する
  • GroupNorm: マルチスケールの減衰率による出力スケール差を正規化。RetNetの学習安定性に不可欠
  • ゲートメカニズム: swish活性化関数によるゲーティングで情報フローを制御
  • チャンクサイズ: チャンクワイズモードのチャンクサイズはハードウェアに応じて調整(A100では512、RTX 4090では256が最適)

研究の発展

RetNetの発表後、以下の派生研究が進んでいる:

  • YOCO(You Only Cache Once): Microsoftが2024年に発表した、Self-DecoderとCross-Decoderを分離するアーキテクチャ。RetNetの状態圧縮とKVキャッシュの効率化を組み合わせる
  • GoldFinch: RetNetベースの多言語モデル
  • RetNet + MoE: Mixture-of-Expertsとの組み合わせによるスケーリング研究

よくある質問(FAQ)

Q1: RetNetはなぜTransformerより速いのか? A: 推論時にリカレントモードを使うことで、各トークンの生成がO(1)の計算量・O(1)のメモリで完了するため。Transformerは過去全トークンのKVキャッシュに対するAttention計算が必要で、コンテキストが長くなるほど遅くなる。RetNetは状態ベクトルのサイズが一定のため、生成速度が系列長に依存しない。

Q2: RetNetの3つのモードはどう使い分けるのか? A: 学習時は並列モード(GPUの並列性を最大活用)、自己回帰推論時はリカレントモード(最小メモリ・最速生成)、長文の入力処理時はチャンクワイズモード(並列性と効率のバランス)を使う。これらは数学的に等価な結果を出力するため、用途に応じて自由に切り替えられる。

Q3: RetNetの限界は何か? A: 主に3点。(1)指数減衰により、非常に遠いトークンの情報が失われやすい(100K+トークンでの性能低下)。(2)Transformerほどのスケーリング実績がなく、最大13Bパラメータにとどまる。(3)FlashAttentionのような成熟した高速化ライブラリが不足している。

まとめ

  • RetNetは並列・リカレント・チャンクワイズの3つの等価な計算モードを提供
  • 学習はTransformer同等の速度、推論は3.5倍高速・定数メモリ
  • Multi-Scale Retentionで異なる時間スケールの依存関係をモデリング
  • Microsoftが主導する研究で、YOCO等の派生アーキテクチャに発展
  • 13Bパラメータまでの検証で、Transformerと同等のPerplexityを達成