記事のサマリー(TL;DR)
- TRL v1.14のAsyncGRPOTrainerが、LoRAアダプタのみを学習し、そのアダプタだけをvLLMに同期する機能に対応しました(PR #7017)。rank-1アダプタは数メガバイト程度のため、NCCLではなく各Jobにマウントされた Storage Bucket を経由してやり取りできます。
- 学習用JobとvLLMレプリカ用Jobを別々のマシン上で稼働させ、両者の間に小さなプロキシを配置して認証ヘッダの付与、KVプレフィックスを保持するレプリカへのルーティング、全レプリカへのアダプタ更新のブロードキャストを行う構成を構築しました。
- AsyncGRPOのメトリクスを見ながら5回の実験を行った結果、同一レシピで500ステップの学習時間が3時間27分から53分まで短縮されました。
詳細
背景:LoRAとRLの相性
LoRAによる学習はRL(強化学習)に特に適しているとされ、Thinking Machinesのブログ「LoRA Without Regret」では、policy-gradient RLにおいてLoRAがrank 1でもフルファインチューニングに匹敵する性能を出せることが示されています。これは、アドバンテージ関数が1エピソードあたり約O(1)ビットの情報しか与えないため、各ステップから学習できる情報量がそれほど多くなく、rank-1アダプタでもそれを吸収するのに十分な容量を持つためだと説明されています。
これにはシステム面での帰結もあります。1.5Bモデル用のrank-1アダプタは数メガバイト程度であるのに対し、フルモデルは約3GBです。更新のたびにフルポリシーを推論ワーカーに送る代わりに、アダプタだけを送ることができます。vLLMは複数のアダプタを同時にロードしておくことも可能で、古いロールアウトは開始時点のポリシーのまま完了し、新しいロールアウトは最新のポリシーを使用します。
課題:Hugging Face Jobs間でのトレーナーと推論サーバの分離
TRLのAsyncGRPOTrainerは既に学習と生成を分離する設計になっており、トレーナーとvLLMは別マシン・別ペースで動作できます。ただしこれは、両プロセスがファイルシステムを共有するか、NCCLグループを形成できる単一ノードまたはクラスタ環境では容易ですが、Hugging Face Jobsでは事情が異なります。
1つのHF Jobは1つのVM上で動く1つのコンテナであり、現時点では1つのJobがトレーナーとvLLMサーバ群を保持するために複数ノードを立ち上げることはできません(1ノードあたり最大8xH200という制約もあります)。AsyncGRPOTrainerはまさにそうした規模のために作られているため、「トレーナーと推論サーバが同一ノードを共有するという前提を外したら、どこまでできるか」が課題となりました。
フルウェイトの同期では、更新のたびにギガバイト単位のデータをマシン間で移動する必要があり、これは密結合クラスタにおけるNCCLの役割ですが、Jobs間ではノード間通信ができず、共有ローカルディスクも共有localhostも存在しません。LoRAであれば同期は数メガバイトで済みます。ファイルシステムについては、HF Jobsは Storage Bucket を裏付けとするボリュームを提供しており、これを各JobにFUSEファイルシステムとしてマウントすることで、ノード間の共有FSとして機能させることができ、Jobs間のネットワークパスは一切不要になります。
構成は次のように小規模なものになりました。
- LoRA(および後述するFSDP)を用いたAsyncGRPOTrainerを実行するトレーナーJob
- ベースモデルとトレーナーが最後に公開したアダプタを配信する2つのvLLM Job
- トレーナーからサーバへアダプタを渡す経路となる、3つ全てに同一パスでマウントされたStorage Bucket
- プロキシサーバ(各ロールアウトをそのKVキャッシュを保持する可能性が高いレプリカにルーティングし、全vLLMレプリカへアダプタ更新をブロードキャストする役割)
アーキテクチャ:Hugging Face JobsとStorage Bucketの活用
AsyncGRPOTrainerの新しいアダプタのみの同期経路は次のように動作します。トレーナーはテンソルをvLLMに送信しません。数optimizerステップごとに、<output_dir>/.vllm_lora/trl-policy-v{N} の下にアダプタを保存し、原子的なリネームでディレクトリを公開した上で、そのパスをvLLMの /v1/load_lora_adapter エンドポイントに送信します。vLLMはディスクからファイルを読み込むため、ロールアウトワーカーは model="trl-policy-v{N}" を指定してリクエストできます。これはvLLMの既存のランタイムアダプタロード機能を利用したものです。
このエンドポイントはテンソルではなくパスを受け取るため、トレーナーとサーバはファイルシステムを共有している必要があります。Slurmクラスタではネットワークファイルシステムがこれに該当しますが、Jobsでは各Jobの同一パスにStorage Bucketをボリュームとしてマウントすることで同じ仕組みを実現しています。内部では hf-mount を使用しており、バケットをコンテナ内のPOSIXファイルシステムとして公開します。
# every Job gets the same bucket at the same absolute path
hf jobs run ... -v hf://buckets/aminediroHF/asyncgrpo-lora-buckets:/lora ...
TRLやvLLM側のコードは一切変更する必要がありませんでした。トレーナーは /lora/<run>/.vllm_lora/ に書き込み、サーバは同じパスから読み込みます。POSTリクエストに含まれるパスは、どのコンテナ内でもそのまま有効です。
なお、チェックポイントと最終アダプタもバケットに保存しています。HF Jobsは一時的な存在ですが、最終アダプタは常にバケットに永続化され、Jobが停止しても失われないため、プリエンプトされたトレーナーは学習を再開できます。
3つのJobの詳細
vLLMレプリカ
各レプリカは1GPUを使用し、標準の vllm/vllm-openai イメージを使います。有効化するのはランタイムLoRAロードと、十分な数のアダプタスロットの確保です。
アダプタスロット数は max_staleness から導かれます。AsyncGRPOTrainerでは、重み同期のたびにポリシーバージョンが1つ増え、max_staleness はロールアウトサンプルがトレーナーに破棄される前に現在のポリシーから何バージョン遅れてよいかを示します。max_staleness=4 の場合、trl-policy-v3 の下で生成されたサンプルは、トレーナーが v7 にあっても学習に使われます。v3 で始まったロールアウトは v3 のまま完了できる必要があります。つまりいかなる時点でも、vLLMは現在のポリシーとその前の4バージョン分をあわせて提供する必要があります。そのためトレーナーは max_staleness + 1 個のアダプタバージョンを登録し続け、それより古いものをアンロードします。同期のたびに新バージョンをロードしてから最も古いものをアンロードするため、入れ替え中はもう1スロット多く必要です。これにより --max-loras 6 となります。5つしかない場合、vLLMは同期のたびに、まだロールアウトが進行中のポリシーを黙って追い出してしまいます。
# --expose 8000 reachable at https://<job_id>--8000.hf.jobs
# -v ...:/lora:ro read-only: the server only reads adapters
# VLLM_ALLOW_RUNTIME_LORA_UPDATING=1 enables /v1/load_lora_adapter
# VLLM_SERVER_DEV_MODE=1 enables /pause, /resume, /server_info (TRL needs all three)
# --max-loras 6 max_staleness=4 -> 4+2 adapter slots
for replica in 1 2; do
hf jobs run --detach --flavor h200 --timeout 8h --secrets HF_TOKEN \
--expose 8000 \
-v "hf://buckets/${BUCKET}:/lora:ro" \
-e VLLM_ALLOW_RUNTIME_LORA_UPDATING=1 \
-e VLLM_SERVER_DEV_MODE=1 \
-- vllm/vllm-openai:v0.27.1 \
vllm serve Qwen/Qwen2.5-Math-1.5B --host 0.0.0.0 --port 8000 \
--max-model-len 4096 --logprobs-mode processed_logprobs --generation-config vllm \
--enable-lora --max-lora-rank 1 --max-loras 6
done
vLLMはバージョン v0.27.1 に固定しています。vLLMの更新は速く、上記のフラグやランタイムLoRAエンドポイントはこのバージョンで公開されているものなので、このバージョンもレシピの一部として扱うべきだとしています。
もう一つの設計として、トレーナーが最新のアダプタだけを保持し、常に同じ名前で公開する方法も考えられます。これを採用しなかった理由は、vLLMがプレフィックスキャッシュをアダプタ名でキーイングしているためです。同一の名前を使い続けると、以前の重みで計算されたKVブロックが入れ替え後も一致してしまい、プリフィルがやり直されず、あるロールアウトがあるポリシーバージョンのプレフィックスと次のバージョンのデコードを混在させて取得してしまう可能性があります。トレーナー側にはそれを検知する手段がなく、ratioが1から乖離するという形で現れます。バージョン付きの名前にすれば、名前は常に一組の重みを指し、キャッシュされたプレフィックスが新しいバージョンと一致することは決してなくなります。
データセットの選定:Sanityセット
使用したのは sail/Sanity-Test-R1D-1.5B で、これは論文「Defeating the Training-Inference Mismatch via FP16」(Qi et al., 2025)由来のデータセットです。再現コードは sail-sg/Precision-RL にあります。著者らはDeepSeek-R1-Distill-Qwen-1.5Bを用いて各MATH問題につき40通りの回答を生成し、成功率が20%から80%の間にある問題を残した結果、1,460問が得られました。
このデータセットは、そのモデルにとって「既に解けている」わけでも「完全に絶望的」でもない問題群であるため、RLの検証に適しており、モデルが学習して改善するための良い初期シグナルが得られるとされています。また、エンドツーエンドのテストとしても頑健で、あるvLLMレプリカがアダプタ名の下で黙ってベースモデルを配信してしまっている場合、数十ステップ以内にカーブにその兆候が現れます。さらに、このデータセットは2時間未満で一巡できる規模の小ささも備えています。
ハイパーパラメータも論文のLoRAスクリプト(oat/scripts/lora)から採用しており、モデルはQwen/Qwen2.5-Math-1.5B、LoRA rank 1・alpha 2、学習率4e-5、プロンプトあたり8サンプル、1ステップあたり128completion、最大生成トークン数3,000、コンテキスト長4,096トークンとなっています。
トレーナー
トレーナーは同じ vllm/vllm-openai:v0.27.1 イメージにTRLをインストールしたものを使用しています。記事執筆時点ではPRブランチを使用しましたが、同じコードは現在TRL v1.14に含まれています。学習スクリプトは通常のAsyncGRPOTrainerスクリプトで、Job固有の値は出力ディレクトリとサーバURLのみです。
from peft import LoraConfig
from trl.experimental.async_grpo import AsyncGRPOConfig, AsyncGRPOTrainer
config = AsyncGRPOConfig(
output_dir="/lora/sanity-lora-r1", # on the bucket: adapters, checkpoints and the final adapter all land here
vllm_server_base_url="http://localhost:8000", # the proxy, not a vLLM Job; TRL never sees the Jobs URLs
max_staleness=4,
weight_sync_steps=4, # publish an adapter every 4 optimizer steps
save_strategy="steps",
save_steps=50, # checkpoints go to the same bucket -> resume after preemption
...
)
trainer = AsyncGRPOTrainer(
model="Qwen/Qwen2.5-Math-1.5B",
args=config,
peft_config=LoraConfig(r=1, lora_alpha=2, target_modules="all-linear"), # plain LoRA vLLM can serve as-is
...
)
初期化時、TRLは /server_info を呼び出し、lora_config が見つかればアダプタのみの同期を使用します。DoRA、modules_to_save、--max-lora-rank を超えるrankなど、vLLMが直接配信できない設定の場合は警告付きでマージ済み重み同期にフォールバックします。正常時はログに「Adapter-only vLLM sync enabled」と表示されます。
プロキシ
トレーナーとvLLM Jobsの間にプロキシを置く理由は2つあります。1つ目は、公開されたJobのポートは全てのリクエストに Authorization: Bearer <HF token> ヘッダが必要であり、プロキシがこのヘッダを付与することで、TRL側はこれを意識する必要がなくなる点です。
2つ目は、複数GPUで生成を行いたいという要件です。単一vLLMサーバでこれを実現する通常の方法は --data-parallel-size > 1 ですが、TRLはこのモードではアダプタのみの同期を拒否します。理由は明確で、/v1/load_lora_adapter の呼び出しは応答したDPランクにしか届かないため、他のランクは新しいポリシー名の下で古いベースモデルを配信し続けてしまうからです。Jobs環境ではそもそも各レプリカが独立したマシンであるため、データ並列は一段上、つまりアダプタロードを全レプリカに展開する仕組みの側で行う必要があります。
そこで、トレーナーJob上の 127.0.0.1:8000 で小さなプロキシを稼働させ、TRLからは単一のvLLMサーバであるかのように扱わせています。プロキシはヘッダの付与に加え、機能的に次の2つを行います。1つは、各completionリクエストを、あるプロンプトの8つのロールアウトがそのプレフィックスが既にキャッシュされているレプリカに着地するよう選んで送ること。もう1つは、アダプタロード、pause、resumeといった状態変更を伴うリクエストを全レプリカにブロードキャストし、ポリシー名がどこでも同じ意味を持つようにすることです。
KVプレフィックスによるロールアウトのルーティング
completionの生成には、ワークロード特性の異なる2つのフェーズがあります。プリフィルはプロンプト全体を一度に処理し、全プロンプトトークンについてattentionのkeyとvalueを計算します。デコードフェーズは1トークンずつ生成し、各新規トークンはそれ以前の全トークンのkeyとvalue(KVキャッシュ)にattendします。attentionが因果的であるため、あるトークンのKVはそれより前のトークンのみに依存し、後に来るものには依存しません。したがって、プレフィックスを共有する2つのリクエストはそのプレフィックスのKVも共有でき、それを既にキャッシュしているレプリカはその部分のプリフィルを丸ごとスキップできます。
vLLMはプレフィックスKVキャッシュを16トークンのブロック単位で保持しています。GRPOの性質上、ロールアウトワーカーは同一プロンプトに対しG個(この場合G=8)のリクエストを送ります。これらが全て同じレプリカに到達すれば、最初のリクエストがプリフィルを計算し、残り7つはそれを再利用できます。ラウンドロビン方式でルーティングすると、半数がプレフィックスをキャッシュしていないレプリカに送られ、それら4リクエストはプリフィル処理をやり直すことになり、GPU計算リソースを無駄にします。
ルーターの役割は、どのレプリカがどのブロックハッシュを見たかを追跡することです。重要な点として、ハッシュは連鎖しており、ブロック3のハッシュはブロック3だけでなくブロック1、2、3を表します。これは因果的attentionを反映しており、ブロック3のKVはブロック1と2が同一である場合にのみ有効だからです。またハッシュ連鎖はアダプタ名でシードしています。KVキャッシュは生成に使われたアダプタにも依存するため、ポリシーv3でキャッシュされたプレフィックスはポリシーv4には使えないからです。
具体的な135トークンのcompletionリクエスト例(Sanityデータセットの問題より)を用いて、レプリカ選択の流れが以下のように説明されています。
- プロンプトをブロックに分割する。 ルーターはトークンIDを受け取り、vLLMと同様に16トークンのブロックに切り分けます。完全なブロックのみをハッシュ化するため、末尾の7トークンは無視されます。
- プレフィックスをハッシュ化する。 各ブロックは前のハッシュと連結してハッシュ化され、アダプタのシードから始まります。同じ最初のkブロックを持つ2つのプロンプトはhkまで同じハッシュを持ちますが、1つのブロックが変わると、それ以降のハッシュも全て変わります。同じプロンプトでも
trl-policy-v4は別のシードから始まるため、v3のエントリとは一致しません。これは古いKVブロックが異なる重みで計算されているため望ましい挙動です。 - 2つのプロンプトを比較する。 問題1は103トークンで、両プロンプトとも同一の23トークンのチャットテンプレートで始まります。最初のブロックは同一ですが、ブロック2には既に問題文が含まれているため、そこからハッシュが分岐します。
- 所有者を記録する。 各ハッシュについて、ルーターはどのレプリカがそれを配信したか、そしてどのハッシュがその後に続いたかを記憶します(後続集合の上限は2とし、あるブロックが1つの続き方をするか複数の続き方をするかだけを知れば十分としています)。数プロンプト経過後、テンプレートブロックh1は両レプリカが所有し複数の後続を持つ一方、h2からh8はレプリカAのみが所有し各々単一の後続を持ち、h2’からh6’はレプリカBのみが所有する、といった状態になります。実際には、あるrun内の全プロンプトが同じトークン列(この場合はチャットテンプレートとシステムプロンプトで、全1,460問の最初の23トークンに相当)で始まります。これらのブロックは数秒以内に全レプリカのキャッシュに入るため、それらが一致してもどのプロンプトがどこにあるかは分かりません。あるブロックが「common(共通)」とみなされるのは、全レプリカがそれを配信済みであるか、複数の後続を持つ場合です。common なブロックは特定のプロンプトを識別しないため、ルーティングの際には無視されます。
- レプリカを選ぶ。 ルーターは各レプリカで一致する先頭ブロック数を数え、共通プレフィックスを除きます。残った数がそのプロンプト固有のブロック数です。あるレプリカが固有ブロックを持ち、かつ最も負荷の低いレプリカより最大8リクエストしか先行していない(swampedでない)場合、そのリクエストはそのレプリカに送られます(affinity hit)。あるレプリカが固有ブロックを持つが8リクエスト以上先行している場合は、キャッシュ活用を諦めて最も負荷の低いレプリカに送ります(spill)。どのレプリカも固有ブロックを持たない場合は新規プロンプトとみなし、最も負荷の低いレプリカ(同点の場合はラウンドロビン)に送ります(unmatched)。
4つのリクエストへのルール適用例も紹介されています。レプリカAとBが共に3リクエストを処理中で、テンプレートブロックh1のみが両者に既知の状態から開始します。
- リクエスト1(問題0、ロールアウト1):両レプリカとも一致するのはテンプレートブロックのみでcommon扱いのため、unmatchedとなりラウンドロビンでAに送られます。h2からh8はAの所有として記録され、Aの処理中リクエストは4になります。
- リクエスト2(問題0、ロールアウト2、同一プロンプト):Aは8ブロック全てが一致、Bはテンプレートのみ一致。common分を除くとAは7個の固有ブロック、Bは0個。Aの先行数はBよりわずか1(上限8以内)のため、affinity hitとしてAに送られ、Aは既にプロンプト全体をKVキャッシュに保持しています。
- リクエスト3(問題1、ロールアウト1、新規プロンプト):両レプリカともテンプレートブロックのみ一致のためunmatchedとなり、この時点で5対3で負荷の低いBに送られます。h2’からh6’はBの所有として記録されます。
- リクエスト4(問題0、ロールアウト9):この時点でAが12リクエスト、Bが3リクエストとします。Aは依然7個の固有ブロックを持ちますが、Bより9リクエスト先行しており上限を超えるため、リクエストはBにspillします。Bは問題0を一度プリフィルし、h2からh8もBの所有として記録され、以降その問題は両レプリカで共通(common)となり、後続のロールアウトは負荷のみで配置されます。
- プリフィルを再利用する。 リクエスト2はこの仕組み全体の存在意義であり、リクエスト1が計算したプリフィルを再利用します。ブロック1から8は既にAのKVキャッシュにあるため、Aは直接デコードに進みます。仮にBに送られていれば、Bは135トークン全体を再度プリフィルすることになり、Aのキャッシュは使われないままになります。リクエスト3はcommonルールが必要な理由を示しており、これがなければ共有チャットテンプレートのせいで、あらゆる新規プロンプトがキャッシュヒットのように見えてしまいます。リクエスト4は負荷を一定範囲に保つ仕組みで、1回のプリフィルを節約するために、あるレプリカを大きく遅れさせる価値はないという判断です。
ルーティングロジックの擬似コードも示されています。
def choose(self, upstreams, model, prompt):
hashes = self.block_hashes(model, prompt) # chained blake2b over 16-token blocks, seeded with `model`
matched = self.matched_prefix(hashes) # per replica: leading blocks it has served
common = self.common_prefix_len(hashes) # leading blocks that identify no prompt (see below)
specific = [max(0, m - common) for m in matched] # what actually distinguishes replicas
least = min(u.inflight for u in upstreams)
best = max(range(self.n), key=lambda i: (specific[i], -upstreams[i].inflight))
if specific[best] > 0 and upstreams[best].inflight - least <= self.cfg.imbalance:
pick = best # affinity: the replica that has this prompt, and is not swamped
else:
candidates = [i for i in range(self.n) if upstreams[i].inflight == least]
pick = candidates[self.rr % len(candidates)] # spill or new prompt: least-loaded, round robin on ties
self.rr += 1
...record `pick` as an owner of every block, and each block's successor...
return upstreams[pick]
共通プレフィックスの扱いが最も厄介な部分だとされています。全てのリクエストが同じシステムプロンプトとチャットテンプレートで始まるため、単純な最長一致では最初のレプリカがほぼ全ての新規プロンプトにマッチしてしまいます。そこで共有プレフィックスをfan-out(分岐数)で検出しています。複数の異なる後続を持つブロックはcommonとみなし、常に同じ後続に繋がるブロックは特定のプロンプトに属するとみなします。この共通プレフィックスより後のブロックのみをaffinityの対象としてカウントします。
アダプタのブロードキャスト
プロキシはアダプタロードを全レプリカにブロードキャストする必要があります。この操作は「オール・オア・ナッシング」として扱われます。各レプリカは独自のバケットマウントを持つため、新しいアダプタを必ずしも同時に認識するとは限りません。No adapter found for <path> というエラーは通常、あるバケットマウントがまだ追いついていないことを意味し、そのレプリカのみリトライします。それ以外のエラーの場合は、既に受け入れたレプリカからアダプタをアンロードし、あるポリシー名が一部のレプリカにしか存在しない状態を避けます。
async def load_one(u):
while True:
status, _, out = await send(u, "POST", "/v1/load_lora_adapter", headers, body)
if status == 200 or "No adapter found" not in out.decode() or time.monotonic() > deadline:
return u, status, out
await asyncio.sleep(cfg.lora_retry_s) # this replica's mount has not seen the directory yet
results = await asyncio.gather(*(load_one(u) for u in ups))
if any(st != 200 for _, st, _ in results):
await asyncio.gather(*(send(u, "POST", "/v1/unload_lora_adapter", headers, unload)
for u, st, _ in results if st == 200))
return web.Response(status=504 if timed_out else st, text="rolled back on the others")
/pause、/resume、/v1/unload_lora_adapter も同様にブロードキャストしています。/health は全レプリカが健全な場合にのみ200を返し、/server_info と /v1/models は1つの応答のみで足りるとしています。
TRLから見ると、プロキシは data_parallel_size=1 の単一サーバであるため、アダプタのみの同期を選択します。当初、Python asyncioベースのプロキシがボトルネックになるのではないかという懸念がありましたが、少なくともこの規模では問題になっていません。同時に処理する非ストリーミングJSONリクエストは最大128件で、ルーティングは数個のハッシュを計算するだけであり、1スレッドで十分処理できています。より多くのトラフィックを扱う必要がある、より高度なルーターであれば、より高速な言語での実装が必要になるだろうとしています。
実行結果全体
以下の数値は、トレーナーがtrackioに記録したメトリクスによるものです。このrunはQwen/Qwen2.5-Math-1.5B、all-linearに対するLoRA r=1、1ステップあたり128 completion、プロンプトあたり8ロールアウトという設定で、500ステップを実行し、50ステップごとにチェックポイントを保存します。トレーナーはh200x2のJob、2つのvLLMレプリカはそれぞれh200のJobを使用します。3つ全てを稼働させるコストは1時間あたり約20ドルとしています。
重み同期
| 同期1回あたり(トレーナー計測) | 126回の同期のうち直近(p50) | |
|---|---|---|
| 同期全体 | 30.8秒 | 8.5秒(最小6.6、最大9.2) |
| うちpause両レプリカ | 0.3秒 | 0.3秒 |
| アダプタのall-gatherとバケットへの保存 | 0.6秒 | 1.1秒 |
| 両レプリカがアダプタを受け入れる | 約29秒 | 約7秒 |
252件のアダプタロード(126回の同期×2レプリカ)は全て成功しました。うち6件が2回目の試行で成功し、246件が3回目の試行で成功しました。
ルーティング
64,728件のロールアウト終了時点でのプロキシのカウンタは次の通りです。
routed [31928, 32800]
affinity 54712
spilled 820
unmatched 9196
プロンプトあたり8ロールアウトの場合、8件のうち少なくとも1件はコールド(キャッシュなし)になるため、理論上の最小unmatched率は12.5%です。実際のルーターの結果はunmatchedが14.2%、affinity hitが84.5%、spillが1.3%でした。実測負荷に基づくより詳細な推論側メトリクスを使わない限り、これ以上の改善余地はあまりないとしています。
ボトルネックの所在
最初の構成には明らかな問題があり、ボトルネックは生成ではなくトレーナー側にありました。500ステップの間の数値は次の通りです。
| p50 | |
|---|---|
| optimizerステップ 1回あたり | 22.9秒 |
| forward + backward | 21.9秒 |
| ロールアウト待ち | 0.02秒 |
| ロールアウトキューの占有 | 512中476 |
| トレーナーMFU | 3.9% |
ロールアウトキューは常に満杯で、ワーカーはほぼバックプレッシャーによってブロックされている状態でした。この構成では2番目のレプリカは実質的に無駄になっていました。この記事では、この後トレーニングと生成の間でボトルネックを移動させながら実行速度を3.9倍にした一連のrunを紹介しています。
Reward(図1: trackio run r1-dp2)
Rewardは500ステップの間に0.15から0.44まで上昇し、ratioは全ステップを通じて0.9993から1.0004の範囲に収まっています。500ステップに要した時間は3時間27分でした。平均rewardは最初の20ステップの0.145から最後の20ステップの0.438まで上昇しています。さらに重要な点として、ratioは全ステップを通じて1.000に維持されており、vLLMが配信するポリシーは常にトレーナーがロールアウトを評価する際に使用したポリシーと一致していました。これは126回の同期すべてで維持され、平均staleness(遅延バージョン数)は1.5、最大は4でした。この結果を「LoRA AsyncGRPOが機能する動かぬ証拠」としています。
ボトルネックの追跡
非同期RLは学習と生成の間のパイプラインであり、片方だけを高速化しても、もう一方が追いつけなければ効果はありません。AsyncGRPOTrainerには、これを直接確認できるだけの詳細なタイミングとメトリクスが追加されています。
主なメトリクスは次の通りです。perf/rollout_wait_s はトレーナーがサンプルを待った時間、rollout/backpressure_s は生成側がロールアウトキューの空きを待った時間を示し、この2つは互いに正反対の関係にあるため、両方が高い状態にはならないはずです。キューサイズと合わせて見ることで、どちら側が遅いかが分かります。
読み解くために継続的に確認するメトリクスとして、次の4グループが挙げられています。
perf/step_sとperf/fwd_bwd_s:optimizerステップの所要時間と、そのうちforward+backwardが占める割合。ステップ時間がforward+backwardとほぼ同じであれば、トレーナーは明らかに計算律速です。perf/rollout_wait_s:トレーナーがステップを開始する前にサンプルを待った時間。ゼロに近ければ生成が学習より先行しており、サンプルはすぐに利用可能です。sample/rollout_queue_size対queue_maxsize:両者の間のバッファ。満杯であれば生成側がスロットリングされており、空であればトレーナー側が枯渇しています。rollout/backpressure_sとrollout/score_block_s:バッファが満杯でロールアウトワーカーがブロックされていた時間。ワーカーは2段階のパイプラインで、生成が完了したグループをスコアリング段階に渡し、スコアリングがスコア済みサンプルをロールアウトバッファに投入します。バッファが満杯だとスコアリングが投入できずブロックされ(rollout/backpressure_s)、その結果スコアリングが自身の入力キューを消化できなくなり、生成も次のグループを渡せなくなります(rollout/score_block_s)。両者は同一の停滞をスコアリング段階と生成段階それぞれで見ているものです。
診断は単純で、キューが満杯かつロールアウト待ちがゼロでbackpressureが高ければトレーナーが遅すぎ、キューが空でロールアウト待ちが上昇しbackpressureがなければ生成が遅すぎることを示します。perf/mfu_wall_clock と perf/mfu_fwd_bwd を比較することでも、トレーナーGPUが学習ではなく待機に費やしている時間の割合が分かります。
5つの実験を行い、それぞれ前のrunのダッシュボードで確認された問題から出発しています。特に断りがない限り、モデル、レシピ、3-Job構成は同一です。以下の名前はtrackioのrun名です。
Run 1、r1-dp2:トレーナーが追いつけない
perf/step_s は22.9秒、perf/fwd_bwd_s は21.9秒で、forwardとbackwardがステップ時間の96%を占めていました。キューは満杯のままで、トレーナーはロールアウト待ちが0.02秒しかない一方、ロールアウトワーカーは1グループあたり15秒間backpressureでブロックされていました。2つのvLLMレプリカはトレーナーが消費するより速く生成しており、報告された4.6k tokens/sは実際の上限ではなく、単に出力の行き場がなかっただけです。
batchメトリクスがMFU 3.9%という低さを説明しています。batch/microbatches_per_step は64、batch/samples_per_row は1.0でした。各ランクは約1.2kトークンの1シーケンスを、1ステップあたり64回処理しており、これは参照レシピの per_device_train_batch_size=1 に由来します。1.5BモデルをH200で動かす場合、これは完全にレイテンシ律速の状態です。
Run 2、r1-dp2-tb16k:マイクロバッチを詰める
バッチサイズは変えず、1ステップあたり128 completionは維持したまま、GPU上でのレイアウトのみを変更しました。1マイクロバッチに1シーケンスではなく、多数のシーケンスを密に1行に詰め込みます。トレーナーはこれを「token-budget batching」として提供しており、token_budget > 0 を設定すると、ランクごとに1行にパディングなしで複数サンプルを詰め込みます。1 optimizerステップは gradient_accumulation_steps 回分の行を処理します。今回は token_budget=16384、gradient_accumulation_steps=6 に設定しました。
batch/samples_per_row は1.0から約12.7に上昇し、マイクロバッチ数は64から6に減少、行の充填率は95%に達しました。forwardとbackwardは21.9秒から5.6秒に短縮し、MFUは3.9%から19%に上昇しました。行がより効率的に詰まったため、平均長からの見積もりよりも多い、1ステップあたり約150サンプルで学習することになりました。vLLM側は何も変更していないにもかかわらず、生成速度も4.6kから25k tokens/sへと跳ね上がりました。キューが常に満杯ではなくなったため、レプリカがようやく本来の性能を発揮できるようになったためです。この結果から、パイプラインの各段階を個別に最適化するのは望ましくなく、遅い段階がそれ以前の全ての段階の実際の性能を隠してしまうため、システム全体を評価する必要があるとしています。
Run 3、r1-dp2-tb16k-nockpt:forwardの再計算を止める
perf/fwd_s は1.34秒であるのに対し perf/fwd_bwd_s は5.6秒でした。通常のbackwardはforwardのおよそ2倍のコストで、ベース重みが凍結されている場合は1倍に近くなるはずです。3.2倍という比率は不自然でした。原因は AsyncGRPOConfig が gradient_checkpointing=True をデフォルトにしていることで、各マイクロバッチがbackward時にforwardを再計算していました。これは、16kトークンの行が141GBのH200上でわずか25GBしか使わなかった理由も説明しています。これはこのトレーナーにとってのメモリ最適化ですが、このケースではモデルが小さく、activationをbackward用に保持したままVRAMに収まるため不要でした。
gradient_checkpointing=False にすると、forwardとbackwardは4.6秒に低下し、ほぼforward1回分減少、MFUは23%に達しました。キューは71まで低下し、ロールアウト待ちは0.04秒から0.6秒に上昇しました。トレーナーが2つのレプリカの生成速度より速くサンプルを消費するようになり、ボトルネックを生成側に移すことに成功しました。これによりさらに2つのコストが明らかになりました。4ステップごとの7.6秒の重み同期は、ウォールクロック時間の25%を占めるようになりました(1ステップが23秒だった時は8%でした)。また、backwardはforwardより依然として2.5倍遅く、ベース重みが凍結されているにもかかわらず、通常のモデル計算らしくない約2秒がステップごとに存在していました。
Run 4、r1-dp3-tb16k-nockpt:3つのレプリカと意外な発見
生成が遅すぎる状態になったため、3つ目のレプリカを追加しました。またプロキシのアダプタリトライ間隔を2秒から0.5秒に短縮し、FSDP2の再gatherがbackwardの余分な2秒の原因かどうかを確認するため fsdp_reshard_after_forward を無効化しました。
重み同期は7.6秒から5.8秒に低下し、短いリトライ間隔が効果を示しました。forwardとbackwardは4.6秒のままで、resharding原因説は否定されました。生成速度は25kから26k tokens/sへとほとんど動かず、3つ目のレプリカは実質的に何もしていませんでした。
原因は rollout/inflight にありました。どのrunでも128で一定であり、プロキシ上ではこれらのリクエストが44+43+41の3レプリカに分割されていました。max_inflight_tasks はロールアウトワーカー全体の並行数を制限するものであり、レプリカごとの制限ではありませんでした。1.5BモデルをH200で動かす場合、43と130の同時シーケンス処理はほぼ同じコストであり、128リクエストを3GPUに分割しても2GPUに分割した場合とほぼ同じスループットになります。つまりvLLM側が上限だったのではなく、クライアント側の定数設定が上限だったということです。この値は、公開のJobsプロキシを介して数百件の長時間HTTPSリクエストがどう振る舞うか分からなかったため、保守的に設定していたものでした。この時点で130,000件のロールアウトcompletionがこの制約を経由しても、転送エラーは1件も発生していませんでした。
Run 5、r1-dp3-inflight384:in-flightリクエスト上限の引き上げ
max_inflight_tasks=384、queue_maxsize=768 とし、他は変更していません。384件のin-flightリクエストにより、各レプリカは128件を担当します。キューは768のうち約690まで速やかに満たされ、その水準を維持しました。backpressureは5秒に戻り、ロールアウト待ちは0.03秒まで低下しました。再びトレーニングがボトルネックになりました。forwardとbackwardは4.6秒、重み同期は償却して1.5秒加算、ステップ時間の中央値は4.8秒でした。平均stalenessはより大きなキューでサンプルが長く待つようになったため1.5から2.0に上昇しましたが、これは max_staleness=4 を下回っており、ratioは1.000近くを維持しました。
なお、Run 2からRun 4はダッシュボードが問いに答えた時点で早期に停止しています。
スコアボード(500ステップ)
| run 1 r1-dp2 | run 5 r1-dp3-inflight384 | |
|---|---|---|
| ウォールクロック | 3時間27分 | 53分 |
| perf/step_s, p50 | 22.9秒 | 4.8秒 |
| perf/fwd_bwd_s, p50 | 21.9秒 | 4.6秒 |
| perf/mfu_fwd_bwd | 3.9% | 23.5% |
| batch/samples_per_step | 128 | 168 |
| 学習に使ったサンプル数 | 64,000 | 84,078 |
| perf/weight_sync_s, p50 | 8.5秒 | 6.2秒 |
| sample/staleness_mean | 1.5 | 2.0 |
| reward(最初の20→最後の20ステップ) | 0.145 → 0.438 | 0.145 → 0.416 |
最終的なrunは3.9倍高速化され、31%多いサンプルで学習しながらも、ほぼ同じrewardカーブを描きました。マイクロバッチの詰め込み、gradient checkpointingの無効化、in-flight上限の引き上げがこの差を生んだ要因であり、いずれの場合もダッシュボードは最初の10分以内にその状況を明らかにしていました。
試すには
git clone https://github.com/AmineDiro/hfjobs-lora-buckets && cd hfjobs-lora-buckets
hf auth login
MAX_STEPS=20 RUN_TAG=smoke ./run_all.sh --wait
# ~15 min, three Jobs, cancels the servers when done
MAX_STEPS=500 ./run_all.sh --wait
# run 1: the reference batch shape, ~3.5 h
TOKEN_BUDGET=16384 GRAD_ACCUM=6 GRADIENT_CHECKPOINTING=0 PROXY_LORA_RETRY_S=0.5 \
MAX_INFLIGHT=384 QUEUE_MAXSIZE=768 MAX_STEPS=500 ./run_all.sh --wait
# run 5: same recipe, ~55 min
参考文献
- John Schulman et al., LoRA Without Regret, Thinking Machines Lab, 2025年9月。policy-gradient RLにおいてrank-1のLoRAがフルファインチューニングに匹敵する根拠を論じたもの。
- TRL、AsyncGRPOTrainerとそのログメトリクス。
- TRL PR #7017:AsyncGRPOTrainer向けのPEFT/LoRA対応(アダプタのみのvLLM同期)、TRL v1.14でリリース。
- Hugging Face JobsとStorage Buckets、hf-mount。
- hf-mount-repro:30秒のネガティブキャッシュ停滞を再現する2スクリプト。
- 本記事の全runのtrackioダッシュボード。
- Penghui Qi, Zichen Liu, Xiangxin Zhou, Tianyu Pang, Chao Du, Wee Sun Lee, Min Lin, Defeating the Training-Inference Mismatch via FP16, arXiv:2510.26788, 2025。Sanityデータセット
sail/Sanity-Test-R1D-1.5BおよびLoRAレシピ(sail-sg/Precision-RL、oat/scripts/lora/bf16_grpo_tis_lora.sh)の出典。
@article{qi2025precisionrl,
title={Defeating the Training-Inference Mismatch via FP16},
author={Qi, Penghui and Liu, Zichen and Zhou, Xiangxin and Pang, Tianyu and Du, Chao and Lee, Wee Sun and Lin, Min},
journal={arXiv preprint arXiv:2510.26788},
year={2025}
}
[“TRL”, “AsyncGRPOTrainer”, “LoRA”, “Hugging Face Jobs”, “vLLM”, “強化学習”, “Storage Bucket”]