Recurrent Looped Transformer、系列長に応じて計算経路を拡張
背景
Transformerは系列を並列に処理できる一方、各トークンに適用される層数は通常固定されています。パリティ計算や状態機械のシミュレーション、置換追跡のように、入力を読むたびに内部状態を更新する課題では、この固定深度が長い系列への外挿を妨げる可能性があります。
この問題に対して、論文は Recurrent Looped Transformer(RLT) を提案しました。固定された層数を単純に増やすのではなく、系列に沿った再帰的な状態更新を導入する設計です。
仕組み
RLTは、8層のネットワークを並列因果エンコーダーと再帰デコーダーに分けます。エンコーダーは現在のトークンと過去の情報から表現を作り、デコーダーはその出力を、直前のトークンで得られた最終デコーダー状態と統合します。
これにより、エンコーダーによる並列処理を維持しながら、状態はトークンから次のトークンへ順番に伝播します。系列が長くなるほど情報が再帰更新を受ける回数は増えますが、各トークンにかかる局所的な計算コストは比較的一定に保たれます。
研究では、8層をエンコーダーとデコーダーにどう配分するかを変え、並列表現能力と状態追跡能力のバランスも調べています。
実験結果
6種類のアルゴリズム課題で、5種類の層分割を3つの乱数シードで評価し、8層の標準Transformerと比較しました。
- パリティ: 最大40ビットで学習した場合でも、2種類のRLT構成は256ビットへ外挿し、全シードで100%の正解率を達成しました。標準Transformerはほぼランダム水準でした。
- S_5置換追跡: 学習長の8倍の長さでは、RLTの最終状態正解率が97%に達した一方、標準Transformerは1%未満でした。デコーダーを深くした構成ほど精度も高くなりました。
- モジュラー算術: 学習長を超える評価で、RLTは最大93%、標準Transformerは33%でした。
- フィードバック除去: 再帰フィードバックを取り除くと、パリティとswap-based S_5の性能は、どの分割でもランダム水準まで低下しました。
並列性とのトレードオフ
フィードバックを4トークンごとに更新するチャンク方式も検証されています。既知のチャンク内のトークンを並列実行できるため、64ビットのパリティでは99%の精度を維持しました。
しかし置換追跡では事情が異なります。長さ64のswap-based S_5では、精度が100%から20%へ下がりました。つまり、パリティのような集約的な計算は遅延更新に耐えやすい一方、順序依存の状態遷移にはトークンごとのフィードバックが必要です。
意義と限界
RLTは、Transformerの深さをただ増やすのではなく、その一部を系列方向の再帰時間へ変換する考え方を示します。長さ外挿にはモデル規模だけでなく、入力に応じて状態を適切な粒度で更新できる構造が重要だという示唆を与えます。
ただし、今回の評価はアルゴリズム課題が中心です。一般的な言語モデリング性能の向上を直接示すものではありません。また、並列性と細かな状態追跡のどちらを優先するかは、今後も設計上の重要な課題になります。
コメント
ログイン状態を確認中…
コメントを読み込み中…