GoogleのTunixシステムは、大規模なエージェント型強化学習(RL)がTPUを効率的に利用することを妨げてきたボトルネックを解消します。インタラクションデータの生成と方策の更新という作業を分離することで、TunixはTPUの利用率を1桁台からほぼフル稼働の状態まで引き上げ、計算資源の無駄を劇的に削減します。
エージェント型RLにおけるボトルネック
エージェント型RLは、より馴染みのある「次トークン予測」型の言語モデル学習とは異なります。エージェントはAPI呼び出しを送信したり、コードを実行したり、シミュレーション環境内をステップ実行したりした後、その結果に反応する必要があります。そのため、学習ループは同期的に動作します。つまり、モデルがアクションを生成し、環境が実行され、結果が返ってきて、その後に初めてモデルが勾配更新を受け取ります。環境の1ステップに数秒かかる場合、高価なTPUハードウェアはアイドル状態となり、報告される利用率は10%を下回ることもあります。この非効率性は、クラウド利用料の高騰と研究サイクルの鈍化に直結します。
Tunixのデカップル(分離)アーキテクチャ
Tunixは、「トラジェトリ生成」と「方策最適化」という2つのステージを、別々のハードウェアプールに分けることでこの問題に対処します。
- 非同期アクター (Asynchronous actors) は、安価なCPUまたはGPU上で動作します。各アクターは割り当てられた環境と継続的にインタラクションを行い、アクションと観測を記録し、得られたトラジェトリを共有ストアにストリーミングします。
- 継続的学習器 (Continuous learners) は、専用のTPU Podを占有します。学習器は中央バッファからバッチを取り出し、単一のアクターのロールアウト完了を待つことなく勾配更新を実行します。
- 高スループット・バッファ (High-throughput buffer) はその中間で、トラジェトリのステージングエリアとして機能します。学習器はバッファがデータを供給できる速度で読み取ることができるため、TPUが停止することはありません。
その結果、TPUがほぼ常に稼働し続け、利用率が100%に近づくトレーニングパイプラインが実現します。
技術的な課題とTunixによる克服
可変長の回(エピソード)とXLAの再コンパイル
JAXのXLAコンパイラは、固定されたテンソル形状に対して最適化を行います。しかし、エージェント型のタスクは長さの異なるシーケンスを生成するため、通常はコストのかかる再コンパイルが発生してしまいます。Tunixは、短いシーケンスをまとめ、長さの近いエピソードをバケット(bucket)にグループ化することで、XLAがコンパイル済みのカーネルを再利用できる程度に形状を安定させます。これにより、パフォーマンスを低下させるコンパイラのオーバーヘッドなしに、安定したスループットを維持できます。
多数のTPUチップへの大規模モデルのスケーリング
700億パラメータを超えるエージェントの学習には、重みとデータを複数のTPUノードに分散させる必要があります。TunixはJAXの ShardMap プリミティブを使用して、モデルのパラメータとアクティベーションの両方をシャード(分割)します。これにより、学習器はモデル全体をメモリに保持しながら、高速にデータを供給し続けることができます。このシャード戦略により、以前は単一のTPU Podでは到達不可能だったモデルの学習が可能になります。
分離されたパイプラインによる勾配の劣化(Stale gradients)
アクターが学習器よりも先に進んでしまうと、それらが供給するデータは現在のモデルの方策に対して「古く(stale)」なる可能性があります。Tunixは2つのメカニズムでこのドリフトを軽減します。一つは、重要度サンプリング(importance-sampling)によって古いサンプルに重み付けを行い、関連性を反映させること。もう一つは、設定可能な鮮度しきい値(staleness threshold)によって、あらかじめ設定された時間を超えたトラジェトリを破棄することです。これらにより、パイプラインが非同期に動作していても、学習の安定性が保たれます。
導入時に注意すべき点
- レイテンシの監査 – デカップルのメリットは、環境のレスポンスタイムに依存します。チームはエンドツーエンドのレイテンシを測定し、バッファが十分に満たされるようにアクタープールの規模を調整する必要があります。
- ワーカープールの設計 – 安価なCPUやGPUで多くのアクターをホストできますが、過剰に割り当てるとネットワークやストレージの競合が発生する可能性があります。バッファの取り込みレートに合わせたバランスの取れたプールが不可欠です。
- バッファの堅牢性 – 中央ストアは、新たなボトルネックになることなく、高い書き込み・読み取りレートを処理できなければなりません。低テイルレイテンシ(low tail latency)と十分な帯域幅を持つストレージシステムを選択することは、アーキテクチャにおいて譲れない要素です。
潜在的なデメリット
分離されたアーキテクチャは、個別のハードウェアフリート、永続的なバッファ、および鮮度制限を適用するための調整ロジックなど、より多くの構成要素(moving parts)を導入することになります。
まとめ
Tunixは、エージェント型RLにおける支配的なコストはモデルそのものではなく、同期的なインタラクションループによって引き起こされるアイドル時間であることを示しています。ロールアウト作業を安価なハードウェアにオフロードし、高スループット・バッファから継続的に学習するTPU Podにデータを供給することで、Googleは利用率10%未満という問題を、ほぼフル稼働のワークフローへと変貌させたのです。
