Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions nemo_rl/algorithms/async_utils/trajectory_collector.py
Original file line number Diff line number Diff line change
Expand Up @@ -547,6 +547,15 @@ def _run_collection_loop(
if not self.running:
break

# Refit and weight updates can wake this loop after validation
# pauses it. Check again before starting a batch.
if not self._manual_pause_cleared.is_set() and self.running:
with (
efficiency_span("idle/validation_pause", tracer=self._tracer),
self._efficiency_timer.time("idle/validation_pause"),
):
self._manual_pause_cleared.wait()

if not self.running:
break

Expand Down
18 changes: 10 additions & 8 deletions nemo_rl/algorithms/grpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

grpo.py:5417

1 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.py and ppo.py share the same AsyncTrajectoryCollector (trajectory_collector.py:126), the same _manual_pause_cleared/_refit_pause_cleared Events, and the same _run_collection_loop. This PR reorders grpo.py so pause.remote() runs before resume_after_refit()/set_weight_version(), closing the race here. ppo.py's async loop still has the old ordering:

ray.get(trajectory_collector.resume_after_refit.remote())   # ppo.py:2790 — wakes the collector
...
if (val_period > 0 and (step + 1) % val_period == 0) or (val_at_end and is_last_step):
    with timer.time("idle/validation"):
        ray.get(trajectory_collector.pause.remote())          # ppo.py:2804 — pause lands after the wake

The new re-check added to trajectory_collector.py only blocks the race if the caller's pause() already landed before the collector wakes — true for grpo.py now, still false for ppo.py. Not dead code either: examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.yaml runs async PPO with val_period: 1.

Action: move the should_run_validation predicate and ray.get(trajectory_collector.pause.remote()) in ppo.py to before the prepare_for_refit/resume_after_refit block, mirroring this PR's grpo.py structure — or scope this PR's description to GRPO and file a fast-follow for PPO.

@jepio jepio Sep 11, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Confirmed — I checked ppo.py at HEAD (commit 25c4b60) and the ordering bug is real, not stale:

  • ppo.py:2790 calls resume_after_refit.remote() unconditionally after refit.
  • ppo.py:2799-2804 computes the validation gate and calls pause.remote() after that resume — same ordering this PR fixes in grpo.py.
  • It's not dead code: examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.yaml sets val_period: 1 and runs async PPO through this exact path (the -automodel- variant also runs it with val_period: 1; the -single-controller variant does not exercise this path — it sets val_period: 0 with validation explicitly unsupported there).

Since this PR's AsyncTrajectoryCollector/re-check changes are shared infra, the fix mirrors cleanly: move the should_run_validation predicate + pause.remote() in ppo.py to before the resume_after_refit/set_weight_version block, same as the grpo.py reorder here.

Would you like to extend this fix to ppo.py in this PR, or should we scope the description to GRPO and track PPO as a fast-follow? Happy to help draft the ppo.py change if useful.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This pause.remote() reorder (the core of this PR's grpo.py fix) has no regression coverage.

1 action item. PR-introduced gap — not a bug in the fix itself, but the fix has no test on this side.

StubAsyncTrajectoryCollector.pause in test_grpo.py:1119-1122 is a bare MagicMock() that bypasses self._remote_method(), unlike its sibling resume_after_refit, which correctly routes through _remote_method("resume_after_refit") and is captured in the events list existing tests assert ordering against (e.g. line 1942: assert events[:3] == ["refit", "set_weight_version", "start_collection"]). Because pause isn't tracked, no test can observe whether this call now fires before refit/weight-sync — so test_collection_loop_defers_new_batch_while_manually_paused covers the trajectory_collector.py-internal re-check, but not this reorder.

Action: in test_grpo.py, route pause through _remote_method like its siblings, then add an assertion for a validation step showing "pause" precedes "refit"/"set_weight_version" in events (pattern already established at lines 1919-1942).

    @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
Expand Down Expand Up @@ -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(
[
Expand Down
47 changes: 47 additions & 0 deletions tests/unit/algorithms/test_async_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1377,6 +1377,53 @@ def test_collection_loop_marks_data_exhausted_on_natural_completion(self):
assert status["errored"] is False
assert status["running"] is False

def test_collection_loop_defers_new_batch_while_manually_paused(self):
"""Keep a weight update from bypassing a manual pause."""
collector = self.create_local_collector()
collector.running = True
processed = []
batch_processed = threading.Event()
should_pause_calls = []

def _fake_should_pause():
# Park on the first check. The weight update below wakes the loop.
should_pause_calls.append(1)
return len(should_pause_calls) == 1

collector._should_pause_for_generation_limits = _fake_should_pause

def _record_batch(batch):
processed.append(batch)
batch_processed.set()

collector._process_batch = _record_batch
collector.dataloader = [{"b": 0}]

loop_thread = threading.Thread(target=collector._collection_loop, daemon=True)
loop_thread.start()

deadline = time.time() + 5.0
while collector._generation_limit_cleared.is_set():
assert time.time() < deadline, "loop never reached generation-limit wait"
time.sleep(0.01)

collector.pause()
collector.set_weight_version(1)

assert not batch_processed.wait(0.5), (
"_process_batch ran while the collector was paused for validation"
)
assert processed == []

collector.resume()
assert batch_processed.wait(5.0), "batch not processed after resume"
assert processed == [{"b": 0}]

loop_thread.join(5.0)
assert not loop_thread.is_alive()
assert collector.data_exhausted is True
assert collector.collection_failed is False

@pytest.mark.asyncio
async def test_drain_payload_metrics_returns_collector_interval(self, monkeypatch):
collector = self.create_local_collector()
Expand Down
Loading