From ee302604536a34d8f991dd7891317e3090dbca64 Mon Sep 17 00:00:00 2001 From: Jeremi Piotrowski Date: Thu, 10 Sep 2026 17:16:52 +0200 Subject: [PATCH] fix(async): gate rollout batches on validation pause Validation paused the async collector only after refit resumed it. A collector waiting on the generation limit could therefore wake below the loop pause check and start a rollout batch during validation. Pause the collector before refit can wake it, and check the manual pause again before starting a batch. Let in-flight rollouts finish, and cover the generation-limit wakeup with a regression test. Signed-off-by: Jeremi Piotrowski --- .../async_utils/trajectory_collector.py | 9 ++++ nemo_rl/algorithms/grpo.py | 18 +++---- tests/unit/algorithms/test_async_utils.py | 47 +++++++++++++++++++ 3 files changed, 66 insertions(+), 8 deletions(-) diff --git a/nemo_rl/algorithms/async_utils/trajectory_collector.py b/nemo_rl/algorithms/async_utils/trajectory_collector.py index 35932424b69..390f7832207 100644 --- a/nemo_rl/algorithms/async_utils/trajectory_collector.py +++ b/nemo_rl/algorithms/async_utils/trajectory_collector.py @@ -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 diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index f31b01007f4..283039f60a5 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -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()) + 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( [ diff --git a/tests/unit/algorithms/test_async_utils.py b/tests/unit/algorithms/test_async_utils.py index 0f53f0f03a7..8e69034eec3 100644 --- a/tests/unit/algorithms/test_async_utils.py +++ b/tests/unit/algorithms/test_async_utils.py @@ -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()