-
Notifications
You must be signed in to change notification settings - Fork 577
fix(async): gate rollout batches on validation pause #4081
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -5406,6 +5406,16 @@ def _flush_collector_telemetry() -> None: | |
| and policy_generation.wake_carries_weight_updates() | ||
| ) | ||
|
|
||
| should_run_validation = ( | ||
| val_period > 0 | ||
| and (step + 1) >= val_start_at | ||
| and (step + 1) % val_period == 0 | ||
| ) or (val_at_end and is_last_step) | ||
| if should_run_validation: | ||
| # Stop dispatch before refit wakes the collector. This also | ||
| # separates the training and validation payload metrics. | ||
| ray.get(trajectory_collector.pause.remote()) | ||
|
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This 1 action item. PR-introduced gap — not a bug in the fix itself, but the fix has no test on this side.
Action: in @property
def pause(self):
"""Pause collection - returns a remote-callable mock"""
return self._remote_method("pause") |
||
|
|
||
| print("🔄 Synchronizing policy weights to trajectory collector…") | ||
| if defer_wake_for_save: | ||
| # Wake-deferral (checkpoint scheduling, which the backend | ||
|
|
@@ -5473,17 +5483,9 @@ def _flush_collector_telemetry() -> None: | |
|
|
||
| # Validation | ||
| val_metrics, validation_timings = None, None | ||
| should_run_validation = ( | ||
| val_period > 0 | ||
| and (step + 1) >= val_start_at | ||
| and (step + 1) % val_period == 0 | ||
| ) or (val_at_end and is_last_step) | ||
|
|
||
| payload_metrics: dict[str, int | float] = {} | ||
| if should_run_validation: | ||
| # Stop new dispatch before separating the training and | ||
| # validation payload-metric intervals. | ||
| ray.get(trajectory_collector.pause.remote()) | ||
| if master_config.grpo.debug_payload_metrics: | ||
| payload_metrics = merge_multimodal_payload_metrics( | ||
| [ | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
grpo.py:54171 action item. Pre-existing bug, not introduced by this PR — but the PR description describes the fix generally ("prevents async training rollouts from starting during validation"), and the same race is still open in PPO's async path.
grpo.pyandppo.pyshare the sameAsyncTrajectoryCollector(trajectory_collector.py:126), the same_manual_pause_cleared/_refit_pause_clearedEvents, and the same_run_collection_loop. This PR reordersgrpo.pysopause.remote()runs beforeresume_after_refit()/set_weight_version(), closing the race here.ppo.py's async loop still has the old ordering:The new re-check added to
trajectory_collector.pyonly blocks the race if the caller'spause()already landed before the collector wakes — true forgrpo.pynow, still false forppo.py. Not dead code either:examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.yamlruns async PPO withval_period: 1.Action: move the
should_run_validationpredicate andray.get(trajectory_collector.pause.remote())inppo.pyto before theprepare_for_refit/resume_after_refitblock, mirroring this PR'sgrpo.pystructure — or scope this PR's description to GRPO and file a fast-follow for PPO.Uh oh!
There was an error while loading. Please reload this page.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Confirmed — I checked
ppo.pyat HEAD (commit25c4b60) and the ordering bug is real, not stale:ppo.py:2790callsresume_after_refit.remote()unconditionally after refit.ppo.py:2799-2804computes the validation gate and callspause.remote()after that resume — same ordering this PR fixes ingrpo.py.examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.yamlsetsval_period: 1and runs async PPO through this exact path (the-automodel-variant also runs it withval_period: 1; the-single-controllervariant does not exercise this path — it setsval_period: 0with validation explicitly unsupported there).Since this PR's
AsyncTrajectoryCollector/re-check changes are shared infra, the fix mirrors cleanly: move theshould_run_validationpredicate +pause.remote()inppo.pyto before theresume_after_refit/set_weight_versionblock, same as thegrpo.pyreorder here.Would you like to extend this fix to
ppo.pyin this PR, or should we scope the description to GRPO and track PPO as a fast-follow? Happy to help draft theppo.pychange if useful.