diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index a9590172a31..f9ed696081e 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -83,6 +83,7 @@ DataPlaneCheckpointBarrier, DataPlaneCheckpointMetadata, DataPlaneMutationCut, + TQReplayGroupMetadata, TQReplayMetadataState, ) from nemo_rl.algorithms.async_utils.staleness_sampler import ( @@ -426,6 +427,7 @@ def __init__( # when Ray deserializes rollout_manager and tq_buffer separately. self._rollout_manager._tq_buffer = self._buffer self._rollout_recovery_ledger = self._rollout_manager.recovery_ledger + self._restored_replay_groups_to_regenerate: list[TQReplayGroupMetadata] = [] # Direct access, deliberately. A getattr default here reads as defensive but # buys a silent failure mode: rename or drop the field and @@ -968,6 +970,28 @@ async def _maybe_restore_replay_buffer(self) -> int: ) await self._validate_replay_inventory(buffer_state) + if self._master_config.checkpointing.get("load_replay_buffer") is False: + # Validate and load the native snapshot before discarding it. This keeps + # replay-free resume fail-closed: a mismatched/corrupt checkpoint must + # not silently turn into a different training stream. + self._restored_replay_groups_to_regenerate = list(groups) + removed = await self._buffer.remove( + list(range(restored)), remove_in_dp=True + ) + if removed != restored: + raise RuntimeError( + "replay-free resume did not discard every restored replay group: " + f"restored={restored}, removed={removed}" + ) + print( + "📦 Discarded " + f"{restored} restored replay group(s); " + "checkpointing.load_replay_buffer=false will regenerate their " + "prompts on the current policy", + flush=True, + ) + return 0 + # Each buffered group holds one _buffer_capacity permit. Restore fails # above if the saved group count exceeds current capacity. assert restored <= self._async_cfg.max_buffered_rollouts @@ -995,6 +1019,9 @@ async def _maybe_restore_rollout_recovery( f"{ROLLOUT_RECOVERY_STATE_FILENAME} exists, but the matching " "native TQ checkpoint does not advertise rollout recovery" ) + if self._restored_replay_groups_to_regenerate: + async with self._data_plane_checkpoint_barrier.mutation() as cut: + await self._queue_restored_replay_groups_for_regeneration(cut) return if not isinstance(expected_payload_sha256, str): raise TypeError( @@ -1068,6 +1095,7 @@ async def _maybe_restore_rollout_recovery( clear_unreferenced=True, ) await self._rehydrate_rollout_recovery_prompts(cut) + await self._queue_restored_replay_groups_for_regeneration(cut) self._sampler_stamps_target_steps = ( parsed_state.sampler_stamps_target_steps if parsed_state.sampler_stamps_target_steps is not None @@ -1089,6 +1117,121 @@ async def _maybe_restore_rollout_recovery( flush=True, ) + async def _queue_restored_replay_groups_for_regeneration( + self, + cut: DataPlaneMutationCut, + ) -> None: + """Convert discarded canonical replay groups into fresh prompt work.""" + groups = self._restored_replay_groups_to_regenerate + if not groups: + return + + for group in groups: + tags = group["meta"].tags or [] + prompt_indices = {tag.get("prompt_idx") for tag in tags} + if len(prompt_indices) != 1: + raise ValueError( + "replay-free resume requires one stable prompt_idx per group: " + f"group={group['group_id']!r}, values={prompt_indices!r}" + ) + prompt_index = next(iter(prompt_indices)) + if isinstance(prompt_index, bool) or not isinstance(prompt_index, int): + raise TypeError( + "replay-free resume requires integer prompt_idx tags: " + f"group={group['group_id']!r}, value={prompt_index!r}" + ) + prompt = await self._load_recovery_prompt( + group_id=group["group_id"], sample_id=str(prompt_index) + ) + restored_prompt_index = prompt.get("idx") + if ( + isinstance(restored_prompt_index, bool) + or not isinstance(restored_prompt_index, int) + or restored_prompt_index != prompt_index + ): + raise ValueError( + "replay-free resume resolved a different prompt identity: " + f"group={group['group_id']!r}, expected={prompt_index!r}, " + f"actual={restored_prompt_index!r}" + ) + self._rollout_manager.reserve_prompt_group( + cut, + prompt, + target_step=group["target_step"], + admitted=True, + ) + + print( + "📦 Queued " + f"{len(groups)} restored replay prompt group(s) for fresh generation", + flush=True, + ) + # Restored ownership must be dispatched even when future checkpoint + # saving is disabled and the native snapshot predates recovery sidecars. + self._rollout_recovery_enabled = True + self._restored_replay_groups_to_regenerate = [] + + async def _load_recovery_prompt( + self, + *, + group_id: str, + sample_id: str, + ) -> DatumSpec: + """Resolve and collate one stable dataset prompt reference.""" + try: + sample_index = int(sample_id) + except ValueError as error: + raise ValueError( + f"recovery group {group_id!r} has a non-integer " + f"dataset sample_id={sample_id!r}" + ) from error + if sample_index < 0 or str(sample_index) != sample_id: + raise ValueError( + f"recovery group {group_id!r} has a non-canonical " + f"dataset sample_id={sample_id!r}" + ) + + dataset = getattr(self._dataloader, "dataset", None) + if dataset is None: + raise RuntimeError( + "cannot restore unfinished rollouts because the dataloader does " + "not expose its source dataset" + ) + try: + dataset_prompt = await asyncio.to_thread(dataset.__getitem__, sample_index) + except (IndexError, KeyError) as error: + raise RuntimeError( + f"cannot rehydrate recovery group {group_id!r}: " + f"dataset sample_id={sample_id!r} is unavailable" + ) from error + if not isinstance(dataset_prompt, dict): + raise TypeError( + f"dataset sample_id={sample_id!r} resolved to " + f"{type(dataset_prompt).__name__}, expected a DatumSpec dictionary" + ) + + collate_fn = getattr(self._dataloader, "collate_fn", None) + if collate_fn is None: + return cast(DatumSpec, dataset_prompt) + prompt_batch = await asyncio.to_thread(collate_fn, [dataset_prompt]) + if isinstance(prompt_batch, BatchedDataDict): + if prompt_batch.size != 1: + raise ValueError( + "recovery collation must return exactly one prompt; " + f"sample_id={sample_id!r}, size={prompt_batch.size}" + ) + return cast( + DatumSpec, + {key: value[0] for key, value in prompt_batch.items()}, + ) + if isinstance(prompt_batch, dict): + return cast(DatumSpec, prompt_batch) + raise TypeError( + "recovery collation for " + f"sample_id={sample_id!r} returned " + f"{type(prompt_batch).__name__}, expected a mapping" + ) + def _validate_restored_sampler_cursor(self) -> None: """Require the sampler cursor to cover every restored target step.""" restored_target_steps = [ @@ -1140,75 +1283,15 @@ async def _rehydrate_rollout_recovery_prompts( if not groups: return - dataset = getattr(self._dataloader, "dataset", None) - if dataset is None: - raise RuntimeError( - "cannot restore unfinished rollouts because the dataloader does " - "not expose its source dataset" - ) - resolved_prompts: dict[str, DatumSpec] = {} for group in groups: sample_id = group.prompt_ref.sample_id - try: - sample_index = int(sample_id) - except ValueError as error: - raise ValueError( - f"recovery group {group.group_id!r} has a non-integer " - f"dataset sample_id={sample_id!r}" - ) from error - if sample_index < 0 or str(sample_index) != sample_id: - raise ValueError( - f"recovery group {group.group_id!r} has a non-canonical " - f"dataset sample_id={sample_id!r}" - ) - prompt = resolved_prompts.get(sample_id) if prompt is None: - try: - dataset_prompt = await asyncio.to_thread( - dataset.__getitem__, sample_index - ) - except (IndexError, KeyError) as error: - raise RuntimeError( - f"cannot rehydrate recovery group {group.group_id!r}: " - f"dataset sample_id={sample_id!r} is unavailable" - ) from error - if not isinstance(dataset_prompt, dict): - raise TypeError( - f"dataset sample_id={sample_id!r} resolved to " - f"{type(dataset_prompt).__name__}, expected a DatumSpec " - "dictionary" - ) - - # Re-run one-row collation to reconstruct the tensor scalars, - # optional fields, and multimodal wrappers expected by RolloutManager. - collate_fn = getattr(self._dataloader, "collate_fn", None) - if collate_fn is None: - prompt = dataset_prompt - else: - prompt_batch = await asyncio.to_thread( - collate_fn, - [dataset_prompt], - ) - if isinstance(prompt_batch, BatchedDataDict): - if prompt_batch.size != 1: - raise ValueError( - "recovery collation must return exactly one prompt; " - f"sample_id={sample_id!r}, size={prompt_batch.size}" - ) - prompt = {key: value[0] for key, value in prompt_batch.items()} - elif isinstance(prompt_batch, dict): - # Identity-style collators used by lightweight/custom - # dataloaders may return the DatumSpec directly. - prompt = prompt_batch - else: - raise TypeError( - "recovery collation for " - f"sample_id={sample_id!r} returned " - f"{type(prompt_batch).__name__}, expected a mapping" - ) - resolved_prompts[sample_id] = cast(DatumSpec, prompt) + prompt = await self._load_recovery_prompt( + group_id=group.group_id, sample_id=sample_id + ) + resolved_prompts[sample_id] = prompt recovery_ledger.bind_runtime_prompt( cut, group.group_id, diff --git a/nemo_rl/utils/checkpoint.py b/nemo_rl/utils/checkpoint.py index f648f97e784..bd7bd3d3a06 100644 --- a/nemo_rl/utils/checkpoint.py +++ b/nemo_rl/utils/checkpoint.py @@ -167,11 +167,11 @@ class CheckpointingConfig(TypedDict): save_data_plane (bool): Whether SingleController checkpoints include the native TQ snapshot and replay-buffer metadata. Supported by the simple and mooncake_cpu backends; no backend-specific checkpoint switch is needed. - load_replay_buffer (bool): Whether async GRPO restores replay-buffer state - when resuming from a checkpoint. Defaults to True. When False the - buffer starts empty and a frontier-aligned resume regenerates the - whole buffered window fresh instead of reusing completed (and - therefore short-rollout-biased) groups. + load_replay_buffer (bool): Whether async GRPO or SingleController restores + replay-buffer state when resuming from a checkpoint. Defaults to True. + When False, a frontier-aligned async resume or a SingleController native + TQ resume regenerates the buffered prompt groups on the current policy + instead of reusing completed groups. """ enabled: bool @@ -186,7 +186,7 @@ class CheckpointingConfig(TypedDict): pretrained_checkpoint: NotRequired[PretrainedCheckpointConfig] save_optimizer: NotRequired[bool] # Default: True save_data_plane: NotRequired[bool] - load_replay_buffer: NotRequired[bool] # Default: True (async GRPO only) + load_replay_buffer: NotRequired[bool] # Default: True _AUTOMODEL_ONLY_CHECKPOINT_FIELDS = frozenset( diff --git a/tests/unit/single_controller/test_checkpoint_borrow_restore.py b/tests/unit/single_controller/test_checkpoint_borrow_restore.py index c78d9fcff32..488756b1090 100644 --- a/tests/unit/single_controller/test_checkpoint_borrow_restore.py +++ b/tests/unit/single_controller/test_checkpoint_borrow_restore.py @@ -227,6 +227,7 @@ def _controller( ) controller._master_config = SimpleNamespace( grpo=controller._algo_cfg, + checkpointing={"load_replay_buffer": True}, token_capture=SimpleNamespace(enabled=False), ) controller._dataloader = loader @@ -242,6 +243,7 @@ def _controller( controller._current_epoch = 0 controller._sampler_stamps_target_steps = False controller._rollout_recovery_enabled = True + controller._restored_replay_groups_to_regenerate = [] controller._batch_shortfall = {} controller._batch_replacements = {} controller._batch_promotions = {} diff --git a/tests/unit/single_controller/test_checkpoint_dispatch_races.py b/tests/unit/single_controller/test_checkpoint_dispatch_races.py index a1de4318507..2e1c75194de 100644 --- a/tests/unit/single_controller/test_checkpoint_dispatch_races.py +++ b/tests/unit/single_controller/test_checkpoint_dispatch_races.py @@ -928,6 +928,7 @@ async def exercise() -> None: controller_cls = SingleControllerActor.__ray_metadata__.modified_class controller = object.__new__(controller_cls) _init_recovery_telemetry(controller, train_steps=7) + controller._restored_replay_groups_to_regenerate = [] controller._sampler = sampler controller._rollout_manager = rollout_manager controller._master_config = SimpleNamespace( @@ -1045,6 +1046,7 @@ async def exercise() -> None: controller_cls = SingleControllerActor.__ray_metadata__.modified_class controller = object.__new__(controller_cls) _init_recovery_telemetry(controller, train_steps=7) + controller._restored_replay_groups_to_regenerate = [] controller._sampler = sampler controller._rollout_manager = rollout_manager controller._master_config = SimpleNamespace( @@ -1203,6 +1205,7 @@ async def exercise() -> None: controller_cls = SingleControllerActor.__ray_metadata__.modified_class controller = object.__new__(controller_cls) controller._data_plane_checkpoint_barrier = DataPlaneCheckpointBarrier() + controller._restored_replay_groups_to_regenerate = [] controller._rollout_manager = rollout_manager controller._master_config = SimpleNamespace( token_capture=SimpleNamespace(enabled=False) diff --git a/tests/unit/single_controller/test_checkpointing.py b/tests/unit/single_controller/test_checkpointing.py index ad5aea228c5..3048a02ec6a 100644 --- a/tests/unit/single_controller/test_checkpointing.py +++ b/tests/unit/single_controller/test_checkpointing.py @@ -59,6 +59,7 @@ REPLAY_BUFFER_METADATA_STORAGE, DataPlaneCheckpointBarrier, DataPlaneCheckpointMetadata, + DataPlaneMutationCut, ) from nemo_rl.algorithms.async_utils.staleness_sampler import ( InOrderSamplerConfig, @@ -503,6 +504,7 @@ def __init__(self, events: Optional[list[str]] = None) -> None: self._tq_buffer = None self.recovery_ledger = RolloutRecoveryLedger() self._events = events + self.reserved_prompts: list[dict[str, Any]] = [] self.telemetry = { "committed_groups": 0, "committed_output_tokens": 0, @@ -516,6 +518,25 @@ def set_data_plane_checkpoint_barrier(self, barrier: Any) -> None: def set_weight_version(self, version: int) -> None: self.weight_versions.append(version) + def reserve_prompt_group( + self, + cut: Any, + input_sample: dict[str, Any], + *, + target_step: Optional[int], + admitted: bool = True, + admission_id: Optional[str] = None, + ) -> str: + del cut, admission_id + self.reserved_prompts.append( + { + "prompt": input_sample, + "target_step": target_step, + "admitted": admitted, + } + ) + return f"regenerated-{len(self.reserved_prompts)}" + def suspend_request_deadlines(self) -> None: if self._events is not None: self._events.append("suspend_deadlines") @@ -560,6 +581,7 @@ def __init__( self.load_return = load_return self.metadata_state_dict_calls: list[int] = [] self.load_calls: list[dict[str, Any]] = [] + self.remove_calls: list[dict[str, Any]] = [] self.checkpoint_barrier: Optional[DataPlaneCheckpointBarrier] = None self.training_claims: list[dict[str, Any]] = [] @@ -650,6 +672,10 @@ async def load_state_dict( ] return self.load_return + async def remove(self, idxs: list[int], remove_in_dp: bool) -> int: + self.remove_calls.append({"idxs": list(idxs), "remove_in_dp": remove_in_dp}) + return len(idxs) + # Default position sentinel the fake dataloader reports via state_dict(). _SENTINEL_DL_STATE = {"fake_position": 42} @@ -663,9 +689,16 @@ class _FakeDataloader(list): train_dataloader.pt. """ - def __init__(self, batches: Any = (), state: Optional[dict[str, Any]] = None): + def __init__( + self, + batches: Any = (), + state: Optional[dict[str, Any]] = None, + dataset: Optional[Any] = None, + ) -> None: super().__init__(batches) self._state = dict(state) if state is not None else dict(_SENTINEL_DL_STATE) + self.dataset = dataset + self.collate_fn = None def state_dict(self) -> dict[str, Any]: return dict(self._state) @@ -690,6 +723,7 @@ def _actor_master_config( data_plane_checkpoint: bool = True, rollout_checkpoint_attempt_interval_s: Optional[float] = None, token_capture_enabled: bool = False, + load_replay_buffer: bool = True, ) -> MasterConfig: """MasterConfig for in-process SingleControllerActor tests. @@ -735,6 +769,7 @@ def _actor_master_config( "save_period": save_period, "save_optimizer": save_optimizer, "save_data_plane": data_plane_checkpoint, + "load_replay_buffer": load_replay_buffer, "checkpoint_must_save_by": checkpoint_must_save_by, "ft_save_period": ft_save_period, }, @@ -3016,6 +3051,148 @@ def test_run_restores_native_tq_replay_metadata_without_payload_reput( # run()'s finally must tear the synchronizer down exactly once. assert actor._weight_synchronizer.shutdown_count == 1 + @pytest.mark.parametrize( + ("checkpointing_enabled", "save_data_plane"), + [(True, True), (False, True), (False, False)], + ) + def test_replay_free_restore_discards_rows_and_queues_original_prompt( + self, tmp_path, checkpointing_enabled, save_data_plane + ): + ckpt_dir = tmp_path / "resume_ckpt" + ckpt_dir.mkdir() + group = { + "meta": KVBatchMeta( + partition_id=_PARTITION_ID, + task_name=None, + sample_ids=["g0-0", "g0-1"], + sequence_lengths=[16, 16], + tags=[ + {"weight_version": 0, "prompt_idx": 1}, + {"weight_version": 0, "prompt_idx": 1}, + ], + ), + "start_weight": 0, + "end_weight": 0, + "target_step": 3, + "group_id": "g0", + } + envelope = {"groups": [group]} + torch.save(envelope, ckpt_dir / REPLAY_BUFFER_METADATA_FILENAME) + mc = _actor_master_config( + tmp_path, + enabled=checkpointing_enabled, + max_num_steps=0, + buffer_checkpoint=True, + data_plane_checkpoint=save_data_plane, + load_replay_buffer=False, + ) + buffer = _FakeTQBuffer(load_return=1) + + class RecordingRolloutManager(_FakeRolloutManager): + def reserve_prompt_group( + self, + cut: DataPlaneMutationCut, + input_sample: dict[str, Any], + *, + target_step: Optional[int], + admitted: bool = True, + admission_id: Optional[str] = None, + ) -> str: + super().reserve_prompt_group( + cut, + input_sample, + target_step=target_step, + admitted=admitted, + admission_id=admission_id, + ) + return self.recovery_ledger.reserve_group( + cut, + prompt_id=str(input_sample["idx"]), + prompt_payload=input_sample, + expected_generations=2, + target_step=target_step, + start_weight_version=0, + admitted=admitted, + admission_id=admission_id, + ).group_id + + rollout_manager = RecordingRolloutManager() + rollout_manager.generate_and_push = AsyncMock() + dataloader = _FakeDataloader( + dataset=[{"idx": 0}, {"idx": 1, "message_log": ["fresh"]}] + ) + + async def _restore() -> Any: + actor = _ACTOR_CLS( + mc, + _make_actor_args( + tq_buffer=buffer, + rollout_manager=rollout_manager, + dataloader=dataloader, + dp_client=_FakeDPClient(sample_ids=["g0-0", "g0-1"]), + last_checkpoint_path=str(ckpt_dir), + data_plane_checkpoint_metadata=( + _data_plane_checkpoint_metadata(group_count=1) + ), + ), + SetupTimingMetrics(), + ) + restored = await actor._maybe_restore_replay_buffer() + await actor._maybe_restore_rollout_recovery(restored_replay_groups=restored) + assert actor._buffer_capacity._value == 4 + actor._rollout_permitted.set() + try: + await asyncio.wait_for(actor._rollout_pump(), timeout=5) + finally: + actor._checkpointer.shutdown() + return actor + + actor = asyncio.run(_restore()) + + assert buffer.remove_calls == [{"idxs": [0], "remove_in_dp": True}] + rollout_manager.generate_and_push.assert_awaited_once() + dispatch = rollout_manager.generate_and_push.await_args + assert dispatch.args == ({"idx": 1, "message_log": ["fresh"]},) + assert dispatch.kwargs["target_step"] == 3 + assert dispatch.kwargs["lineage_group_id"] in { + group.group_id for group in rollout_manager.recovery_ledger.groups() + } + assert actor._buffer_capacity._value == 3 + assert rollout_manager.reserved_prompts == [ + { + "prompt": {"idx": 1, "message_log": ["fresh"]}, + "target_step": 3, + "admitted": True, + } + ] + + @pytest.mark.parametrize("dataset_idx", [99, True]) + def test_replay_free_restore_rejects_changed_prompt_identity(self, dataset_idx): + actor = object.__new__(_ACTOR_CLS) + actor._dataloader = _FakeDataloader(dataset=[{"idx": 0}, {"idx": dataset_idx}]) + actor._data_plane_checkpoint_barrier = DataPlaneCheckpointBarrier() + actor._rollout_manager = _FakeRolloutManager() + actor._restored_replay_groups_to_regenerate = [ + { + "meta": KVBatchMeta( + partition_id=_PARTITION_ID, + task_name=None, + sample_ids=["g0-0", "g0-1"], + tags=[{"prompt_idx": 1}, {"prompt_idx": 1}], + ), + "target_step": 3, + "group_id": "g0", + } + ] + + async def regenerate() -> None: + async with actor._data_plane_checkpoint_barrier.mutation() as cut: + await actor._queue_restored_replay_groups_for_regeneration(cut) + + with pytest.raises(ValueError, match="prompt"): + asyncio.run(regenerate()) + assert actor._rollout_manager.reserved_prompts == [] + def test_restored_permits_are_released_by_a_live_pump(self, tmp_path): # The restore takes one capacity permit per group; a running pump must # give them all back. Every other restore test uses max_num_steps=0,