From 77bc2ecf795eae6c6ecd5d437c81fc1bb137ee08 Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Wed, 15 Jul 2026 08:51:22 -0700 Subject: [PATCH 01/11] =?UTF-8?q?feat(sc):=20rollout=20path=20=E2=80=94=20?= =?UTF-8?q?TQReplayBuffer,=20rollout=5Fmanager,=20=5Frollout=5Fpump?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Yuki Huang --- .../algorithms/async_utils/replay_buffer.py | 219 +++-- nemo_rl/algorithms/single_controller.py | 135 +-- nemo_rl/experience/payload.py | 117 +++ nemo_rl/experience/rollout_manager.py | 88 +- tests/unit/algorithms/test_async_utils.py | 159 +--- tests/unit/experience/test_rollout_manager.py | 799 ++++++++++++++++++ tests/unit/experience/test_rollouts.py | 555 +----------- tests/unit/single_controller/__init__.py | 13 + .../single_controller/test_rollout_pump.py | 317 +++++++ .../test_tq_replay_buffer.py | 324 +++++++ 10 files changed, 1858 insertions(+), 868 deletions(-) create mode 100644 nemo_rl/experience/payload.py create mode 100644 tests/unit/experience/test_rollout_manager.py create mode 100644 tests/unit/single_controller/test_rollout_pump.py create mode 100644 tests/unit/single_controller/test_tq_replay_buffer.py diff --git a/nemo_rl/algorithms/async_utils/replay_buffer.py b/nemo_rl/algorithms/async_utils/replay_buffer.py index b6dc16d236..3264ed7722 100644 --- a/nemo_rl/algorithms/async_utils/replay_buffer.py +++ b/nemo_rl/algorithms/async_utils/replay_buffer.py @@ -12,14 +12,20 @@ # See the License for the specific language governing permissions and # limitations under the License. +import asyncio import statistics import threading as _threading +import uuid from collections import Counter +from collections.abc import Mapping from typing import Any, Iterable, Optional import ray from nemo_rl.algorithms.async_utils.interfaces import ReplayBufferProtocol +from nemo_rl.data_plane import KVBatchMeta +from nemo_rl.experience.interfaces import PromptGroupRecord +from nemo_rl.experience.payload import pack_payload, record_to_train_batch # Classes with @ray.remote can't be inherited from, so we split the implementation out. @@ -629,93 +635,164 @@ class ReplayBuffer(ReplayBufferImpl): pass -# WIP: DO NOT USE - This class is WIP and may be changed without notice, please DO NOT USE it. -# Will be replaced by TQReplayBuffer once TQ is ready. -@ray.remote # pragma: no cover -class ReplayBufferNew(ReplayBufferImpl): - """Staleness-window replay buffer. - - -- WIP: DO NOT USE -- - This class is WIP and may be changed without notice, please DO NOT USE it. - - Differences from ReplayBuffer: - - _evict(): Stale rows (trainer_version - weight_version > max_staleness) are evicted - at the start of every sample() call. - - sample(): selects trajectories in freshest-first order (default) or FIFO order, - controlled by the sample_freshest_first flag, from whatever remains in the buffer - after eviction. - - TODO: remove when cleaning up - - max_age_steps won't be used in ReplayBufferNew; - - self.target_weight_versions won't be used in ReplayBufferNew and will be removed - when cleaning up. target_weight_versions gates generation on specific trainer steps, - which causes generation pauses; ReplayBufferNew intentionally avoids this. - - add this class to nemo_rl/algorithms/async_utils/__init__.py +class TQReplayBuffer: + """Meta cache + TQ writer with reserve-then-commit slot semantics. + + meta_list, weight_list, ready_list, _group_ids are parallel; a slot stays + ready=False until commit fills it. """ def __init__( - self, max_size: int, max_staleness: int, sample_freshest_first: bool = True + self, + dp_client: Any, + partition_id: str, + *, + pad_value_dict: Mapping[str, int], ): - super().__init__(max_size) - if max_staleness < 0: - raise ValueError(f"max_staleness must be non-negative, got {max_staleness}") - self.max_staleness = max_staleness - # will move to StalenessSampler when we implement it - self.sample_freshest_first = sample_freshest_first + self._dp_client = dp_client + self._partition_id = partition_id + self._pad_value_dict = dict(pad_value_dict) + self.meta_list: list[Optional[KVBatchMeta]] = [] + self.start_weight_list: list[int] = [] + self.end_weight_list: list[int] = [] + # Per-slot target training step (set when force_in_order=True, else None). + self.target_step_list: list[Optional[int]] = [] + self.ready_list: list[bool] = [] + self._group_ids: list[str] = [] + + def reserve( + self, + *, + weight_version: int, + target_step: Optional[int] = None, + group_id: Optional[str] = None, + ) -> str: + """Append an unready slot tagged with weight_version. - def _evict(self, current_weight_version: int) -> None: - """Evict rows where trainer_version - weight_version > max_staleness. + Args: + weight_version: Weight version stamped on the slot. + target_step: Training step this slot targets; only consulted by StalenessSampler.force_in_order. + group_id: Per-group sample_id prefix; defaults to a fresh uuid4. - Must be called with self._lock held. + Returns: + group_id used by the matching commit. """ - min_valid = current_weight_version - self.max_staleness - stale = [i for i, v in enumerate(self.trajectory_versions) if v < min_valid] - self._remove_indices(stale) - - def sample( + if group_id is None: + group_id = str(uuid.uuid4()) + self.meta_list.append(None) + self.start_weight_list.append(weight_version) + self.end_weight_list.append(-1) + self.target_step_list.append(target_step) + self.ready_list.append(False) + self._group_ids.append(group_id) + return group_id + + async def commit( self, - num_prompt_groups: int, - current_weight_version: int, - max_age_steps: int, - ) -> Optional[dict[str, Any]]: - """Sample num_prompt_groups trajectories, freshest-first. + group_id: str, + record: PromptGroupRecord, + start_weight_version: int, + end_weight_version: int, + ) -> KVBatchMeta: + """Tensorize record, write N rows to TQ, and mark the slot ready. - Will evict stale rows before sampling, so we will get [current_weight_version - self.max_staleness, current_weight_version] valid trajectories. + Args: + group_id: group_id returned by the matching reserve call. + record: PromptGroupRecord to tensorize. + start_weight_version: Weight version stamped on the slot before rollout. + The same as the one from reserve, passed again to avoid race condition when lookup. + end_weight_version: Weight version stamped on the slot after rollout. Returns: - Dictionary with 'trajectories' and 'avg_trajectory_age' keys, or None. + KVBatchMeta for the committed group. + + Raises: + ValueError: group_id has no live slot (removed or never reserved). """ - with self._lock: - self._evict(current_weight_version) + train_batch = record_to_train_batch(record, pad_value_dict=self._pad_value_dict) + sample_ids, fields, tags = pack_payload( + train_batch, weight_version=start_weight_version, group_id=group_id + ) + await self._call_dp( + "put_samples", + sample_ids=sample_ids, + partition_id=self._partition_id, + fields=fields, + tags=tags, + ) - if not self.trajectories: - return None + # mirrors kv_first_write + lengths = train_batch["input_lengths"] + meta = KVBatchMeta( + partition_id=self._partition_id, + task_name="train", + sample_ids=list(sample_ids), + fields=list(fields.keys()), + sequence_lengths=[int(s) for s in lengths.tolist()], + tags=[dict(t) for t in tags], + ) - all_indices = range(len(self.trajectory_versions)) - if self.sample_freshest_first: - all_indices = sorted( - all_indices, - key=lambda i: self.trajectory_versions[i], - reverse=True, - ) + idx = self._group_ids.index(group_id) + self.meta_list[idx] = meta + self.end_weight_list[idx] = end_weight_version + self.ready_list[idx] = True + return meta - if len(all_indices) < num_prompt_groups: - print( - f"Insufficient trajectories: have {len(all_indices)}, " - f"need {num_prompt_groups}. Waiting." - ) - return None + async def remove(self, idxs: list[int], remove_in_dp: bool) -> int: + """Drop entries at the given indices and optionally clear them from DataPlane. - selected = all_indices[:num_prompt_groups] - sampled_weights = [self.trajectory_versions[i] for i in selected] - avg_trajectory_age = current_weight_version - sum(sampled_weights) / len( - sampled_weights + Args: + idxs: Entry indices to drop. Must be within [0, size). + remove_in_dp: If True, also clear the dropped rows from DataPlane. + + Returns: + Number of group entries removed from the buffer. + """ + if len(idxs) == 0: + return 0 + + drop_idxs = sorted(idxs, reverse=True) + if drop_idxs[0] >= len(self.meta_list): + raise IndexError( + f"TQReplayBuffer.remove: indices out of range: {drop_idxs[0]}; " + f"size={len(self.meta_list)}" ) - sampled_items = [self.trajectories[i] for i in selected] - self._remove_indices(selected) + dropped_sample_ids: list[str] = [] + for i in drop_idxs: + meta = self.meta_list[i] + if meta is not None: + dropped_sample_ids.extend(meta.sample_ids) + del self.meta_list[i] + del self.start_weight_list[i] + del self.end_weight_list[i] + del self.target_step_list[i] + del self.ready_list[i] + del self._group_ids[i] + + if remove_in_dp: + await self._call_dp( + "clear_samples", + sample_ids=dropped_sample_ids, + partition_id=self._partition_id, + ) - return { - "trajectories": sampled_items, - "avg_trajectory_age": avg_trajectory_age, - } + return len(drop_idxs) + + def size(self) -> int: + """Return the number of prompt-group entries currently held.""" + return len(self.meta_list) + + def __len__(self) -> int: + return len(self.meta_list) + + async def _call_dp(self, method_name: str, **kwargs: Any) -> Any: + """Call a DataPlaneClient method, awaiting Ray remotes if needed.""" + method = getattr(self._dp_client, method_name) + remote = getattr(method, "remote", None) + if remote is not None: + return await remote(**kwargs) + result = method(**kwargs) + if asyncio.iscoroutine(result): + return await result + return result diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index 6b94452c77..44763c71bb 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -397,74 +397,83 @@ async def _call_dp(self, method_name: str, **kwargs) -> Any: # ── the three pumps + the inline advantage stage ─────────────────────── async def _rollout_pump(self) -> None: - """Dispatch prompts as concurrent coroutines, one per prompt group. + """Continuously dispatch rollout tasks until cancellation. - Flow per prompt: + Per batch (over_sampling=False): + 0. Wait while _max_rollout_version >= trainer_version + max_staleness, + then claim the next step by incrementing _max_rollout_version. + + Per prompt: 1. Acquire _buffer_capacity slot (backpressure) - 2. Wait for _rollout_permitted (paused during weight sync) - 3. Call gen.generate_and_push(prompt, dp_client) — RPC to GenWorker - GenWorker generates and calls DataPlane put_samples directly - 4. Decrement _inflight_rollouts + 2. Acquire sem (cap concurrent in-flight rollouts) + 3. Wait for _rollout_permitted (paused during weight sync) + 4. Call rollout_manager.generate_and_push(prompt) — local async + RolloutManager reserves a slot, runs the rollout, then commits the + group via TQReplayBuffer (→ dp_client.put_samples + mark ready) + 5. Decrement _inflight_rollouts """ - n = self._cfg.max_rollout_prompts - max_epochs = self._cfg.max_num_epochs - sem = asyncio.Semaphore(self._cfg.max_inflight_prompts) - - start = time.monotonic() - print(f"rollout_pump: dispatching {n} prompts", flush=True) - - async def _one_group(prompt: str) -> None: - await self._buffer_capacity.acquire() - await self._rollout_permitted.wait() - async with sem: - self._inflight_rollouts += 1 - try: - await self._ray_get( - self._gen.generate_and_push.remote(prompt, self._dp_client) - ) - if self._cfg.diagnostics: - print( - f" rollout done for prompt='{prompt[:20]}...'", - flush=True, - ) - finally: - self._inflight_rollouts -= 1 - - dispatched = 0 - if max_epochs is None: - # Unbounded-epoch path: max_rollout_prompts alone caps dispatch, - # all prompts in flight together (cycling through the list). - tasks = [ - asyncio.ensure_future(_one_group(self._prompts[i % len(self._prompts)])) - for i in range(n) - ] - await asyncio.gather(*tasks) - dispatched = n - else: - # Epoch-bounded path: one gather per pass over the prompt list, - # mirroring grpo.py's per-epoch dataset iteration. Ends at - # whichever bound is hit first (epochs or total prompt budget). - while dispatched < n and self._current_epoch < max_epochs: - k = min(len(self._prompts), n - dispatched) - tasks = [ - asyncio.ensure_future(_one_group(self._prompts[i])) - for i in range(k) - ] - await asyncio.gather(*tasks) - dispatched += k - self._current_epoch += 1 - print( - f"rollout_pump: epoch {self._current_epoch}/{max_epochs} " - f"complete ({dispatched}/{n} prompts)", - flush=True, + sem = asyncio.Semaphore(self._async_cfg.max_inflight_prompts) + over_sampling = self._async_cfg.over_sampling + max_staleness = self._async_cfg.max_weight_staleness_versions + force_in_order = self._async_cfg.force_in_order + print("rollout_pump: starting", flush=True) + + async def _dispatch_one_prompt( + prompt: DatumSpec, target_step: Optional[int] + ) -> None: + self._inflight_rollouts += 1 + try: + await self._rollout_manager.generate_and_push( + prompt, target_step=target_step ) + if self._diagnostics: + content = "" + for i in range(len(prompt["message_log"])): + if prompt["message_log"][i]["role"] == "user": + content = prompt["message_log"][i]["content"] + break + print(f" rollout done for prompt='{content[:20]}...'", flush=True) + finally: + self._inflight_rollouts -= 1 + sem.release() + + max_epochs = self._master_config.grpo["max_num_epochs"] + epoch = 0 + while max_epochs is None or epoch < max_epochs: + for prompt_batch in self._dataloader: + # over_sampling=False: batch-level gate on max_rollout_version. + if not over_sampling: + while ( + self._max_rollout_version + >= self._trainer_version + max_staleness + ): + await asyncio.sleep(0.005) + self._max_rollout_version += 1 + + # target_step = batch dispatch index when force_in_order is on. + target_step = self._max_rollout_version if force_in_order else None + + for prompt_idx in range(prompt_batch.size): + prompt: DatumSpec = { # type: ignore + k: v[prompt_idx] for k, v in prompt_batch.items() + } + + # check if buffer is full + await self._buffer_capacity.acquire() + # check if inflight rollouts is full + await sem.acquire() + # wait for rollout to be permitted + await self._rollout_permitted.wait() + + # dispatch rollout + task = asyncio.create_task( + _dispatch_one_prompt(prompt, target_step) + ) + self._dispatched_rollouts.add(task) + task.add_done_callback(self._dispatched_rollouts.discard) + epoch += 1 - self._rollout_done = True - print( - f"rollout_pump: finished {dispatched} prompts in " - f"{time.monotonic() - start:.2f}s", - flush=True, - ) + print(f"rollout_pump: completed {epoch} epoch(s)", flush=True) async def _train_pump(self) -> None: """Per-prompt-group streaming train loop. diff --git a/nemo_rl/experience/payload.py b/nemo_rl/experience/payload.py new file mode 100644 index 0000000000..bdf95ab290 --- /dev/null +++ b/nemo_rl/experience/payload.py @@ -0,0 +1,117 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Producer-side payload helpers for the async-RL TQ path.""" + +from collections.abc import Mapping +from typing import Any + +import numpy as np +import torch +from tensordict import TensorDict + +from nemo_rl.data_plane.codec import pack_jagged_fields +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.experience.interfaces import PromptGroupRecord + + +def record_to_train_batch( + record: PromptGroupRecord, + *, + pad_value_dict: Mapping[str, int], +) -> BatchedDataDict[Any]: + """Convert one prompt group's record into a packed BatchedDataDict of N rows. + + Args: + record: Rollout's PromptGroupRecord with N completions to flatten into rows. + pad_value_dict: Field-name → pad value used by batched_message_log_to_flat_message. + + Returns: + BatchedDataDict with input_ids, input_lengths, generation_logprobs, token_mask, + sample_mask, prompt_ids_for_adv, and total_reward. + """ + # Lazy imports: grpo and llm_message_utils transitively pull + # experience.rollouts, so importing at module top risks a cycle. + from nemo_rl.algorithms.grpo import ( + add_grpo_token_loss_masks_and_generation_logprobs, + extract_initial_prompt_messages, + ) + from nemo_rl.data.llm_message_utils import batched_message_log_to_flat_message + + completions = record.completions + n = len(completions) + assert n > 0, "PromptGroupRecord has no completions" + + message_logs = [c.message_log for c in completions] + prompt_token_count = sum(len(m["token_ids"]) for m in record.prompt) + prompt_lengths = torch.full((n,), prompt_token_count, dtype=torch.long) + + prompt_message_logs = extract_initial_prompt_messages(message_logs, prompt_lengths) + prompt_flat, _ = batched_message_log_to_flat_message( + prompt_message_logs, + pad_value_dict=dict(pad_value_dict), # type: ignore + ) + + add_grpo_token_loss_masks_and_generation_logprobs(message_logs) + flat, input_lengths = batched_message_log_to_flat_message( + message_logs, # type: ignore + pad_value_dict=dict(pad_value_dict), # type: ignore + ) + + total_reward = torch.tensor( + [float(c.reward) for c in completions], dtype=torch.float32 + ) + sample_mask = torch.ones(n, dtype=torch.float32) + + return BatchedDataDict[Any]( + { + "input_ids": flat["token_ids"], + "input_lengths": input_lengths, + "generation_logprobs": flat["generation_logprobs"], + "token_mask": flat["token_loss_mask"], + "sample_mask": sample_mask, + "prompt_ids_for_adv": prompt_flat["token_ids"], + "total_reward": total_reward, + } + ) + + +def pack_payload( + train_batch: Mapping[str, Any], + *, + weight_version: int, + group_id: str, +) -> tuple[list[str], TensorDict, list[dict[str, Any]]]: + """Pack a producer batch into (sample_ids, fields, tags) for put_samples. + + Args: + train_batch: Mapping with at least input_lengths plus the tensor/object fields to send. + weight_version: Trainer weight version stamped on every row's tag. + group_id: Per-group identifier used as the sample_id prefix; the caller owns uniqueness. + + Returns: + sample_ids of the form {group_id}_g{i}, a jagged-packed TensorDict, and per-row tags. + """ + lengths = train_batch["input_lengths"] + n = int(lengths.shape[0]) + tensor_fields: dict[str, torch.Tensor | np.ndarray] = { + k: v + for k, v in train_batch.items() + if isinstance(v, torch.Tensor) + or (isinstance(v, np.ndarray) and v.dtype == object) + } + fields_td = pack_jagged_fields(tensor_fields, lengths=lengths) + sample_ids = [f"{group_id}_g{i}" for i in range(n)] + tags = [{"weight_version": weight_version} for _ in range(n)] + return sample_ids, fields_td, tags diff --git a/nemo_rl/experience/rollout_manager.py b/nemo_rl/experience/rollout_manager.py index 234ab881f8..6bf2773828 100644 --- a/nemo_rl/experience/rollout_manager.py +++ b/nemo_rl/experience/rollout_manager.py @@ -21,7 +21,8 @@ from transformers import PreTrainedTokenizerBase from wandb import Table -from nemo_rl.data.interfaces import DatumSpec +from nemo_rl.algorithms.async_utils.replay_buffer import TQReplayBuffer +from nemo_rl.data.interfaces import DatumSpec, LLMMessageLogType from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.environments.interfaces import EnvironmentInterface from nemo_rl.experience.interfaces import Completion, PromptGroupRecord @@ -47,7 +48,7 @@ class AsyncRolloutImpl: def __init__( self, tokenizer: TokenizerType, - task_to_env: dict[str, EnvironmentInterface], + env_handles: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, policy_generation: GenerationInterface, @@ -55,7 +56,7 @@ def __init__( **kwargs: Any, ) -> None: self._tokenizer = tokenizer - self._task_to_env = task_to_env + self._env_handles = env_handles self._num_generations_per_prompt = num_generations_per_prompt self._max_seq_len = max_seq_len self._max_rollout_turns = max_rollout_turns @@ -188,7 +189,7 @@ async def _run_single_rollout( # step. In this case, need to wrap with asyncio.to_thread to make # this function yieldable. env_output = await asyncio.to_thread( - calculate_rewards, sample_batch, self._task_to_env + calculate_rewards, sample_batch, self._env_handles ) # Update reward and termination statistics @@ -399,15 +400,15 @@ class AsyncNemoGymRolloutImpl: def __init__( self, tokenizer: TokenizerType, - task_to_env: dict[str, EnvironmentInterface], + env_handles: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, generation_config: GenerationConfig, - max_rollout_turns: Optional[int] = None, + max_rollout_turns: int, **kwargs: Any, ) -> None: self._tokenizer = tokenizer - self._task_to_env = task_to_env + self._env_handles = env_handles self._num_generations_per_prompt = num_generations_per_prompt self._max_seq_len = max_seq_len self._max_rollout_turns = max_rollout_turns @@ -429,7 +430,7 @@ async def run_rollout(self, input_sample: DatumSpec) -> PromptGroupRecord: timer.start(f"{timer_prefix}/total") rollout_inputs = self._build_inputs(input_sample) - completions, rollout_metrics = await self._run_rollouts( + completions, prompt_message_log, rollout_metrics = await self._run_rollouts( rollout_inputs, timer, timer_prefix ) @@ -438,7 +439,7 @@ async def run_rollout(self, input_sample: DatumSpec) -> PromptGroupRecord: return PromptGroupRecord( prompt_idx=input_sample["idx"], - prompt=input_sample["message_log"], + prompt=prompt_message_log, extra_env_info=input_sample["extra_env_info"], metadata={"task_name": "nemo_gym"}, completions=completions, @@ -454,8 +455,9 @@ def _validate_init_params(self) -> None: ) # Validate max_rollout_turns. - assert self._max_rollout_turns is None, ( - "`max_rollout_turns` is not supported in NeMo-Gym path!" + assert self._max_rollout_turns == 1, ( + "`max_rollout_turns` is not supported in NeMo-Gym path! " + "Please set `max_rollout_turns` to 1." ) def _build_inputs(self, input_sample: DatumSpec) -> list[dict]: @@ -488,9 +490,9 @@ def _build_inputs(self, input_sample: DatumSpec) -> list[dict]: async def _run_rollouts( self, inputs: list[dict], timer: Timer, timer_prefix: str - ) -> tuple[list[Completion], dict[str, Any]]: - """Dispatch rows to NeMo-Gym and return completions + metrics.""" - nemo_gym_env = self._task_to_env["nemo_gym"] + ) -> tuple[list[Completion], LLMMessageLogType, dict[str, Any]]: + """Dispatch rows to NeMo-Gym; return completions, prompt, and metrics.""" + nemo_gym_env = self._env_handles["nemo_gym"] # Run generation and restore input order as results stream back. with timer.time(f"{timer_prefix}/run_rollouts"): @@ -517,11 +519,14 @@ async def _run_rollouts( raise RuntimeError( "NeMo-Gym rollout stream ended before all rows arrived" ) + + completed_results = [result for result in results if result is not None] + # All N rollouts share the same input prompt; tensorize one copy. + prompt_message_log = completed_results[0]["input_message_log"] + _tensorize_by_key(prompt_message_log, "token_ids") # Convert results to completions. completions = [ - self._result_to_completion(result) - for result in results - if result is not None + self._result_to_completion(result) for result in completed_results ] # Compute rollout metrics. @@ -532,12 +537,11 @@ async def _run_rollouts( rollout_metrics.update(env_timing_metrics) - return completions, rollout_metrics + return completions, prompt_message_log, rollout_metrics def _result_to_completion(self, result: dict) -> Completion: """Convert one run_rollouts result dict into a Completion.""" # Tensorize token fields. - _tensorize_by_key(result["input_message_log"], "token_ids") _tensorize_by_key(result["message_log"], "token_ids") _tensorize_by_key( [m for m in result["message_log"] if m["role"] == "assistant"], @@ -635,18 +639,19 @@ def _compute_rollout_metrics( class RolloutManager: - """Factory that routes to AsyncRolloutImpl (native async) or AsyncNemoGymRolloutImpl (NeMo-Gym).""" + """Routes to AsyncRolloutImpl (native async) or AsyncNemoGymRolloutImpl (NeMo-Gym), and pushes results to a TQReplayBuffer.""" def __init__( self, tokenizer: TokenizerType, - task_to_env: dict[str, EnvironmentInterface], + env_handles: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, max_rollout_turns: Optional[int] = None, policy_generation: Optional[GenerationInterface] = None, generation_config: Optional[GenerationConfig] = None, use_nemo_gym: bool = False, + tq_buffer: Optional[TQReplayBuffer] = None, ) -> None: assert num_generations_per_prompt >= 1, ( "num_generations_per_prompt must be >= 1" @@ -667,13 +672,52 @@ def __init__( self._impl: AsyncRolloutImpl | AsyncNemoGymRolloutImpl = rollout_cls( tokenizer=tokenizer, - task_to_env=task_to_env, + env_handles=env_handles, num_generations_per_prompt=num_generations_per_prompt, max_seq_len=max_seq_len, max_rollout_turns=max_rollout_turns, # type: ignore policy_generation=policy_generation, # type: ignore generation_config=generation_config, ) + self._tokenizer = tokenizer + self._num_generations_per_prompt = num_generations_per_prompt + self._tq_buffer = tq_buffer + self._weight_version: int = 0 + + def set_weight_version(self, version: int) -> None: + """Set the weight_version used for rollout tags. + + Args: + version: Trainer weight version to stamp on future rollout tags. + """ + self._weight_version = int(version) async def run_rollout(self, input_sample: DatumSpec) -> PromptGroupRecord: return await self._impl.run_rollout(input_sample) + + async def generate_and_push( + self, input_sample: DatumSpec, *, target_step: Optional[int] = None + ) -> None: + """Reserve a buffer slot, run one prompt's rollout, then commit the slot. + + Args: + input_sample: A single prompt (one DatumSpec entry). + target_step: Training step this rollout targets; stamped on the buffer slot for StalenessSampler.force_in_order. + """ + assert self._tq_buffer is not None, ( + "generate_and_push requires tq_buffer to be set at __init__" + ) + start_version = self._weight_version + group_id = self._tq_buffer.reserve( + weight_version=start_version, target_step=target_step + ) + + record = await self.run_rollout(input_sample) + end_version = self._weight_version + + await self._tq_buffer.commit( + group_id, + record, + start_weight_version=start_version, + end_weight_version=end_version, + ) diff --git a/tests/unit/algorithms/test_async_utils.py b/tests/unit/algorithms/test_async_utils.py index 099328ddbb..aecf5ff1eb 100644 --- a/tests/unit/algorithms/test_async_utils.py +++ b/tests/unit/algorithms/test_async_utils.py @@ -36,10 +36,7 @@ AsyncTrajectoryCollector, ReplayBuffer, ) -from nemo_rl.algorithms.async_utils.replay_buffer import ( - ReplayBufferImpl, - ReplayBufferNew, -) +from nemo_rl.algorithms.async_utils.replay_buffer import ReplayBufferImpl from nemo_rl.algorithms.grpo import ( MasterConfig, _get_next_nemo_gym_task_index, @@ -1043,160 +1040,6 @@ def test_replay_buffer_checkpoint_with_torch_save(self): ray.kill(buffer2) -class TestReplayBufferNew: - """Tests for ReplayBufferNew: staleness-window sampling via _evict + sample.""" - - def _make_traj(self, label: str) -> dict: - return {"batch": {"data": label}, "rollout_metrics": {}} - - def _add(self, buf, label: str, weight_version: int): - return ray.get( - buf.add.remote( - self._make_traj(label), - weight_version=weight_version, - target_weight_version=0, # unused in ReplayBufferNew - ) - ) - - def _sample(self, buf, num_groups: int, trainer_version: int): - return ray.get( - buf.sample.remote( - num_prompt_groups=num_groups, - current_weight_version=trainer_version, - max_age_steps=0, # unused in ReplayBufferNew - ) - ) - - # ------------------------------------------------------------------ - # Construction - # ------------------------------------------------------------------ - - def test_invalid_max_staleness_raises(self): - with pytest.raises(Exception): - buf = ReplayBufferNew.remote(max_size=10, max_staleness=-1) - ray.get(buf.size.remote()) - - # ------------------------------------------------------------------ - # _evict (via sample) - # ------------------------------------------------------------------ - - def test_stale_rows_evicted_before_sampling(self): - """Rows with age > max_staleness are removed before sample() selects.""" - buf = ReplayBufferNew.remote(max_size=10, max_staleness=2) - # age at trainer=4: gen_v=1 → 3 > 2 (stale), gen_v=3 → 1 ≤ 2 (valid) - self._add(buf, "stale", weight_version=1) - self._add(buf, "fresh", weight_version=3) - - result = self._sample(buf, num_groups=1, trainer_version=4) - - assert result is not None - assert result["trajectories"][0]["batch"]["data"] == "fresh" - assert ray.get(buf.size.remote()) == 0 # stale row also gone - ray.kill(buf) - - def test_all_stale_returns_none(self): - """sample() returns None when all rows are evicted as stale.""" - buf = ReplayBufferNew.remote(max_size=10, max_staleness=1) - self._add(buf, "a", weight_version=0) - self._add(buf, "b", weight_version=1) - - # trainer=5: both ages > 1 - result = self._sample(buf, num_groups=1, trainer_version=5) - - assert result is None - assert ray.get(buf.size.remote()) == 0 - ray.kill(buf) - - def test_eviction_frees_capacity(self): - """Evicting a stale row allows a subsequent add() to succeed.""" - buf = ReplayBufferNew.remote(max_size=1, max_staleness=1) - self._add(buf, "x", weight_version=1) - assert self._add(buf, "x", weight_version=1) == "full" - - # sample() at trainer=5 evicts the stale row (age 4 > 1) - self._sample(buf, num_groups=1, trainer_version=5) - - assert self._add(buf, "y", weight_version=4) == "success" - ray.kill(buf) - - def test_within_window_not_evicted(self): - """Rows whose age is within max_staleness are not evicted.""" - buf = ReplayBufferNew.remote(max_size=10, max_staleness=3) - self._add(buf, "x", weight_version=4) - - # trainer=6: age = 6 - 4 = 2 ≤ 3 → should survive - # should return None since there is only 1 row - result = self._sample(buf, num_groups=2, trainer_version=6) - assert result is None - - # this sample should still be there - assert ray.get(buf.size.remote()) == 1 - ray.kill(buf) - - # ------------------------------------------------------------------ - # sample() - # ------------------------------------------------------------------ - - @pytest.mark.parametrize("sample_freshest_first", [True, False]) - def test_sample_freshest_first(self, sample_freshest_first): - """sample() returns the freshest trajectories first.""" - buf = ReplayBufferNew.remote( - max_size=10, max_staleness=5, sample_freshest_first=sample_freshest_first - ) - for gen_v in [3, 4, 5]: - self._add(buf, f"v{gen_v}", weight_version=gen_v) - - result = self._sample(buf, num_groups=2, trainer_version=6) - - assert result is not None - data = [t["batch"]["data"] for t in result["trajectories"]] - if sample_freshest_first: - assert data == ["v5", "v4"] - else: - assert data == ["v3", "v4"] - ray.kill(buf) - - def test_sample_returns_none_when_insufficient(self): - """sample() returns None when fewer rows than requested remain after eviction.""" - buf = ReplayBufferNew.remote(max_size=10, max_staleness=5) - self._add(buf, "only", weight_version=1) - - result = self._sample(buf, num_groups=3, trainer_version=2) - - assert result is None - ray.kill(buf) - - def test_sample_returns_none_on_empty_buffer(self): - buf = ReplayBufferNew.remote(max_size=10, max_staleness=5) - result = self._sample(buf, num_groups=1, trainer_version=1) - assert result is None - ray.kill(buf) - - def test_sample_avg_trajectory_age(self): - """avg_trajectory_age is computed from the sampled generation versions.""" - buf = ReplayBufferNew.remote(max_size=10, max_staleness=5) - # freshest first: gen 8 (age 2), gen 6 (age 4) → avg = 3.0 - for gen_v in [6, 8]: - self._add(buf, f"v{gen_v}", weight_version=gen_v) - - result = self._sample(buf, num_groups=2, trainer_version=10) - - assert result is not None - assert abs(result["avg_trajectory_age"] - 3.0) < 1e-6 - ray.kill(buf) - - def test_sample_consumes_selected_rows(self): - """Rows returned by sample() are removed from the buffer.""" - buf = ReplayBufferNew.remote(max_size=10, max_staleness=5) - for gen_v in [1, 2, 3]: - self._add(buf, f"v{gen_v}", weight_version=gen_v) - - self._sample(buf, num_groups=2, trainer_version=4) - - assert ray.get(buf.size.remote()) == 1 - ray.kill(buf) - - class TestAsyncTrajectoryCollector: """Test cases for AsyncTrajectoryCollector.""" diff --git a/tests/unit/experience/test_rollout_manager.py b/tests/unit/experience/test_rollout_manager.py new file mode 100644 index 0000000000..63a590e6a1 --- /dev/null +++ b/tests/unit/experience/test_rollout_manager.py @@ -0,0 +1,799 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for RolloutManager. + +Two groups: + +* TestGenerateAndPushFlow — lightweight unit tests for the reserve→run→commit + flow in generate_and_push (no Ray/vLLM; fakes for impl + tq_buffer). +* AsyncRollout / AsyncNemoGymRollout tests — vLLM/Ray-backed end-to-end checks + for the underlying run_rollout paths (AsyncRolloutImpl / AsyncNemoGymRolloutImpl). +""" + +from __future__ import annotations + +import asyncio +import json +import tempfile +import uuid +from copy import deepcopy + +import pytest +import torch + +from nemo_rl.data.collate_fn import rl_collate_fn +from nemo_rl.data.datasets.response_datasets import NemoGymDataset +from nemo_rl.data.interfaces import DatumSpec +from nemo_rl.data.processors import nemo_gym_data_processor +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.experience.interfaces import Completion, PromptGroupRecord +from nemo_rl.experience.rollout_manager import RolloutManager +from nemo_rl.experience.rollouts import ( + run_async_multi_turn_rollout, + run_async_nemo_gym_rollout, +) + +# Fixtures shared with the heavyweight rollout tests. +from tests.unit.environments.test_nemo_gym import ( + cluster, # noqa: F401 + nemo_gym, # noqa: F401 + nemo_gym_sanity_test_data, # noqa: F401 + nemo_gym_tokenizer, # noqa: F401 + nemo_gym_vllm_generation, # noqa: F401 +) +from tests.unit.experience.test_rollouts import ( + initial_multi_step_calculator_batch, # noqa: F401 + multi_step_calculator_environment, # noqa: F401 + multi_step_setup_vllm_async, # noqa: F401 + rollout_cluster, # noqa: F401 + rollout_tokenizer, # noqa: F401 +) +from tests.unit.test_envs import MultiStepCalcMetadata + + +def _run(coro): + return asyncio.run(coro) + + +class _FakeBuffer: + """Minimal TQReplayBuffer stand-in that records reserve/commit calls.""" + + def __init__(self) -> None: + self.reserve_calls: list[int] = [] # weight_versions passed to reserve + self.commit_calls: list[tuple[str, object, int, int]] = [] + # reserve(weight_version=X) -> group_id; commit fills the slot. + self._slots: list[str] = [] + + def reserve(self, *, weight_version: int, group_id: str | None = None) -> str: + if group_id is None: + group_id = str(uuid.uuid4()) + self.reserve_calls.append(weight_version) + self._slots.append(group_id) + return group_id + + async def commit( + self, + group_id: str, + record, + start_weight_version: int, + end_weight_version: int, + ): + self.commit_calls.append( + (group_id, record, start_weight_version, end_weight_version) + ) + return record + + +class _FakeImpl: + """Stand-in for AsyncRolloutImpl that returns a sentinel record.""" + + def __init__(self, record="sentinel-record", on_run=None) -> None: + self._record = record + self._on_run = on_run + + async def run_rollout(self, input_sample): + if self._on_run is not None: + await self._on_run(input_sample) + return self._record + + +def _make_manager(buffer: _FakeBuffer, impl: _FakeImpl) -> RolloutManager: + """Build a RolloutManager without firing the real __init__.""" + mgr = object.__new__(RolloutManager) + mgr._impl = impl + mgr._tokenizer = None + mgr._num_generations_per_prompt = 1 + mgr._tq_buffer = buffer + mgr._weight_version = 0 + return mgr + + +class TestGenerateAndPushFlow: + def test_reserves_then_runs_then_commits(self): + events: list[str] = [] + buf = _FakeBuffer() + + async def _track_run(_sample): + events.append("run") + + impl = _FakeImpl(record="r0", on_run=_track_run) + mgr = _make_manager(buf, impl) + + # Wrap reserve/commit to log ordering. + original_reserve = buf.reserve + original_commit = buf.commit + + def _logged_reserve(**kwargs): + events.append("reserve") + return original_reserve(**kwargs) + + async def _logged_commit(*args, **kwargs): + events.append("commit") + return await original_commit(*args, **kwargs) + + buf.reserve = _logged_reserve # type: ignore[method-assign] + buf.commit = _logged_commit # type: ignore[method-assign] + + _run(mgr.generate_and_push({"prompt": "p"})) + + assert events == ["reserve", "run", "commit"] + assert buf.reserve_calls == [0] + assert len(buf.commit_calls) == 1 + gid, record, start_v, end_v = buf.commit_calls[0] + assert gid in buf._slots + assert record == "r0" + assert start_v == 0 + assert end_v == 0 + + def test_start_weight_version_pinned_at_reserve_time(self): + """If set_weight_version is called mid-rollout, start != end.""" + buf = _FakeBuffer() + + async def _bump_weight_mid_rollout(_sample): + # Simulate a sync_weights bump during the rollout. + mgr.set_weight_version(5) + + impl = _FakeImpl(record="r0", on_run=_bump_weight_mid_rollout) + mgr = _make_manager(buf, impl) + mgr.set_weight_version(3) + + _run(mgr.generate_and_push({"prompt": "p"})) + + # reserve happened before run_rollout → captured weight 3. + assert buf.reserve_calls == [3] + # commit's start is the same dispatch-time value; end reflects the post-rollout weight. + _, _, start_v, end_v = buf.commit_calls[0] + assert start_v == 3 + assert end_v == 5 + + def test_no_weight_change_means_start_equals_end(self): + buf = _FakeBuffer() + impl = _FakeImpl(record="r0") + mgr = _make_manager(buf, impl) + mgr.set_weight_version(7) + + _run(mgr.generate_and_push({"prompt": "p"})) + + _, _, start_v, end_v = buf.commit_calls[0] + assert start_v == 7 + assert end_v == 7 + + def test_concurrent_dispatch_preserves_reserve_order(self): + """Two concurrent generate_and_push calls must reserve before either commits. + + The contract: reserve order == dispatch order, even if rollouts finish + out of order. Slot order in the buffer reflects the order reserve was + called (not the order run_rollout completed). + """ + buf = _FakeBuffer() + + # First call's rollout blocks until second call has reserved. + first_reserved = asyncio.Event() + second_reserved = asyncio.Event() + + async def _first_run(_sample): + first_reserved.set() + await second_reserved.wait() + + async def _second_run(_sample): + # Second is dispatched only after first reserves, so by the time + # second's reserve fires, slots[0] == first's gid. + second_reserved.set() + + first_impl = _FakeImpl(record="r0", on_run=_first_run) + second_impl = _FakeImpl(record="r1", on_run=_second_run) + + first_mgr = _make_manager(buf, first_impl) + # Share buffer across two managers (mimics two dispatches from one pump). + second_mgr = object.__new__(RolloutManager) + second_mgr._impl = second_impl + second_mgr._tokenizer = None + second_mgr._num_generations_per_prompt = 1 + second_mgr._tq_buffer = buf + second_mgr._weight_version = 0 + + async def _drive(): + t1 = asyncio.create_task(first_mgr.generate_and_push({"prompt": "p1"})) + # Wait until first has reserved before kicking off second so the + # reserve ordering is deterministic. + await first_reserved.wait() + t2 = asyncio.create_task(second_mgr.generate_and_push({"prompt": "p2"})) + await asyncio.gather(t1, t2) + + _run(_drive()) + + # Slots in buffer == reserve order. + first_gid, second_gid = buf._slots + # Commit recorded both, in either order, but each maps to its own gid. + commit_gids = [c[0] for c in buf.commit_calls] + assert set(commit_gids) == {first_gid, second_gid} + assert buf.reserve_calls == [0, 0] + + def test_requires_tq_buffer(self): + mgr = _make_manager(_FakeBuffer(), _FakeImpl()) + mgr._tq_buffer = None + with pytest.raises(AssertionError, match="tq_buffer"): + _run(mgr.generate_and_push({"prompt": "p"})) + + +# --------------------------------------------------------------------------- +# Tests for RolloutManager +# --------------------------------------------------------------------------- + + +def test_rollout_manager_raises_without_impl_params(): + """RolloutManager raises AssertionError when required params are missing.""" + common = { + "tokenizer": None, + "env_handles": {}, + "num_generations_per_prompt": 1, + "max_seq_len": 1, + } + + with pytest.raises(AssertionError, match="num_generations_per_prompt must be >= 1"): + updated_common = common.copy() + updated_common["num_generations_per_prompt"] = 0 + RolloutManager(**updated_common, use_nemo_gym=False) + + with pytest.raises(AssertionError, match="policy_generation is required"): + RolloutManager(**common, use_nemo_gym=False) + + with pytest.raises(AssertionError, match="generation_config is required"): + RolloutManager(**common, use_nemo_gym=True) + + +# --------------------------------------------------------------------------- +# Tests for AsyncRolloutManager (native async path) +# --------------------------------------------------------------------------- + + +@pytest.fixture(scope="function") +def single_multi_step_calculator_input_sample(rollout_tokenizer): # noqa: F811 + """Returns a single DatumSpec prompt dict (problem 0) for AsyncRolloutManager tests.""" + problem_text = "(5 + 3) * 2" + expected_answer = 16.0 + max_steps = 5 + + tool_instructions = ( + "You have a calculator tool. To use it, respond with:\n" + "'[operand1, operand2, operation_name]'\n" + "The valid 'operation_name' values are exactly: 'sum', 'diff', 'prod', 'div'.\n" + "Example: [5, 3, sum]\n" + "You will receive the result of your calculation as ...\n" + "Use this result to make the next calculation if needed.\n" + "IMPORTANT: Only perform one calculation step (one tool call) before waiting for a result and making a new tool call.\n" + "IMPORTANT: Do not perform any other calculations or operations aside from the tool call and result. Doing so will result in failure.\n" + "To give the final answer, just output the number. numbers inside of don't count, so output just the final number yourself outside of this.\n" + "Example full output: [2, 4, sum]\n6.0\n[6, 6, diff]\n0.0 0\n(note how you have to output the final 0 outside of the tags)" + "------\n" + f"Solve: {problem_text}" + ) + + initial_prompt_content = rollout_tokenizer.apply_chat_template( + [{"role": "user", "content": tool_instructions}], + tokenize=False, + add_system_prompt=False, + add_generation_prompt=True, + add_special_tokens=False, + ) + tokenized_prompt = rollout_tokenizer( + initial_prompt_content, return_tensors="pt", add_special_tokens=False + )["input_ids"][0] + message_log = [ + { + "role": "user", + "content": initial_prompt_content, + "token_ids": tokenized_prompt, + } + ] + metadata = MultiStepCalcMetadata( + problem=problem_text, + expected_final_answer=expected_answer, + max_steps=max_steps, + current_step=0, + ) + return { + "message_log": message_log, + "extra_env_info": metadata, + "task_name": "multi_step_calculator_game", + "stop_strings": [""], + "idx": 0, + } + + +@pytest.mark.vllm +def test_async_rollout_manager( + multi_step_setup_vllm_async, # noqa: F811 + single_multi_step_calculator_input_sample, +): + """Standalone test for AsyncRolloutManager. + + Given 1 prompt with num_generations_per_prompt=N, asserts: + - output is a PromptGroupRecord with N Completion objects + - each Completion has a reward (float) and a non-empty message_log + - rollout_metrics has the expected keys with correct types + - completions hold independent (not aliased) message_log objects + """ + vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + input_sample = single_multi_step_calculator_input_sample + num_generations = 2 + max_seq_len = 1024 + max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 + + manager = RolloutManager( + use_nemo_gym=False, + tokenizer=tokenizer, + env_handles=env_handles, + num_generations_per_prompt=num_generations, + max_seq_len=max_seq_len, + max_rollout_turns=max_rollout_turns, + policy_generation=vllm_generation, + ) + + vllm_generation.prepare_for_generation() + record = asyncio.run(manager.run_rollout(input_sample)) + vllm_generation.finish_generation() + + assert isinstance(record, PromptGroupRecord) + assert len(record.completions) == num_generations, ( + f"Expected {num_generations} completions, got {len(record.completions)}" + ) + assert record.prompt_idx == input_sample["idx"] + + for i, completion in enumerate(record.completions): + assert isinstance(completion, Completion) + + # 1. message_log length + assert len(completion.message_log) >= 4, ( + f"Completion {i}: expected >= 4 messages, got {len(completion.message_log)}" + ) + + # 2. last assistant content + last_assistant = next( + (m for m in reversed(completion.message_log) if m["role"] == "assistant"), + None, + ) + assert last_assistant is not None, f"Completion {i}: no assistant message found" + assert last_assistant["content"].strip() == "16", ( + f"Completion {i}: last assistant content {last_assistant['content']!r} != '16'" + ) + + # 3. reward + assert completion.reward == 1.0, ( + f"Completion {i}: reward {completion.reward} != 1.0" + ) + + # completions must be independent objects + assert record.completions[0].message_log is not record.completions[1].message_log + + +@pytest.mark.vllm +def test_async_rollout_manager_truncation( + multi_step_setup_vllm_async, # noqa: F811 + single_multi_step_calculator_input_sample, +): + """Small max_seq_len forces truncation and truncation_rate=1.0.""" + vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + input_sample = single_multi_step_calculator_input_sample + num_generations = 2 + max_seq_len = 290 + max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 + + manager = RolloutManager( + use_nemo_gym=False, + tokenizer=tokenizer, + env_handles=env_handles, + num_generations_per_prompt=num_generations, + max_seq_len=max_seq_len, + max_rollout_turns=max_rollout_turns, + policy_generation=vllm_generation, + ) + vllm_generation.prepare_for_generation() + record = asyncio.run(manager.run_rollout(input_sample)) + vllm_generation.finish_generation() + + assert len(record.completions) == num_generations + assert all(c.truncated for c in record.completions) + assert record.rollout_metrics["truncation_rate"] == 1.0 + assert record.rollout_metrics["natural_termination_rate"] == 0.0 + + +@pytest.mark.vllm +def test_async_rollout_manager_matches_original( + multi_step_setup_vllm_async, # noqa: F811 + single_multi_step_calculator_input_sample, +): + """Comparison test: AsyncRolloutManager output is structurally equivalent to the original. + + Calls run_async_multi_turn_rollout with a batch of N identical prompts, + then calls AsyncRolloutManager with 1 prompt and N generations. + Asserts that both produce N results with matching message-log depth, rewards, + and rollout_metrics numeric values. + + TODO: remove this test together with run_async_multi_turn_rollout when the legacy path is deleted. + """ + vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + input_sample = single_multi_step_calculator_input_sample + num_generations = 2 + max_seq_len = 1024 + max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 + + # Build a batch of N identical prompts for the original function + batch = BatchedDataDict( + { + "message_log": [ + deepcopy(input_sample["message_log"]) for _ in range(num_generations) + ], + "extra_env_info": [ + deepcopy(input_sample["extra_env_info"]) for _ in range(num_generations) + ], + "task_name": [input_sample["task_name"]] * num_generations, + "stop_strings": [input_sample["stop_strings"]] * num_generations, + "idx": list(range(num_generations)), + "loss_multiplier": [1.0] * num_generations, + } + ) + + vllm_generation.prepare_for_generation() + original_batch, original_metrics = run_async_multi_turn_rollout( + policy_generation=vllm_generation, + input_batch=batch, + tokenizer=tokenizer, + task_to_env=env_handles, + max_seq_len=max_seq_len, + max_rollout_turns=max_rollout_turns, + ) + + manager = RolloutManager( + use_nemo_gym=False, + tokenizer=tokenizer, + env_handles=env_handles, + num_generations_per_prompt=num_generations, + max_seq_len=max_seq_len, + max_rollout_turns=max_rollout_turns, + policy_generation=vllm_generation, + ) + record = asyncio.run(manager.run_rollout(input_sample)) + vllm_generation.finish_generation() + + # Both should produce N results + assert len(original_batch["message_log"]) == num_generations + assert len(record.completions) == num_generations + + for i in range(num_generations): + orig_msg_log = original_batch["message_log"][i] + new_msg_log = record.completions[i].message_log + + # 1. message_log length matches + assert len(orig_msg_log) == len(new_msg_log), ( + f"Completion {i}: message_log length {len(new_msg_log)} != original {len(orig_msg_log)}" + ) + + # 2. last assistant content matches + def _last_assistant_content(msg_log): + for m in reversed(msg_log): + if m["role"] == "assistant": + return m.get("content", "") + return "" + + orig_last = _last_assistant_content(orig_msg_log) + new_last = _last_assistant_content(new_msg_log) + assert orig_last == new_last, ( + f"Completion {i}: last assistant content mismatch\n" + f" original: {orig_last!r}\n" + f" manager: {new_last!r}" + ) + + # 3. reward matches + orig_reward = original_batch["total_reward"][i].item() + new_reward = record.completions[i].reward + assert orig_reward == new_reward, ( + f"Completion {i}: reward mismatch — original {orig_reward}, manager {new_reward}" + ) + + # 4. rollout_metrics numeric values match (timing and histogram fields are excluded). + # The new impl emits slash-style keys (X/mean, X/max, X/min) via calculate_single_metric; + # translate the legacy prefix-style keys before comparing. + def _translate_legacy_key(key: str) -> str: + if key == "avg_turns_per_sample": + return "turns_per_sample/mean" + if key == "max_turns_reached_rate": + return key + # Keys already in slash-style (e.g. turns_per_sample/p95, max_gen_tokens_per_turn/max) + # are new-style and should not be re-translated by the prefix-strip logic. + if "/" in key: + return key + for prefix, suffix in (("mean_", "/mean"), ("max_", "/max"), ("min_", "/min")): + if key.startswith(prefix): + return f"{key[len(prefix) :]}{suffix}" + return key + + new_metrics = record.rollout_metrics + for key in original_metrics.keys(): + if key.startswith("timing/") or key.startswith("histogram/"): + continue + + new_key = _translate_legacy_key(key) + assert new_key in new_metrics, ( + f"rollout_metrics[{new_key!r}] missing from manager" + ) + + orig_val = original_metrics[key] + new_val = new_metrics[new_key] + + assert type(orig_val) == type(new_val), ( + f"rollout_metrics[{key!r}] type mismatch: {type(orig_val)} != {type(new_val)}" + ) + if not isinstance(orig_val, (bool, int, float)): + continue + + assert orig_val == pytest.approx(new_val), ( + f"rollout_metrics[{key!r}] mismatch — original {orig_val}, manager {new_val}" + ) + + +# --------------------------------------------------------------------------- +# Tests for AsyncNemoGymRolloutManager +# --------------------------------------------------------------------------- + + +@pytest.mark.nemo_gym +def test_async_nemo_gym_rollout_manager( + nemo_gym, # noqa: F811 + nemo_gym_vllm_generation, # noqa: F811 + nemo_gym_sanity_test_data, # noqa: F811 + nemo_gym_tokenizer, # noqa: F811 +): + """Standalone test for AsyncNemoGymRolloutManager. + + Given 1 prompt with num_generations_per_prompt=N, asserts: + - output is a PromptGroupRecord with N Completion objects + - each Completion has a reward (float) and a non-empty message_log + - completions hold independent message_log objects + + If the result here does not match, please check the following: + 1. Test data changed: re-run test_nemo_gym_sanity (tests/unit/environments/test_nemo_gym.py) + and use _write_actual_test_data output to refresh test_nemo_gym_sanity.json. + 2. Logic changed: inspect recent changes to AsyncNemoGymRolloutManager or the gym env. + """ + with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: + for data in nemo_gym_sanity_test_data["input"]: + f.write(json.dumps(data) + "\n") + data_path = f.name + + dataset = NemoGymDataset(data_path) + examples = [ + nemo_gym_data_processor(dataset.dataset[idx], None, None, None, idx) + for idx in range(len(dataset.dataset)) + ] + input_batch: BatchedDataDict[DatumSpec] = rl_collate_fn(examples) + + # Use only the first prompt + single_prompt = { + "message_log": input_batch["message_log"][0], + "extra_env_info": input_batch["extra_env_info"][0], + "task_name": "nemo_gym", + "idx": 0, + "loss_multiplier": float(input_batch["loss_multiplier"][0]), + } + num_generations = 2 + + manager = RolloutManager( + use_nemo_gym=True, + tokenizer=nemo_gym_tokenizer, + env_handles={"nemo_gym": nemo_gym}, + num_generations_per_prompt=num_generations, + max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], + generation_config=nemo_gym_vllm_generation.cfg, + ) + record = asyncio.run(manager.run_rollout(single_prompt)) + + assert isinstance(record, PromptGroupRecord) + assert len(record.completions) == num_generations, ( + f"Expected {num_generations} completions, got {len(record.completions)}" + ) + assert record.prompt_idx == 0 + + for i, completion in enumerate(record.completions): + assert isinstance(completion, Completion) + + # 1. message_log length + assert len(completion.message_log) == 2, ( + f"Completion {i}: expected 2 messages, got {len(completion.message_log)}" + ) + + # 2. last assistant token_ids + last_assistant = next( + (m for m in reversed(completion.message_log) if m["role"] == "assistant"), + None, + ) + assert last_assistant is not None, f"Completion {i}: no assistant message found" + assert torch.equal( + last_assistant["token_ids"], + torch.tensor([151667, 198, 32313, 11, 1077]), + ), ( + f"Completion {i}: last assistant token_ids {last_assistant['token_ids'].tolist()} " + f"!= [151667, 198, 32313, 11, 1077]" + ) + + # 3. reward + assert completion.reward == 0.0, ( + f"Completion {i}: reward {completion.reward} != 0.0" + ) + + # completions must be independent objects + assert record.completions[0].message_log is not record.completions[1].message_log + + +@pytest.mark.nemo_gym +def test_async_nemo_gym_rollout_manager_matches_original( + nemo_gym, # noqa: F811 + nemo_gym_vllm_generation, # noqa: F811 + nemo_gym_sanity_test_data, # noqa: F811 + nemo_gym_tokenizer, # noqa: F811 +): + """Comparison test: AsyncNemoGymRolloutManager output is structurally equivalent to the original. + + Calls run_async_nemo_gym_rollout with a batch of N identical rows, + then calls AsyncNemoGymRolloutManager with 1 prompt, N generations. + Asserts that both produce N results and rewards are in the same numeric domain. + + TODO: remove this test together with run_async_nemo_gym_rollout when the legacy path is deleted. + """ + with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: + for data in nemo_gym_sanity_test_data["input"]: + f.write(json.dumps(data) + "\n") + data_path = f.name + + dataset = NemoGymDataset(data_path) + examples = [ + nemo_gym_data_processor(dataset.dataset[idx], None, None, None, idx) + for idx in range(len(dataset.dataset)) + ] + input_batch: BatchedDataDict[DatumSpec] = rl_collate_fn(examples) + + num_generations = 2 + single_prompt = { + "message_log": input_batch["message_log"][0], + "extra_env_info": input_batch["extra_env_info"][0], + "task_name": "nemo_gym", + "idx": 0, + "loss_multiplier": float(input_batch["loss_multiplier"][0]), + } + + # Build a batch of N identical rows for the original function + repeated_batch = BatchedDataDict( + { + "message_log": [ + deepcopy(input_batch["message_log"][0]) for _ in range(num_generations) + ], + "extra_env_info": [ + deepcopy(input_batch["extra_env_info"][0]) + for _ in range(num_generations) + ], + "loss_multiplier": input_batch["loss_multiplier"][0:1].repeat( + num_generations + ), + "idx": list(range(num_generations)), + "task_name": ["nemo_gym"] * num_generations, + } + ) + + original_result = run_async_nemo_gym_rollout( + policy_generation=nemo_gym_vllm_generation, + input_batch=repeated_batch, + tokenizer=nemo_gym_tokenizer, + task_to_env={"nemo_gym": nemo_gym}, + generation_config=nemo_gym_vllm_generation.cfg, + max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], + max_rollout_turns=None, + ) + + manager = RolloutManager( + use_nemo_gym=True, + tokenizer=nemo_gym_tokenizer, + env_handles={"nemo_gym": nemo_gym}, + num_generations_per_prompt=num_generations, + max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], + generation_config=nemo_gym_vllm_generation.cfg, + ) + record = asyncio.run(manager.run_rollout(single_prompt)) + + # Both should produce N completions + assert len(original_result.final_batch["message_log"]) == num_generations + assert len(record.completions) == num_generations + + for i in range(num_generations): + orig_msg_log = original_result.final_batch["message_log"][i] + new_msg_log = record.completions[i].message_log + + # 1. message_log length matches + assert len(orig_msg_log) == len(new_msg_log), ( + f"Completion {i}: message_log length {len(new_msg_log)} != original {len(orig_msg_log)}" + ) + + # 2. last assistant token_ids match + def _last_assistant_token_ids(msg_log): + for m in reversed(msg_log): + if m["role"] == "assistant": + return m.get("token_ids") + return None + + orig_token_ids = _last_assistant_token_ids(orig_msg_log) + new_token_ids = _last_assistant_token_ids(new_msg_log) + assert orig_token_ids is not None, ( + f"Completion {i}: no assistant message in original" + ) + assert new_token_ids is not None, ( + f"Completion {i}: no assistant message in manager" + ) + assert torch.equal(orig_token_ids, new_token_ids), ( + f"Completion {i}: last assistant token_ids mismatch\n" + f" original: {orig_token_ids.tolist()}\n" + f" manager: {new_token_ids.tolist()}" + ) + + # 3. reward matches + orig_reward = original_result.final_batch["total_reward"][i].item() + new_reward = record.completions[i].reward + assert orig_reward == new_reward, ( + f"Completion {i}: reward mismatch — original {orig_reward}, manager {new_reward}" + ) + + # 4. rollout_metrics numeric values match (timing and Table fields are excluded) + orig_metrics = original_result.rollout_metrics + new_metrics = record.rollout_metrics + for key in orig_metrics.keys(): + # Skip timing and full_result fields + if key.startswith("timing/") or key.endswith("/full_result"): + continue + + # Check that the key is present in the new metrics + assert key in new_metrics, f"rollout_metrics[{key!r}] missing from manager" + + orig_val = orig_metrics[key] + new_val = new_metrics[key] + + # Skip non-numeric fields + assert type(orig_val) == type(new_val), ( + f"rollout_metrics[{key!r}] type mismatch: {type(orig_val)} != {type(new_val)}" + ) + if not isinstance(orig_val, (bool, int, float)): + continue + + # Check equal + assert orig_val == pytest.approx(new_val), ( + f"rollout_metrics[{key!r}] mismatch — original {orig_val}, manager {new_val}" + ) diff --git a/tests/unit/experience/test_rollouts.py b/tests/unit/experience/test_rollouts.py index 4db6090daa..04d42c006f 100644 --- a/tests/unit/experience/test_rollouts.py +++ b/tests/unit/experience/test_rollouts.py @@ -38,9 +38,8 @@ SlidingPuzzleGameLogic, SlidingPuzzleMetadata, ) -from nemo_rl.experience.interfaces import Completion, PromptGroupRecord from nemo_rl.experience.metric_utils import calculate_single_metric, pct -from nemo_rl.experience.rollout_manager import AsyncNemoGymRolloutImpl, RolloutManager +from nemo_rl.experience.rollout_manager import AsyncNemoGymRolloutImpl from nemo_rl.experience.rollouts import ( generate_responses_async, run_async_multi_turn_rollout, @@ -1635,555 +1634,3 @@ def _standardize(d: dict) -> dict: 1. In nemo_rl/experience/rollouts.py::run_async_nemo_gym_rollout, the sampling params are passed appropriately 2. In nemo_rl/models/generation/vllm/vllm_worker_async.py::VllmAsyncGenerationWorker::_setup_vllm_server::create_chat_completion, the sampling params (like top_k) are set as appropriate """ - - -# --------------------------------------------------------------------------- -# Tests for RolloutManager -# --------------------------------------------------------------------------- - - -def test_rollout_manager_raises_without_impl_params(): - """RolloutManager raises AssertionError when required params are missing.""" - common = { - "tokenizer": None, - "task_to_env": {}, - "num_generations_per_prompt": 1, - "max_seq_len": 1, - } - - with pytest.raises(AssertionError, match="num_generations_per_prompt must be >= 1"): - updated_common = common.copy() - updated_common["num_generations_per_prompt"] = 0 - RolloutManager(**updated_common, use_nemo_gym=False) - - with pytest.raises(AssertionError, match="policy_generation is required"): - RolloutManager(**common, use_nemo_gym=False) - - with pytest.raises(AssertionError, match="generation_config is required"): - RolloutManager(**common, use_nemo_gym=True) - - -# --------------------------------------------------------------------------- -# Tests for AsyncRolloutManager (native async path) -# --------------------------------------------------------------------------- - - -@pytest.fixture(scope="function") -def single_multi_step_calculator_input_sample(rollout_tokenizer): - """Returns a single DatumSpec prompt dict (problem 0) for AsyncRolloutManager tests.""" - problem_text = "(5 + 3) * 2" - expected_answer = 16.0 - max_steps = 5 - - tool_instructions = ( - "You have a calculator tool. To use it, respond with:\n" - "'[operand1, operand2, operation_name]'\n" - "The valid 'operation_name' values are exactly: 'sum', 'diff', 'prod', 'div'.\n" - "Example: [5, 3, sum]\n" - "You will receive the result of your calculation as ...\n" - "Use this result to make the next calculation if needed.\n" - "IMPORTANT: Only perform one calculation step (one tool call) before waiting for a result and making a new tool call.\n" - "IMPORTANT: Do not perform any other calculations or operations aside from the tool call and result. Doing so will result in failure.\n" - "To give the final answer, just output the number. numbers inside of don't count, so output just the final number yourself outside of this.\n" - "Example full output: [2, 4, sum]\n6.0\n[6, 6, diff]\n0.0 0\n(note how you have to output the final 0 outside of the tags)" - "------\n" - f"Solve: {problem_text}" - ) - - initial_prompt_content = rollout_tokenizer.apply_chat_template( - [{"role": "user", "content": tool_instructions}], - tokenize=False, - add_system_prompt=False, - add_generation_prompt=True, - add_special_tokens=False, - ) - tokenized_prompt = rollout_tokenizer( - initial_prompt_content, return_tensors="pt", add_special_tokens=False - )["input_ids"][0] - message_log = [ - { - "role": "user", - "content": initial_prompt_content, - "token_ids": tokenized_prompt, - } - ] - metadata = MultiStepCalcMetadata( - problem=problem_text, - expected_final_answer=expected_answer, - max_steps=max_steps, - current_step=0, - ) - return { - "message_log": message_log, - "extra_env_info": metadata, - "task_name": "multi_step_calculator_game", - "stop_strings": [""], - "idx": 0, - } - - -@pytest.mark.vllm -def test_async_rollout_manager( - multi_step_setup_vllm_async, - single_multi_step_calculator_input_sample, -): - """Standalone test for AsyncRolloutManager. - - Given 1 prompt with num_generations_per_prompt=N, asserts: - - output is a PromptGroupRecord with N Completion objects - - each Completion has a reward (float) and a non-empty message_log - - rollout_metrics has the expected keys with correct types - - completions hold independent (not aliased) message_log objects - """ - vllm_generation, rollout_tokenizer, task_to_env, _, _ = multi_step_setup_vllm_async - input_sample = single_multi_step_calculator_input_sample - num_generations = 2 - max_seq_len = 1024 - max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 - - manager = RolloutManager( - use_nemo_gym=False, - tokenizer=rollout_tokenizer, - task_to_env=task_to_env, - num_generations_per_prompt=num_generations, - max_seq_len=max_seq_len, - max_rollout_turns=max_rollout_turns, - policy_generation=vllm_generation, - ) - - vllm_generation.prepare_for_generation() - record = asyncio.run(manager.run_rollout(input_sample)) - vllm_generation.finish_generation() - - assert isinstance(record, PromptGroupRecord) - assert len(record.completions) == num_generations, ( - f"Expected {num_generations} completions, got {len(record.completions)}" - ) - assert record.prompt_idx == input_sample["idx"] - - for i, completion in enumerate(record.completions): - assert isinstance(completion, Completion) - - # 1. message_log length - assert len(completion.message_log) >= 4, ( - f"Completion {i}: expected >= 4 messages, got {len(completion.message_log)}" - ) - - # 2. last assistant content - last_assistant = next( - (m for m in reversed(completion.message_log) if m["role"] == "assistant"), - None, - ) - assert last_assistant is not None, f"Completion {i}: no assistant message found" - assert last_assistant["content"].strip() == "16", ( - f"Completion {i}: last assistant content {last_assistant['content']!r} != '16'" - ) - - # 3. reward - assert completion.reward == 1.0, ( - f"Completion {i}: reward {completion.reward} != 1.0" - ) - - # completions must be independent objects - assert record.completions[0].message_log is not record.completions[1].message_log - - -@pytest.mark.vllm -def test_async_rollout_manager_truncation( - multi_step_setup_vllm_async, - single_multi_step_calculator_input_sample, -): - """Small max_seq_len forces truncation and truncation_rate=1.0.""" - vllm_generation, rollout_tokenizer, task_to_env, _, _ = multi_step_setup_vllm_async - input_sample = single_multi_step_calculator_input_sample - num_generations = 2 - max_seq_len = 290 - max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 - - manager = RolloutManager( - use_nemo_gym=False, - tokenizer=rollout_tokenizer, - task_to_env=task_to_env, - num_generations_per_prompt=num_generations, - max_seq_len=max_seq_len, - max_rollout_turns=max_rollout_turns, - policy_generation=vllm_generation, - ) - vllm_generation.prepare_for_generation() - record = asyncio.run(manager.run_rollout(input_sample)) - vllm_generation.finish_generation() - - assert len(record.completions) == num_generations - assert all(c.truncated for c in record.completions) - assert record.rollout_metrics["truncation_rate"] == 1.0 - assert record.rollout_metrics["natural_termination_rate"] == 0.0 - - -@pytest.mark.vllm -def test_async_rollout_manager_matches_original( - multi_step_setup_vllm_async, - single_multi_step_calculator_input_sample, -): - """Comparison test: AsyncRolloutManager output is structurally equivalent to the original. - - Calls run_async_multi_turn_rollout with a batch of N identical prompts, - then calls AsyncRolloutManager with 1 prompt and N generations. - Asserts that both produce N results with matching message-log depth, rewards, - and rollout_metrics numeric values. - - TODO: remove this test together with run_async_multi_turn_rollout when the legacy path is deleted. - """ - vllm_generation, rollout_tokenizer, task_to_env, _, _ = multi_step_setup_vllm_async - input_sample = single_multi_step_calculator_input_sample - num_generations = 2 - max_seq_len = 1024 - max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 - - # Build a batch of N identical prompts for the original function - batch = BatchedDataDict( - { - "message_log": [ - deepcopy(input_sample["message_log"]) for _ in range(num_generations) - ], - "extra_env_info": [ - deepcopy(input_sample["extra_env_info"]) for _ in range(num_generations) - ], - "task_name": [input_sample["task_name"]] * num_generations, - "stop_strings": [input_sample["stop_strings"]] * num_generations, - "idx": list(range(num_generations)), - "loss_multiplier": [1.0] * num_generations, - } - ) - - vllm_generation.prepare_for_generation() - original_batch, original_metrics = run_async_multi_turn_rollout( - policy_generation=vllm_generation, - input_batch=batch, - tokenizer=rollout_tokenizer, - task_to_env=task_to_env, - max_seq_len=max_seq_len, - max_rollout_turns=max_rollout_turns, - ) - - manager = RolloutManager( - use_nemo_gym=False, - tokenizer=rollout_tokenizer, - task_to_env=task_to_env, - num_generations_per_prompt=num_generations, - max_seq_len=max_seq_len, - max_rollout_turns=max_rollout_turns, - policy_generation=vllm_generation, - ) - record = asyncio.run(manager.run_rollout(input_sample)) - vllm_generation.finish_generation() - - # Both should produce N results - assert len(original_batch["message_log"]) == num_generations - assert len(record.completions) == num_generations - - for i in range(num_generations): - orig_msg_log = original_batch["message_log"][i] - new_msg_log = record.completions[i].message_log - - # 1. message_log length matches - assert len(orig_msg_log) == len(new_msg_log), ( - f"Completion {i}: message_log length {len(new_msg_log)} != original {len(orig_msg_log)}" - ) - - # 2. last assistant content matches - def _last_assistant_content(msg_log): - for m in reversed(msg_log): - if m["role"] == "assistant": - return m.get("content", "") - return "" - - orig_last = _last_assistant_content(orig_msg_log) - new_last = _last_assistant_content(new_msg_log) - assert orig_last == new_last, ( - f"Completion {i}: last assistant content mismatch\n" - f" original: {orig_last!r}\n" - f" manager: {new_last!r}" - ) - - # 3. reward matches - orig_reward = original_batch["total_reward"][i].item() - new_reward = record.completions[i].reward - assert orig_reward == new_reward, ( - f"Completion {i}: reward mismatch — original {orig_reward}, manager {new_reward}" - ) - - # 4. rollout_metrics numeric values match (timing and histogram fields are excluded). - # The new impl emits slash-style keys (X/mean, X/max, X/min) via calculate_single_metric; - # translate the legacy prefix-style keys before comparing. - def _translate_legacy_key(key: str) -> str: - if key == "avg_turns_per_sample": - return "turns_per_sample/mean" - if key == "max_turns_reached_rate": - return key - # Keys already in slash-style (e.g. turns_per_sample/p95, max_gen_tokens_per_turn/max) - # are new-style and should not be re-translated by the prefix-strip logic. - if "/" in key: - return key - for prefix, suffix in (("mean_", "/mean"), ("max_", "/max"), ("min_", "/min")): - if key.startswith(prefix): - return f"{key[len(prefix) :]}{suffix}" - return key - - new_metrics = record.rollout_metrics - for key in original_metrics.keys(): - if key.startswith("timing/") or key.startswith("histogram/"): - continue - - new_key = _translate_legacy_key(key) - assert new_key in new_metrics, ( - f"rollout_metrics[{new_key!r}] missing from manager" - ) - - orig_val = original_metrics[key] - new_val = new_metrics[new_key] - - assert type(orig_val) == type(new_val), ( - f"rollout_metrics[{key!r}] type mismatch: {type(orig_val)} != {type(new_val)}" - ) - if not isinstance(orig_val, (bool, int, float)): - continue - - assert orig_val == pytest.approx(new_val), ( - f"rollout_metrics[{key!r}] mismatch — original {orig_val}, manager {new_val}" - ) - - -# --------------------------------------------------------------------------- -# Tests for AsyncNemoGymRolloutManager -# --------------------------------------------------------------------------- - - -@pytest.mark.nemo_gym -def test_async_nemo_gym_rollout_manager( - nemo_gym, # noqa: F811 - nemo_gym_vllm_generation, # noqa: F811 - nemo_gym_sanity_test_data, # noqa: F811 - nemo_gym_tokenizer, # noqa: F811 -): - """Standalone test for AsyncNemoGymRolloutManager. - - Given 1 prompt with num_generations_per_prompt=N, asserts: - - output is a PromptGroupRecord with N Completion objects - - each Completion has a reward (float) and a non-empty message_log - - completions hold independent message_log objects - - If the result here does not match, please check the following: - 1. Test data changed: re-run test_nemo_gym_sanity (tests/unit/environments/test_nemo_gym.py) - and use _write_actual_test_data output to refresh test_nemo_gym_sanity.json. - 2. Logic changed: inspect recent changes to AsyncNemoGymRolloutManager or the gym env. - """ - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - for data in nemo_gym_sanity_test_data["input"]: - f.write(json.dumps(data) + "\n") - data_path = f.name - - dataset = NemoGymDataset(data_path) - examples = [ - nemo_gym_data_processor(dataset.dataset[idx], None, None, None, idx) - for idx in range(len(dataset.dataset)) - ] - input_batch: BatchedDataDict[DatumSpec] = rl_collate_fn(examples) - - # Use only the first prompt - single_prompt = { - "message_log": input_batch["message_log"][0], - "extra_env_info": input_batch["extra_env_info"][0], - "task_name": "nemo_gym", - "idx": 0, - "loss_multiplier": float(input_batch["loss_multiplier"][0]), - } - num_generations = 2 - - manager = RolloutManager( - use_nemo_gym=True, - tokenizer=nemo_gym_tokenizer, - task_to_env={"nemo_gym": nemo_gym}, - num_generations_per_prompt=num_generations, - max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], - generation_config=nemo_gym_vllm_generation.cfg, - ) - record = asyncio.run(manager.run_rollout(single_prompt)) - - assert isinstance(record, PromptGroupRecord) - assert len(record.completions) == num_generations, ( - f"Expected {num_generations} completions, got {len(record.completions)}" - ) - assert record.prompt_idx == 0 - - for i, completion in enumerate(record.completions): - assert isinstance(completion, Completion) - - # 1. message_log length - assert len(completion.message_log) == 2, ( - f"Completion {i}: expected 2 messages, got {len(completion.message_log)}" - ) - - # 2. last assistant token_ids - last_assistant = next( - (m for m in reversed(completion.message_log) if m["role"] == "assistant"), - None, - ) - assert last_assistant is not None, f"Completion {i}: no assistant message found" - assert torch.equal( - last_assistant["token_ids"], - torch.tensor([151667, 198, 32313, 11, 1077]), - ), ( - f"Completion {i}: last assistant token_ids {last_assistant['token_ids'].tolist()} " - f"!= [151667, 198, 32313, 11, 1077]" - ) - - # 3. reward - assert completion.reward == 0.0, ( - f"Completion {i}: reward {completion.reward} != 0.0" - ) - - # completions must be independent objects - assert record.completions[0].message_log is not record.completions[1].message_log - - -@pytest.mark.nemo_gym -def test_async_nemo_gym_rollout_manager_matches_original( - nemo_gym, # noqa: F811 - nemo_gym_vllm_generation, # noqa: F811 - nemo_gym_sanity_test_data, # noqa: F811 - nemo_gym_tokenizer, # noqa: F811 -): - """Comparison test: AsyncNemoGymRolloutManager output is structurally equivalent to the original. - - Calls run_async_nemo_gym_rollout with a batch of N identical rows, - then calls AsyncNemoGymRolloutManager with 1 prompt, N generations. - Asserts that both produce N results and rewards are in the same numeric domain. - - TODO: remove this test together with run_async_nemo_gym_rollout when the legacy path is deleted. - """ - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - for data in nemo_gym_sanity_test_data["input"]: - f.write(json.dumps(data) + "\n") - data_path = f.name - - dataset = NemoGymDataset(data_path) - examples = [ - nemo_gym_data_processor(dataset.dataset[idx], None, None, None, idx) - for idx in range(len(dataset.dataset)) - ] - input_batch: BatchedDataDict[DatumSpec] = rl_collate_fn(examples) - - num_generations = 2 - single_prompt = { - "message_log": input_batch["message_log"][0], - "extra_env_info": input_batch["extra_env_info"][0], - "task_name": "nemo_gym", - "idx": 0, - "loss_multiplier": float(input_batch["loss_multiplier"][0]), - } - - # Build a batch of N identical rows for the original function - repeated_batch = BatchedDataDict( - { - "message_log": [ - deepcopy(input_batch["message_log"][0]) for _ in range(num_generations) - ], - "extra_env_info": [ - deepcopy(input_batch["extra_env_info"][0]) - for _ in range(num_generations) - ], - "loss_multiplier": input_batch["loss_multiplier"][0:1].repeat( - num_generations - ), - "idx": list(range(num_generations)), - "task_name": ["nemo_gym"] * num_generations, - } - ) - - original_result = run_nemo_gym_rollout_sync( - policy_generation=nemo_gym_vllm_generation, - input_batch=repeated_batch, - tokenizer=nemo_gym_tokenizer, - task_to_env={"nemo_gym": nemo_gym}, - generation_config=nemo_gym_vllm_generation.cfg, - log_full_result_tables=False, - max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], - max_rollout_turns=None, - ) - - manager = RolloutManager( - use_nemo_gym=True, - tokenizer=nemo_gym_tokenizer, - task_to_env={"nemo_gym": nemo_gym}, - num_generations_per_prompt=num_generations, - max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], - generation_config=nemo_gym_vllm_generation.cfg, - ) - record = asyncio.run(manager.run_rollout(single_prompt)) - - # Both should produce N completions - assert len(original_result.final_batch["message_log"]) == num_generations - assert len(record.completions) == num_generations - - for i in range(num_generations): - orig_msg_log = original_result.final_batch["message_log"][i] - new_msg_log = record.completions[i].message_log - - # 1. message_log length matches - assert len(orig_msg_log) == len(new_msg_log), ( - f"Completion {i}: message_log length {len(new_msg_log)} != original {len(orig_msg_log)}" - ) - - # 2. last assistant token_ids match - def _last_assistant_token_ids(msg_log): - for m in reversed(msg_log): - if m["role"] == "assistant": - return m.get("token_ids") - return None - - orig_token_ids = _last_assistant_token_ids(orig_msg_log) - new_token_ids = _last_assistant_token_ids(new_msg_log) - assert orig_token_ids is not None, ( - f"Completion {i}: no assistant message in original" - ) - assert new_token_ids is not None, ( - f"Completion {i}: no assistant message in manager" - ) - assert torch.equal(orig_token_ids, new_token_ids), ( - f"Completion {i}: last assistant token_ids mismatch\n" - f" original: {orig_token_ids.tolist()}\n" - f" manager: {new_token_ids.tolist()}" - ) - - # 3. reward matches - orig_reward = original_result.final_batch["total_reward"][i].item() - new_reward = record.completions[i].reward - assert orig_reward == new_reward, ( - f"Completion {i}: reward mismatch — original {orig_reward}, manager {new_reward}" - ) - - # 4. rollout_metrics numeric values match (timing and Table fields are excluded) - orig_metrics = original_result.rollout_metrics - new_metrics = record.rollout_metrics - for key in orig_metrics.keys(): - # Skip timing and full_result fields - if key.startswith("timing/") or key.endswith("/full_result"): - continue - - # Check that the key is present in the new metrics - assert key in new_metrics, f"rollout_metrics[{key!r}] missing from manager" - - orig_val = orig_metrics[key] - new_val = new_metrics[key] - - # Skip non-numeric fields - assert type(orig_val) == type(new_val), ( - f"rollout_metrics[{key!r}] type mismatch: {type(orig_val)} != {type(new_val)}" - ) - if not isinstance(orig_val, (bool, int, float)): - continue - - # Check equal - assert orig_val == pytest.approx(new_val), ( - f"rollout_metrics[{key!r}] mismatch — original {orig_val}, manager {new_val}" - ) diff --git a/tests/unit/single_controller/__init__.py b/tests/unit/single_controller/__init__.py index e69de29bb2..4fc25d0d3c 100644 --- a/tests/unit/single_controller/__init__.py +++ b/tests/unit/single_controller/__init__.py @@ -0,0 +1,13 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. diff --git a/tests/unit/single_controller/test_rollout_pump.py b/tests/unit/single_controller/test_rollout_pump.py new file mode 100644 index 0000000000..b4d5f291b3 --- /dev/null +++ b/tests/unit/single_controller/test_rollout_pump.py @@ -0,0 +1,317 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""End-to-end test: SC._rollout_pump writes the expected rows to TQ.""" + +from __future__ import annotations + +import time +from typing import Any + +import ray +import torch +from tensordict import TensorDict + +from nemo_rl.algorithms.async_utils.replay_buffer import TQReplayBuffer +from nemo_rl.algorithms.single_controller import SingleControllerActor +from nemo_rl.algorithms.single_controller_utils import ( + AsyncRLConfig, + MasterConfig, + SingleControllerBundle, +) +from nemo_rl.data_plane.adapters.noop import NoOpDataPlaneClient +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.experience.rollout_manager import RolloutManager + +# Reuse fixtures from the experience tests; same shape as test_async_rollout_manager. +from tests.unit.experience.test_rollout_manager import ( + single_multi_step_calculator_input_sample, # noqa: F401 +) +from tests.unit.experience.test_rollouts import ( + initial_multi_step_calculator_batch, # noqa: F401 + multi_step_calculator_environment, # noqa: F401 + multi_step_setup_vllm_async, # noqa: F401 + rollout_cluster, # noqa: F401 + rollout_tokenizer, # noqa: F401 +) + +_PARTITION_ID = "rollout_data" +# TQReplayBuffer.add tensorizes each PromptGroupRecord and writes +# ``generations_per_prompt`` training rows directly to TQ. +_BULK_FIELDS = [ + "input_ids", + "input_lengths", + "generation_logprobs", + "token_mask", + "sample_mask", + "prompt_ids_for_adv", + "total_reward", +] + + +@ray.remote(num_cpus=0) +class _TQActor: + """Ray-wrapped NoOpDataPlaneClient for cross-process TQ inspection.""" + + def __init__( + self, + partition_id: str, + fields: list[str], + num_samples: int, + consumer_tasks: list[str], + ) -> None: + self._client = NoOpDataPlaneClient() + self._client.register_partition( + partition_id=partition_id, + fields=list(fields), + num_samples=int(num_samples), + consumer_tasks=list(consumer_tasks), + ) + + def put_samples( + self, + sample_ids: list[str], + partition_id: str, + fields: TensorDict | None = None, + tags: list[dict[str, Any]] | None = None, + ) -> Any: + return self._client.put_samples( + sample_ids=sample_ids, + partition_id=partition_id, + fields=fields, + tags=tags, + ) + + def claim_meta(self, **kwargs: Any) -> Any: + return self._client.claim_meta(**kwargs) + + def get_samples( + self, + sample_ids: list[str], + partition_id: str, + select_fields: list[str], + ) -> TensorDict: + return self._client.get_samples( + sample_ids=sample_ids, + partition_id=partition_id, + select_fields=list(select_fields), + ) + + def get_tags( + self, partition_id: str, sample_ids: list[str] + ) -> list[dict[str, Any]]: + rec = self._client._partitions[partition_id] + return [dict(rec.tags.get(sid, {})) for sid in sample_ids] + + def peek_count(self, partition_id: str) -> int: + return len(self._client._partitions[partition_id].rows) + + +class _SyncDPAdapter: + """Sync DataPlaneClient over a Ray actor handle. Pads nested tensors before transport.""" + + def __init__(self, handle: Any) -> None: + self._handle = handle + + def put_samples( + self, + sample_ids: list[str], + partition_id: str, + fields: TensorDict | None = None, + tags: list[dict[str, Any]] | None = None, + ) -> Any: + if fields is not None: + fields = self._padded(fields) + return ray.get( + self._handle.put_samples.remote( + sample_ids=sample_ids, + partition_id=partition_id, + fields=fields, + tags=tags, + ) + ) + + @staticmethod + def _padded(td: TensorDict) -> TensorDict: + out: dict[str, torch.Tensor] = {} + for k in td.keys(): + v = td.get(k) + if isinstance(v, torch.Tensor) and v.is_nested: + v = torch.nested.to_padded_tensor(v, padding=0) + out[k] = v + return TensorDict(out, batch_size=td.batch_size) + + +def test_rollout_pump_writes_expected_tq_data( + multi_step_setup_vllm_async, # noqa: F811 + single_multi_step_calculator_input_sample, # noqa: F811 +): + """SC._rollout_pump writes max_rollout_prompts * num_generations rows to TQ with the expected fields and tags.""" + vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + input_sample = single_multi_step_calculator_input_sample + + num_generations = 2 + max_rollout_prompts = 2 + # TQReplayBuffer.add writes ``num_generations`` training rows per prompt. + expected_samples = max_rollout_prompts * num_generations + max_seq_len = 1024 + max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 + + tq_actor = _TQActor.remote( + partition_id=_PARTITION_ID, + fields=_BULK_FIELDS, + num_samples=expected_samples * 4, + consumer_tasks=["train"], + ) + dp_adapter = _SyncDPAdapter(tq_actor) + + mc = MasterConfig.model_construct( + grpo={ + "max_num_steps": 1, + "max_num_epochs": None, + "num_generations_per_prompt": num_generations, + }, + async_rl=AsyncRLConfig( + batch_selection_strategy="strict_on_policy", + max_weight_staleness_versions=0, + min_prompt_groups_per_batch=1, + max_inflight_prompts=max_rollout_prompts, + max_buffered_rollouts=max_rollout_prompts, + ), + ) + # Wrap each value in a single-element list so size==1 and v[0] returns the original field. + batched_sample = BatchedDataDict({k: [v] for k, v in input_sample.items()}) + dataloader = [batched_sample] * max_rollout_prompts + + tq_buffer = TQReplayBuffer( + dp_adapter, + partition_id=_PARTITION_ID, + pad_value_dict={"token_ids": int(tokenizer.pad_token_id or 0)}, + ) + rollout_manager = RolloutManager( + tokenizer=tokenizer, + env_handles=env_handles, + num_generations_per_prompt=num_generations, + max_seq_len=max_seq_len, + max_rollout_turns=max_rollout_turns, + policy_generation=vllm_generation, + use_nemo_gym=False, + tq_buffer=tq_buffer, + ) + bundle = SingleControllerBundle( + gen_handle=vllm_generation, + trainer_handle=object(), + env_handles=env_handles, + train_cluster=None, + inference_cluster=None, + dp_client=dp_adapter, + dataloader=dataloader, + weight_synchronizer=object(), + advantage_estimator=None, + loss_fn=None, + rollout_manager=rollout_manager, + tq_buffer=tq_buffer, + partition_id=_PARTITION_ID, + ) + ctrl = SingleControllerActor.remote(master_config=mc, bundle=bundle) + + vllm_generation.prepare_for_generation() + + # _rollout_pump runs until cancelled, so poll TQ then cancel. + pump_ref = ctrl._rollout_pump.remote() + deadline = time.monotonic() + 120.0 + while time.monotonic() < deadline: + if ray.get(tq_actor.peek_count.remote(_PARTITION_ID)) >= expected_samples: + break + time.sleep(0.5) + assert ray.get(tq_actor.peek_count.remote(_PARTITION_ID)) >= expected_samples, ( + "rollout_pump did not push expected_samples within timeout" + ) + ray.cancel(pump_ref) + try: + ray.get(pump_ref) + except (ray.exceptions.RayTaskError, ray.exceptions.TaskCancelledError): + pass + + vllm_generation.finish_generation() + + meta = ray.get( + tq_actor.claim_meta.remote( + partition_id=_PARTITION_ID, + task_name="train", + required_fields=_BULK_FIELDS, + batch_size=expected_samples * 4, + blocking=False, + timeout_s=0.0, + ) + ) + assert meta.size == expected_samples + + # pack_payload stamps sample_ids as ``{group_uuid}_g{i}``. + group_ids: set[str] = set() + for sid in meta.sample_ids: + head, _, tail = sid.rpartition("_g") + assert head and tail.isdigit(), f"unexpected sample_id: {sid}" + group_ids.add(head) + assert len(group_ids) == max_rollout_prompts + + bulk = ray.get( + tq_actor.get_samples.remote( + sample_ids=meta.sample_ids, + partition_id=_PARTITION_ID, + select_fields=_BULK_FIELDS, + ) + ) + assert set(bulk.keys()) >= set(_BULK_FIELDS), ( + f"missing bulk fields: {set(_BULK_FIELDS) - set(bulk.keys())}" + ) + + input_lengths = bulk["input_lengths"].long() + assert input_lengths.shape[0] == expected_samples + assert torch.all(input_lengths > 0) + assert torch.allclose( + bulk["sample_mask"].float(), + torch.ones(expected_samples, dtype=torch.float32), + ) + + # Same deterministic prompt as test_async_rollout_manager: the model + # solves the calculator task every time -> reward == 1.0 and decoded + # tail contains " 16". + rewards = bulk["total_reward"].float().flatten() + assert rewards.shape == (expected_samples,) + assert torch.allclose(rewards, torch.ones(expected_samples)), ( + f"expected all rewards == 1.0, got {rewards.tolist()}" + ) + + input_ids = bulk["input_ids"] + token_mask = bulk["token_mask"] + for i in range(expected_samples): + length = int(input_lengths[i]) + decoded = tokenizer.decode( + input_ids[i, :length].tolist(), skip_special_tokens=False + ) + assert " 16" in decoded[-64:], ( + f"sample {i}: decoded tail {decoded[-64:]!r} missing ' 16'" + ) + assert int(token_mask[i, :length].sum().item()) > 0, ( + f"sample {i}: token_mask has no assistant tokens" + ) + + tags = ray.get( + tq_actor.get_tags.remote(partition_id=_PARTITION_ID, sample_ids=meta.sample_ids) + ) + for tag in tags: + assert tag["weight_version"] == 0 + # Slim tag schema: weight_version is the only field producers stamp. + assert set(tag) == {"weight_version"} diff --git a/tests/unit/single_controller/test_tq_replay_buffer.py b/tests/unit/single_controller/test_tq_replay_buffer.py new file mode 100644 index 0000000000..55b45c9ece --- /dev/null +++ b/tests/unit/single_controller/test_tq_replay_buffer.py @@ -0,0 +1,324 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for TQReplayBuffer (plain SC-process buffer + TQ proxy).""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest +import torch + +import nemo_rl.algorithms.async_utils.replay_buffer as _replay_buffer_module +from nemo_rl.algorithms.async_utils.replay_buffer import TQReplayBuffer +from nemo_rl.data_plane import KVBatchMeta +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.experience.interfaces import PromptGroupRecord + +# Each record yields _N_GENS training rows. +_N_GENS = 2 + + +def _stub_record_to_train_batch( + record: PromptGroupRecord, *, pad_value_dict: Any +) -> BatchedDataDict[Any]: + del record, pad_value_dict + return BatchedDataDict[Any]( + { + "input_ids": torch.ones((_N_GENS, 3), dtype=torch.long), + "input_lengths": torch.full((_N_GENS,), 3, dtype=torch.long), + "total_reward": torch.zeros(_N_GENS, dtype=torch.float32), + } + ) + + +@pytest.fixture(autouse=True) +def _patch_converter(monkeypatch): + """Bypass the real ``record_to_train_batch`` so tests can use empty records.""" + monkeypatch.setattr( + _replay_buffer_module, + "record_to_train_batch", + _stub_record_to_train_batch, + ) + + +class FakeDataPlaneClient: + """Sync in-memory DataPlaneClient stub used by TQReplayBuffer tests.""" + + def __init__(self, partition_id: str = "rollout_data") -> None: + self._partition_id = partition_id + self._rows: dict[str, dict[str, Any]] = {} + self.put_calls: list[dict[str, Any]] = [] + self.clear_calls: list[list[str]] = [] + + def put_samples( + self, + sample_ids: list[str], + partition_id: str, + fields: Any = None, + tags: list[dict[str, Any]] | None = None, + ) -> KVBatchMeta: + assert partition_id == self._partition_id + self.put_calls.append( + { + "sample_ids": list(sample_ids), + "fields": fields, + "tags": [dict(t) for t in tags] if tags is not None else None, + } + ) + for i, sid in enumerate(sample_ids): + self._rows[sid] = { + "tag": dict(tags[i]) if tags is not None else {}, + } + return KVBatchMeta( + partition_id=partition_id, + task_name=None, + sample_ids=list(sample_ids), + fields=None, + tags=[dict(t) for t in tags] if tags is not None else None, + ) + + def clear_samples(self, sample_ids: list[str] | None, partition_id: str) -> None: + assert partition_id == self._partition_id + ids = list(sample_ids) if sample_ids is not None else list(self._rows) + self.clear_calls.append(list(ids)) + for sid in ids: + self._rows.pop(sid, None) + + def depth(self) -> int: + return len(self._rows) + + +def _run(coro): + return asyncio.run(coro) + + +def _make_record() -> PromptGroupRecord: + """Opaque PromptGroupRecord — converter is stubbed, so contents are unused.""" + return PromptGroupRecord( + prompt_idx=0, + prompt=[], + extra_env_info=None, + metadata={}, + completions=[], + rollout_metrics={}, + ) + + +def _make_buffer(dp: FakeDataPlaneClient) -> TQReplayBuffer: + return TQReplayBuffer( + dp, partition_id="rollout_data", pad_value_dict={"token_ids": 0} + ) + + +def _add_group( + buf: TQReplayBuffer, weight: int, end_weight: int | None = None +) -> KVBatchMeta: + if end_weight is None: + end_weight = weight + group_id = buf.reserve(weight_version=weight) + return _run( + buf.commit( + group_id, + _make_record(), + start_weight_version=weight, + end_weight_version=end_weight, + ) + ) + + +class TestTQReplayBufferReserveCommit: + def test_reserve_appends_placeholder_unready(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + + group_id = buf.reserve(weight_version=3) + + assert isinstance(group_id, str) and group_id + assert buf.size() == 1 + assert buf.start_weight_list == [3] + assert buf.end_weight_list == [-1] + assert buf.ready_list == [False] + assert buf.meta_list == [None] + assert dp.depth() == 0 + assert dp.put_calls == [] + + def test_commit_writes_tq_then_fills_meta(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + + group_id = buf.reserve(weight_version=3) + meta = _run( + buf.commit( + group_id, + _make_record(), + start_weight_version=3, + end_weight_version=4, + ) + ) + + # pack_payload stamps sample_ids as ``{group_uuid}_g{i}``. + assert len(meta.sample_ids) == _N_GENS + head, _, idx = meta.sample_ids[0].rpartition("_g") + assert head == group_id and idx == "0" + assert all(sid.startswith(group_id + "_g") for sid in meta.sample_ids) + assert dp.depth() == _N_GENS + assert buf.size() == 1 + assert buf.start_weight_list == [3] + assert buf.end_weight_list == [4] + assert buf.ready_list == [True] + assert buf.meta_list[0].sample_ids == meta.sample_ids + # TQ tag uses start_weight_version (dispatch time). + assert meta.tags == [{"weight_version": 3}] * _N_GENS + assert len(dp.put_calls) == 1 + + def test_commit_raises_for_unknown_group_id(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + buf.reserve(weight_version=3) + + with pytest.raises(ValueError): + _run( + buf.commit( + "not-a-real-id", + _make_record(), + start_weight_version=3, + end_weight_version=3, + ) + ) + + def test_reserve_then_commit_preserves_dispatch_order(self): + """Reserve in dispatch order, commit out of order; insertion order holds.""" + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + + weights = (1, 2, 3) + gids = [buf.reserve(weight_version=w) for w in weights] + # Commit out of order: 2, 0, 1 — buffer order must still match reserve order. + for i in (2, 0, 1): + _run( + buf.commit( + gids[i], + _make_record(), + start_weight_version=weights[i], + end_weight_version=weights[i], + ) + ) + + assert buf.size() == 3 + assert buf.start_weight_list == [1, 2, 3] + assert buf.end_weight_list == [1, 2, 3] + assert buf.ready_list == [True, True, True] + # sample_id head equals reserved group_id at each slot. + for i, gid in enumerate(gids): + assert buf.meta_list[i] is not None + assert buf.meta_list[i].sample_ids[0].startswith(gid + "_g") + + def test_commit_appends_multiple_records_in_order(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + + metas = [_add_group(buf, weight=w) for w in (1, 2, 3)] + + assert buf.size() == 3 + assert buf.start_weight_list == [1, 2, 3] + assert buf.end_weight_list == [1, 2, 3] + assert [m.sample_ids for m in buf.meta_list] == [ + list(metas[0].sample_ids), + list(metas[1].sample_ids), + list(metas[2].sample_ids), + ] + + +class TestTQReplayBufferRemove: + def test_remove_drops_indices_and_clears_dp_when_requested(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + metas = [_add_group(buf, weight=g) for g in range(3)] + + n = _run(buf.remove([0, 2], remove_in_dp=True)) + + assert n == 2 + assert buf.size() == 1 + assert buf.start_weight_list == [1] + assert buf.end_weight_list == [1] + assert buf.meta_list[0].sample_ids == list(metas[1].sample_ids) + assert dp.depth() == _N_GENS + assert set(dp._rows) == set(metas[1].sample_ids) + + def test_remove_without_dp_keeps_rows(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + metas = [_add_group(buf, weight=g) for g in range(2)] + + n = _run(buf.remove([0], remove_in_dp=False)) + + assert n == 1 + assert buf.size() == 1 + assert buf.start_weight_list == [1] + assert buf.end_weight_list == [1] + assert buf.meta_list[0].sample_ids == list(metas[1].sample_ids) + assert dp.clear_calls == [] + assert dp.depth() == 2 * _N_GENS + + def test_remove_rejects_out_of_range_before_mutating(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + metas = [_add_group(buf, weight=g) for g in range(2)] + + with pytest.raises(IndexError, match=r"out of range: 5; size=2"): + _run(buf.remove([0, 5], remove_in_dp=True)) + + assert buf.size() == 2 + assert [m.sample_ids for m in buf.meta_list] == [ + list(metas[0].sample_ids), + list(metas[1].sample_ids), + ] + assert dp.depth() == 2 * _N_GENS + assert dp.clear_calls == [] + + def test_remove_empty_is_noop(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + _add_group(buf, weight=0) + _add_group(buf, weight=0) + + n = _run(buf.remove([], remove_in_dp=True)) + + assert n == 0 + assert buf.size() == 2 + assert dp.depth() == 2 * _N_GENS + assert dp.clear_calls == [] + + +class TestTQReplayBufferSize: + def test_size_and_len(self): + dp = FakeDataPlaneClient() + buf = _make_buffer(dp) + assert buf.size() == 0 + assert len(buf) == 0 + + _add_group(buf, weight=0) + assert buf.size() == 1 + assert len(buf) == 1 + + _add_group(buf, weight=0) + assert buf.size() == 2 + assert len(buf) == 2 + + _run(buf.remove([0], remove_in_dp=True)) + assert buf.size() == 1 + assert len(buf) == 1 From b7783e92e66ca75d30d11e61d16298311f3ddbdd Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Fri, 17 Jul 2026 06:40:21 -0700 Subject: [PATCH 02/11] lint Signed-off-by: Yuki Huang --- nemo_rl/algorithms/single_controller.py | 1 + pyrefly.toml | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index 44763c71bb..c3df3bef16 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -58,6 +58,7 @@ incomplete_group_indices, min_weight_version, ) +from nemo_rl.data.interfaces import DatumSpec from nemo_rl.data_plane import KVBatchMeta from nemo_rl.utils.logger import Logger from nemo_rl.utils.timer import Timer diff --git a/pyrefly.toml b/pyrefly.toml index 71a0fb94c3..f24606c677 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -138,6 +138,7 @@ project-includes = [ "nemo_rl/experience/__init__.py", "nemo_rl/experience/interfaces.py", "nemo_rl/experience/metric_utils.py", + "nemo_rl/experience/payload.py", "nemo_rl/experience/rollout_manager.py", "nemo_rl/experience/rollouts.py", "nemo_rl/modelopt/__init__.py", @@ -182,7 +183,6 @@ project-includes = [ "nemo_rl/models/generation/vllm/worker_utils.py", "nemo_rl/models/huggingface/__init__.py", "nemo_rl/models/megatron/__init__.py", - "nemo_rl/models/megatron/draft/__init__.py", "nemo_rl/models/policy/__init__.py", "nemo_rl/models/policy/interfaces.py", "nemo_rl/models/policy/utils.py", From ea4b476f1e058fd375c4b48a711371e8788f2c0c Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Fri, 17 Jul 2026 07:41:43 -0700 Subject: [PATCH 03/11] fix(sc): wire rollout_manager/dataloader into SingleControllerActor, add target_step to _FakeBuffer.reserve Signed-off-by: Yuki Huang --- nemo_rl/algorithms/single_controller.py | 31 +++++++++-- tests/unit/experience/test_rollout_manager.py | 8 ++- .../single_controller/test_rollout_pump.py | 55 ++++++++++--------- 3 files changed, 60 insertions(+), 34 deletions(-) diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index c3df3bef16..c46f5f7f8f 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -116,6 +116,13 @@ class SingleControllerConfig(BaseModel, extra="allow"): max_inflight_prompts: int = 8 max_buffered_rollouts: int = 8 # _buffer_capacity semaphore size + # Rollout dispatch gating (read by _rollout_pump). over_sampling=False + # gates each batch on max_rollout_version vs trainer_version; force_in_order + # stamps target_step on each dispatch so downstream consumers can match + # rollout batches to trainer steps exactly. + over_sampling: bool = False + force_in_order: bool = False + # Training max_train_steps: int = 10 max_rollout_prompts: int = 32 @@ -236,6 +243,8 @@ def __init__( weight_synchronizer: Any, loss_fn: Any, advantage_estimator: Any | None = None, + rollout_manager: Any = None, + dataloader: Any = None, ) -> None: self._cfg = cfg self._prompts = prompts @@ -245,6 +254,8 @@ def __init__( self._weight_synchronizer = weight_synchronizer self._loss_fn = loss_fn self._advantage_estimator = advantage_estimator + self._rollout_manager = rollout_manager + self._dataloader = dataloader # Built here, not on the driver: Logger backends (wandb/tb/...) hold # _thread.lock that Ray can't cloudpickle into the actor. @@ -314,6 +325,14 @@ def __init__( # Count of in-flight generate_and_push calls self._inflight_rollouts: int = 0 + # Rollout batch counter — pre-incremented before each dispatch, so start + # at -1 to allow the first batch through the strict_on_policy gate. + self._max_rollout_version: int = -1 + + # Strong refs to dispatched rollout tasks so asyncio doesn't GC them + # while they're still running (removed via done_callback on completion). + self._dispatched_rollouts: set[asyncio.Task] = set() + # Backpressure valve: max unconsumed rollout groups allowed in DataPlane. # Acquired before each rollout dispatch; released after clear_samples. self._buffer_capacity: asyncio.Semaphore = asyncio.Semaphore( @@ -413,10 +432,10 @@ async def _rollout_pump(self) -> None: group via TQReplayBuffer (→ dp_client.put_samples + mark ready) 5. Decrement _inflight_rollouts """ - sem = asyncio.Semaphore(self._async_cfg.max_inflight_prompts) - over_sampling = self._async_cfg.over_sampling - max_staleness = self._async_cfg.max_weight_staleness_versions - force_in_order = self._async_cfg.force_in_order + sem = asyncio.Semaphore(self._cfg.max_inflight_prompts) + over_sampling = self._cfg.over_sampling + max_staleness = self._cfg.max_weight_staleness_versions + force_in_order = self._cfg.force_in_order print("rollout_pump: starting", flush=True) async def _dispatch_one_prompt( @@ -427,7 +446,7 @@ async def _dispatch_one_prompt( await self._rollout_manager.generate_and_push( prompt, target_step=target_step ) - if self._diagnostics: + if self._cfg.diagnostics: content = "" for i in range(len(prompt["message_log"])): if prompt["message_log"][i]["role"] == "user": @@ -438,7 +457,7 @@ async def _dispatch_one_prompt( self._inflight_rollouts -= 1 sem.release() - max_epochs = self._master_config.grpo["max_num_epochs"] + max_epochs = self._cfg.max_num_epochs epoch = 0 while max_epochs is None or epoch < max_epochs: for prompt_batch in self._dataloader: diff --git a/tests/unit/experience/test_rollout_manager.py b/tests/unit/experience/test_rollout_manager.py index 63a590e6a1..0aa24a4364 100644 --- a/tests/unit/experience/test_rollout_manager.py +++ b/tests/unit/experience/test_rollout_manager.py @@ -76,7 +76,13 @@ def __init__(self) -> None: # reserve(weight_version=X) -> group_id; commit fills the slot. self._slots: list[str] = [] - def reserve(self, *, weight_version: int, group_id: str | None = None) -> str: + def reserve( + self, + *, + weight_version: int, + target_step: int | None = None, + group_id: str | None = None, + ) -> str: if group_id is None: group_id = str(uuid.uuid4()) self.reserve_calls.append(weight_version) diff --git a/tests/unit/single_controller/test_rollout_pump.py b/tests/unit/single_controller/test_rollout_pump.py index b4d5f291b3..7130d69edb 100644 --- a/tests/unit/single_controller/test_rollout_pump.py +++ b/tests/unit/single_controller/test_rollout_pump.py @@ -24,11 +24,9 @@ from tensordict import TensorDict from nemo_rl.algorithms.async_utils.replay_buffer import TQReplayBuffer -from nemo_rl.algorithms.single_controller import SingleControllerActor -from nemo_rl.algorithms.single_controller_utils import ( - AsyncRLConfig, - MasterConfig, - SingleControllerBundle, +from nemo_rl.algorithms.single_controller import ( + SingleControllerActor, + SingleControllerConfig, ) from nemo_rl.data_plane.adapters.noop import NoOpDataPlaneClient from nemo_rl.distributed.batched_data_dict import BatchedDataDict @@ -156,6 +154,7 @@ def _padded(td: TensorDict) -> TensorDict: def test_rollout_pump_writes_expected_tq_data( multi_step_setup_vllm_async, # noqa: F811 single_multi_step_calculator_input_sample, # noqa: F811 + tmp_path, ): """SC._rollout_pump writes max_rollout_prompts * num_generations rows to TQ with the expected fields and tags.""" vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async @@ -176,19 +175,25 @@ def test_rollout_pump_writes_expected_tq_data( ) dp_adapter = _SyncDPAdapter(tq_actor) - mc = MasterConfig.model_construct( - grpo={ - "max_num_steps": 1, - "max_num_epochs": None, - "num_generations_per_prompt": num_generations, + cfg = SingleControllerConfig.model_construct( + batch_selection_strategy="strict_on_policy", + max_weight_staleness_versions=0, + min_groups_per_batch=1, + group_size=num_generations, + max_inflight_prompts=max_rollout_prompts, + max_buffered_rollouts=max_rollout_prompts, + max_train_steps=1, + max_num_epochs=None, + over_sampling=True, + partition_id=_PARTITION_ID, + logger={ + "log_dir": str(tmp_path / "logs"), + "wandb_enabled": False, + "swanlab_enabled": False, + "tensorboard_enabled": False, + "mlflow_enabled": False, + "monitor_gpus": False, }, - async_rl=AsyncRLConfig( - batch_selection_strategy="strict_on_policy", - max_weight_staleness_versions=0, - min_prompt_groups_per_batch=1, - max_inflight_prompts=max_rollout_prompts, - max_buffered_rollouts=max_rollout_prompts, - ), ) # Wrap each value in a single-element list so size==1 and v[0] returns the original field. batched_sample = BatchedDataDict({k: [v] for k, v in input_sample.items()}) @@ -209,22 +214,18 @@ def test_rollout_pump_writes_expected_tq_data( use_nemo_gym=False, tq_buffer=tq_buffer, ) - bundle = SingleControllerBundle( + ctrl = SingleControllerActor.remote( + cfg=cfg, + prompts=[], + dp_client_handle=dp_adapter, gen_handle=vllm_generation, trainer_handle=object(), - env_handles=env_handles, - train_cluster=None, - inference_cluster=None, - dp_client=dp_adapter, - dataloader=dataloader, weight_synchronizer=object(), - advantage_estimator=None, loss_fn=None, + advantage_estimator=None, rollout_manager=rollout_manager, - tq_buffer=tq_buffer, - partition_id=_PARTITION_ID, + dataloader=dataloader, ) - ctrl = SingleControllerActor.remote(master_config=mc, bundle=bundle) vllm_generation.prepare_for_generation() From 246a84a406864f6b1de86150e5167b0722896889 Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Fri, 17 Jul 2026 07:56:04 -0700 Subject: [PATCH 04/11] refactor(sc): drop unused _rollout_done and _flush_incomplete_groups shutdown path Signed-off-by: Yuki Huang --- nemo_rl/algorithms/single_controller.py | 62 ++----------------------- 1 file changed, 4 insertions(+), 58 deletions(-) diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index c46f5f7f8f..f8b84c4b36 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -55,7 +55,6 @@ from nemo_rl.algorithms.staleness_sampler import ( StalenessSampler, count_groups, - incomplete_group_indices, min_weight_version, ) from nemo_rl.data.interfaces import DatumSpec @@ -341,7 +340,6 @@ def __init__( self._trainer_version: int = 0 self._train_steps: int = 0 - self._rollout_done: bool = False # Completed prompt-list passes; only advances when # cfg.max_num_epochs is set (see _rollout_pump). self._current_epoch: int = 0 @@ -458,8 +456,7 @@ async def _dispatch_one_prompt( sem.release() max_epochs = self._cfg.max_num_epochs - epoch = 0 - while max_epochs is None or epoch < max_epochs: + while max_epochs is None or self._current_epoch < max_epochs: for prompt_batch in self._dataloader: # over_sampling=False: batch-level gate on max_rollout_version. if not over_sampling: @@ -491,9 +488,10 @@ async def _dispatch_one_prompt( ) self._dispatched_rollouts.add(task) task.add_done_callback(self._dispatched_rollouts.discard) - epoch += 1 - print(f"rollout_pump: completed {epoch} epoch(s)", flush=True) + self._current_epoch += 1 + + print(f"rollout_pump: completed {self._current_epoch} epoch(s)", flush=True) async def _train_pump(self) -> None: """Per-prompt-group streaming train loop. @@ -568,18 +566,6 @@ async def _train_pump(self) -> None: ) if group_indices is None: - if self._rollout_done: - # No group is selectable and no more samples - # will arrive: flush groups that can never - # complete so the emptiness check fires - # instead of spinning forever. - await self._flush_incomplete_groups() - if ( - self._claimed_meta is None - or self._claimed_meta.size == 0 - ): - rollout_exhausted = True - break await asyncio.sleep(0.005) continue @@ -879,46 +865,6 @@ async def _evict_stale_claimed(self) -> KVBatchMeta | None: self._claimed_meta = self._claimed_meta.drop(indices) return evicted_meta - async def _flush_incomplete_groups(self) -> None: - """Drop groups that can never become selectable after rollout shutdown. - - With ``_rollout_done`` set, a group that is uncommitted or short of - ``expected_num_samples`` will never receive more rows, yet the sampler - neither selects nor evicts it — without this flush the train pump - spins forever and ``run()`` hangs. Drain rows still unclaimed at - DataPlane first: a group can straddle a ``claim_meta`` batch boundary, - so incomplete-in-``_claimed_meta`` does not yet prove - incomplete-in-partition. - """ - while True: - before = self._claimed_meta.size if self._claimed_meta is not None else 0 - await self._claim_available_meta() - after = self._claimed_meta.size if self._claimed_meta is not None else 0 - if after == before: - break - if self._claimed_meta is None or self._claimed_meta.size == 0: - return - indices = incomplete_group_indices( - self._claimed_meta, - group_size=self._cfg.group_size, - ) - if not indices: - return - dropped = self._claimed_meta.subset(indices) - print( - f"WARNING: rollout done: dropping {dropped.size} sample(s) from " - f"incomplete prompt group(s) that can no longer complete", - flush=True, - ) - await self._call_dp( - "clear_samples", - sample_ids=dropped.sample_ids, - partition_id=dropped.partition_id, - ) - self._claimed_meta = self._claimed_meta.drop(indices) - # No _buffer_capacity release: the rollout pump has exited, so no - # dispatcher will acquire again this run. - def _tensor_field(data: TensorDict, field_name: str) -> torch.Tensor: value = data[field_name] From beb3df4c3ebb839e5e71dd66f109e9fa84d8361a Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Fri, 17 Jul 2026 08:02:11 -0700 Subject: [PATCH 05/11] refactor: place max_rollout_turns right after max_seq_len in rollout_manager signatures Signed-off-by: Yuki Huang --- nemo_rl/experience/rollout_manager.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/nemo_rl/experience/rollout_manager.py b/nemo_rl/experience/rollout_manager.py index 6bf2773828..7d120a399d 100644 --- a/nemo_rl/experience/rollout_manager.py +++ b/nemo_rl/experience/rollout_manager.py @@ -51,8 +51,8 @@ def __init__( env_handles: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, + max_rollout_turns: int, policy_generation: GenerationInterface, - max_rollout_turns: int = 999999, **kwargs: Any, ) -> None: self._tokenizer = tokenizer @@ -403,8 +403,8 @@ def __init__( env_handles: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, - generation_config: GenerationConfig, max_rollout_turns: int, + generation_config: GenerationConfig, **kwargs: Any, ) -> None: self._tokenizer = tokenizer @@ -647,7 +647,7 @@ def __init__( env_handles: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, - max_rollout_turns: Optional[int] = None, + max_rollout_turns: int = 1, policy_generation: Optional[GenerationInterface] = None, generation_config: Optional[GenerationConfig] = None, use_nemo_gym: bool = False, @@ -662,8 +662,6 @@ def __init__( assert policy_generation is not None, ( "policy_generation is required for the native async path" ) - if max_rollout_turns is None: - max_rollout_turns = 999999 # use AsyncRolloutImpl's default value else: rollout_cls = AsyncNemoGymRolloutImpl assert generation_config is not None, ( @@ -675,7 +673,7 @@ def __init__( env_handles=env_handles, num_generations_per_prompt=num_generations_per_prompt, max_seq_len=max_seq_len, - max_rollout_turns=max_rollout_turns, # type: ignore + max_rollout_turns=max_rollout_turns, policy_generation=policy_generation, # type: ignore generation_config=generation_config, ) From 7fbc54fea655331f86ad09cfa484ed8ac5ae13b5 Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Fri, 17 Jul 2026 08:08:55 -0700 Subject: [PATCH 06/11] fix(sc): guard TQReplayBuffer.commit against unknown group_id Signed-off-by: Yuki Huang --- nemo_rl/algorithms/async_utils/replay_buffer.py | 7 +++++++ tests/unit/single_controller/test_tq_replay_buffer.py | 4 ++++ 2 files changed, 11 insertions(+) diff --git a/nemo_rl/algorithms/async_utils/replay_buffer.py b/nemo_rl/algorithms/async_utils/replay_buffer.py index 3264ed7722..ba2a1ba87e 100644 --- a/nemo_rl/algorithms/async_utils/replay_buffer.py +++ b/nemo_rl/algorithms/async_utils/replay_buffer.py @@ -709,6 +709,13 @@ async def commit( Raises: ValueError: group_id has no live slot (removed or never reserved). """ + # Precondition: reserve() must have registered this group_id. Raise + # before any side effects so a stray commit doesn't leak orphan DP rows. + if group_id not in self._group_ids: + raise ValueError( + f"commit called with unknown group_id={group_id!r}; " + f"reserve() must precede commit() (or the slot was already removed)" + ) train_batch = record_to_train_batch(record, pad_value_dict=self._pad_value_dict) sample_ids, fields, tags = pack_payload( train_batch, weight_version=start_weight_version, group_id=group_id diff --git a/tests/unit/single_controller/test_tq_replay_buffer.py b/tests/unit/single_controller/test_tq_replay_buffer.py index 55b45c9ece..12e8249a9f 100644 --- a/tests/unit/single_controller/test_tq_replay_buffer.py +++ b/tests/unit/single_controller/test_tq_replay_buffer.py @@ -200,6 +200,10 @@ def test_commit_raises_for_unknown_group_id(self): ) ) + # No orphan rows in DataPlane: commit must validate group_id before writing. + assert dp.depth() == 0 + assert dp.put_calls == [] + def test_reserve_then_commit_preserves_dispatch_order(self): """Reserve in dispatch order, commit out of order; insertion order holds.""" dp = FakeDataPlaneClient() From cb7767693c501bb3413bf36699f5b2b78818b718 Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Fri, 17 Jul 2026 08:18:49 -0700 Subject: [PATCH 07/11] fix(sc): pass token_aligned_fields in pack_payload Signed-off-by: Yuki Huang --- nemo_rl/experience/payload.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/nemo_rl/experience/payload.py b/nemo_rl/experience/payload.py index bdf95ab290..9f0dbfc987 100644 --- a/nemo_rl/experience/payload.py +++ b/nemo_rl/experience/payload.py @@ -22,6 +22,7 @@ from tensordict import TensorDict from nemo_rl.data_plane.codec import pack_jagged_fields +from nemo_rl.data_plane.column_io import TOKEN_ALIGNED_FIELDS from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.experience.interfaces import PromptGroupRecord @@ -111,7 +112,9 @@ def pack_payload( if isinstance(v, torch.Tensor) or (isinstance(v, np.ndarray) and v.dtype == object) } - fields_td = pack_jagged_fields(tensor_fields, lengths=lengths) + fields_td = pack_jagged_fields( + tensor_fields, lengths=lengths, token_aligned_fields=TOKEN_ALIGNED_FIELDS + ) sample_ids = [f"{group_id}_g{i}" for i in range(n)] tags = [{"weight_version": weight_version} for _ in range(n)] return sample_ids, fields_td, tags From ba226abaed58a91dc909454eb1845d287ecd48a4 Mon Sep 17 00:00:00 2001 From: Yuki Huang Date: Sat, 18 Jul 2026 06:18:20 -0700 Subject: [PATCH 08/11] fix(sc): reject over_sampling=True with strict_on_policy/staleness=0/force_in_order Signed-off-by: Yuki Huang --- nemo_rl/algorithms/single_controller.py | 17 +++++++++++++++-- .../unit/single_controller/test_rollout_pump.py | 4 ++-- 2 files changed, 17 insertions(+), 4 deletions(-) diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index f8b84c4b36..088aba3ff5 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -270,11 +270,24 @@ def __init__( # values at config construction, so no runtime assert is needed here. if cfg.batch_selection_strategy == "strict_on_policy": cfg.max_weight_staleness_versions = 0 + cfg.over_sampling = False print( - "Using strict_on_policy, auto setting " - "max_weight_staleness_versions to 0.", + "Using strict_on_policy, auto setting max_weight_staleness_versions " + "to 0 and over_sampling to False.", flush=True, ) + if cfg.max_weight_staleness_versions == 0 and cfg.over_sampling: + raise ValueError( + "max_weight_staleness_versions=0 requires over_sampling=False: " + "with zero staleness the dispatch gate needs to advance one batch " + "per trainer_version, which over_sampling=True bypasses." + ) + if cfg.force_in_order and cfg.over_sampling: + raise ValueError( + "force_in_order=True requires over_sampling=False so that each " + "dispatched batch corresponds to exactly one target training step." + ) + if cfg.target_groups_per_step is None: cfg.target_groups_per_step = cfg.min_groups_per_batch if cfg.target_groups_per_step < cfg.min_groups_per_batch: diff --git a/tests/unit/single_controller/test_rollout_pump.py b/tests/unit/single_controller/test_rollout_pump.py index 7130d69edb..87054c31fd 100644 --- a/tests/unit/single_controller/test_rollout_pump.py +++ b/tests/unit/single_controller/test_rollout_pump.py @@ -176,8 +176,8 @@ def test_rollout_pump_writes_expected_tq_data( dp_adapter = _SyncDPAdapter(tq_actor) cfg = SingleControllerConfig.model_construct( - batch_selection_strategy="strict_on_policy", - max_weight_staleness_versions=0, + batch_selection_strategy="staleness_window", + max_weight_staleness_versions=1, min_groups_per_batch=1, group_size=num_generations, max_inflight_prompts=max_rollout_prompts, From fce9ad526b3115f173c8d03cdd6575e5849557b2 Mon Sep 17 00:00:00 2001 From: ruit Date: Tue, 21 Jul 2026 00:53:39 -0700 Subject: [PATCH 09/11] fix(sc): address rollout path review feedback Signed-off-by: ruit --- .../algorithms/async_utils/replay_buffer.py | 79 +++++-- nemo_rl/algorithms/single_controller.py | 136 ++++++----- nemo_rl/experience/rollout_manager.py | 40 ++-- pyrefly.toml | 1 + tests/unit/excluded_unit_tests.sh | 6 +- tests/unit/experience/test_rollout_manager.py | 42 +++- .../single_controller/test_rollout_pump.py | 211 +++++++++++++++++- .../test_tq_replay_buffer.py | 37 +++ 8 files changed, 444 insertions(+), 108 deletions(-) diff --git a/nemo_rl/algorithms/async_utils/replay_buffer.py b/nemo_rl/algorithms/async_utils/replay_buffer.py index ba2a1ba87e..61b1f11f8a 100644 --- a/nemo_rl/algorithms/async_utils/replay_buffer.py +++ b/nemo_rl/algorithms/async_utils/replay_buffer.py @@ -720,30 +720,65 @@ async def commit( sample_ids, fields, tags = pack_payload( train_batch, weight_version=start_weight_version, group_id=group_id ) - await self._call_dp( - "put_samples", - sample_ids=sample_ids, - partition_id=self._partition_id, - fields=fields, - tags=tags, - ) + try: + await self._call_dp( + "put_samples", + sample_ids=sample_ids, + partition_id=self._partition_id, + fields=fields, + tags=tags, + ) - # mirrors kv_first_write - lengths = train_batch["input_lengths"] - meta = KVBatchMeta( - partition_id=self._partition_id, - task_name="train", - sample_ids=list(sample_ids), - fields=list(fields.keys()), - sequence_lengths=[int(s) for s in lengths.tolist()], - tags=[dict(t) for t in tags], - ) + # mirrors kv_first_write + lengths = train_batch["input_lengths"] + meta = KVBatchMeta( + partition_id=self._partition_id, + task_name="train", + sample_ids=list(sample_ids), + fields=list(fields.keys()), + sequence_lengths=[int(s) for s in lengths.tolist()], + tags=[dict(t) for t in tags], + ) - idx = self._group_ids.index(group_id) - self.meta_list[idx] = meta - self.end_weight_list[idx] = end_weight_version - self.ready_list[idx] = True - return meta + idx = self._group_ids.index(group_id) + self.meta_list[idx] = meta + self.end_weight_list[idx] = end_weight_version + self.ready_list[idx] = True + return meta + except BaseException as commit_error: + # put_samples may have written rows before raising. Roll back by the + # deterministic IDs known here; the caller removes the reserved slot. + try: + await self._call_dp( + "clear_samples", + sample_ids=list(sample_ids), + partition_id=self._partition_id, + ) + except BaseException as rollback_error: + raise BaseExceptionGroup( + f"commit and rollback both failed for group_id={group_id!r}", + [commit_error, rollback_error], + ) + raise + + async def remove_group(self, group_id: str, *, remove_in_dp: bool = False) -> int: + """Remove the live slot identified by ``group_id``. + + Args: + group_id: Group identifier returned by :meth:`reserve`. + remove_in_dp: Whether to clear rows referenced by a committed slot. + + Returns: + Number of removed slots (always one on success). + + Raises: + ValueError: ``group_id`` has no live slot. + """ + try: + idx = self._group_ids.index(group_id) + except ValueError as error: + raise ValueError(f"unknown group_id={group_id!r}") from error + return await self.remove([idx], remove_in_dp=remove_in_dp) async def remove(self, idxs: list[int], remove_in_dp: bool) -> int: """Drop entries at the given indices and optionally clear them from DataPlane. diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index 088aba3ff5..16a0461f31 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -44,6 +44,7 @@ import asyncio import time from contextlib import nullcontext +from functools import partial from typing import Any, Literal, Optional import ray @@ -341,9 +342,9 @@ def __init__( # at -1 to allow the first batch through the strict_on_policy gate. self._max_rollout_version: int = -1 - # Strong refs to dispatched rollout tasks so asyncio doesn't GC them - # while they're still running (removed via done_callback on completion). - self._dispatched_rollouts: set[asyncio.Task] = set() + # Active rollout tasks used by downstream synchronization/drain paths. + # TaskGroup remains responsible for task ownership and cancellation. + self._dispatched_rollouts: set[asyncio.Task[None]] = set() # Backpressure valve: max unconsumed rollout groups allowed in DataPlane. # Acquired before each rollout dispatch; released after clear_samples. @@ -382,14 +383,19 @@ async def run(self) -> dict[str, Any]: """Main entry point. Runs until max_train_steps is reached.""" rollout_task = asyncio.create_task(self._rollout_pump()) train_task = asyncio.create_task(self._train_pump()) - - await train_task - - rollout_task.cancel() try: - await rollout_task - except asyncio.CancelledError: - pass + done, _ = await asyncio.wait( + {rollout_task, train_task}, return_when=asyncio.FIRST_COMPLETED + ) + if rollout_task in done: + # Propagate rollout failures immediately. A normally exhausted + # rollout pump leaves the train pump to drain committed groups. + await rollout_task + await train_task + finally: + rollout_task.cancel() + train_task.cancel() + await asyncio.gather(rollout_task, train_task, return_exceptions=True) return { "train_steps": self._train_steps, @@ -450,59 +456,87 @@ async def _rollout_pump(self) -> None: print("rollout_pump: starting", flush=True) async def _dispatch_one_prompt( - prompt: DatumSpec, target_step: Optional[int] + prompt: DatumSpec, + target_step: Optional[int], + task_started_event: asyncio.Event, ) -> None: + task_started_event.set() self._inflight_rollouts += 1 try: await self._rollout_manager.generate_and_push( prompt, target_step=target_step ) - if self._cfg.diagnostics: - content = "" - for i in range(len(prompt["message_log"])): - if prompt["message_log"][i]["role"] == "user": - content = prompt["message_log"][i]["content"] - break - print(f" rollout done for prompt='{content[:20]}...'", flush=True) + except BaseException: + # On success ownership transfers to the train pump, which + # releases this permit after consuming the committed group. + self._buffer_capacity.release() + raise finally: self._inflight_rollouts -= 1 sem.release() + if self._cfg.diagnostics: + content = "" + for i in range(len(prompt["message_log"])): + if prompt["message_log"][i]["role"] == "user": + content = prompt["message_log"][i]["content"] + break + print(f" rollout done for prompt='{content[:20]}...'", flush=True) + + def _release_permits_if_task_not_started( + _: asyncio.Task[Any], + *, + task_started_event: asyncio.Event, + ) -> None: + if not task_started_event.is_set(): + self._buffer_capacity.release() + sem.release() + max_epochs = self._cfg.max_num_epochs - while max_epochs is None or self._current_epoch < max_epochs: - for prompt_batch in self._dataloader: - # over_sampling=False: batch-level gate on max_rollout_version. - if not over_sampling: - while ( - self._max_rollout_version - >= self._trainer_version + max_staleness - ): - await asyncio.sleep(0.005) - self._max_rollout_version += 1 - - # target_step = batch dispatch index when force_in_order is on. - target_step = self._max_rollout_version if force_in_order else None - - for prompt_idx in range(prompt_batch.size): - prompt: DatumSpec = { # type: ignore - k: v[prompt_idx] for k, v in prompt_batch.items() - } - - # check if buffer is full - await self._buffer_capacity.acquire() - # check if inflight rollouts is full - await sem.acquire() - # wait for rollout to be permitted - await self._rollout_permitted.wait() - - # dispatch rollout - task = asyncio.create_task( - _dispatch_one_prompt(prompt, target_step) - ) - self._dispatched_rollouts.add(task) - task.add_done_callback(self._dispatched_rollouts.discard) + async with asyncio.TaskGroup() as rollout_tasks: + while max_epochs is None or self._current_epoch < max_epochs: + for prompt_batch in self._dataloader: + # over_sampling=False: batch-level gate on max_rollout_version. + if not over_sampling: + while ( + self._max_rollout_version + >= self._trainer_version + max_staleness + ): + await asyncio.sleep(0.005) + self._max_rollout_version += 1 + + # target_step = batch dispatch index when force_in_order is on. + target_step = self._max_rollout_version if force_in_order else None + + for prompt_idx in range(prompt_batch.size): + prompt: DatumSpec = { # type: ignore + k: v[prompt_idx] for k, v in prompt_batch.items() + } + + # check if buffer is full + await self._buffer_capacity.acquire() + # check if inflight rollouts is full + await sem.acquire() + # wait for rollout to be permitted + await self._rollout_permitted.wait() + + task_started_event = asyncio.Event() + # dispatch rollout + task = rollout_tasks.create_task( + _dispatch_one_prompt( + prompt, target_step, task_started_event + ) + ) + self._dispatched_rollouts.add(task) + task.add_done_callback(self._dispatched_rollouts.discard) + task.add_done_callback( + partial( + _release_permits_if_task_not_started, + task_started_event=task_started_event, + ) + ) - self._current_epoch += 1 + self._current_epoch += 1 print(f"rollout_pump: completed {self._current_epoch} epoch(s)", flush=True) diff --git a/nemo_rl/experience/rollout_manager.py b/nemo_rl/experience/rollout_manager.py index 7d120a399d..a1f485bf5e 100644 --- a/nemo_rl/experience/rollout_manager.py +++ b/nemo_rl/experience/rollout_manager.py @@ -48,7 +48,7 @@ class AsyncRolloutImpl: def __init__( self, tokenizer: TokenizerType, - env_handles: dict[str, EnvironmentInterface], + task_to_env: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, max_rollout_turns: int, @@ -56,7 +56,7 @@ def __init__( **kwargs: Any, ) -> None: self._tokenizer = tokenizer - self._env_handles = env_handles + self._task_to_env = task_to_env self._num_generations_per_prompt = num_generations_per_prompt self._max_seq_len = max_seq_len self._max_rollout_turns = max_rollout_turns @@ -189,7 +189,7 @@ async def _run_single_rollout( # step. In this case, need to wrap with asyncio.to_thread to make # this function yieldable. env_output = await asyncio.to_thread( - calculate_rewards, sample_batch, self._env_handles + calculate_rewards, sample_batch, self._task_to_env ) # Update reward and termination statistics @@ -400,7 +400,7 @@ class AsyncNemoGymRolloutImpl: def __init__( self, tokenizer: TokenizerType, - env_handles: dict[str, EnvironmentInterface], + task_to_env: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, max_rollout_turns: int, @@ -408,7 +408,7 @@ def __init__( **kwargs: Any, ) -> None: self._tokenizer = tokenizer - self._env_handles = env_handles + self._task_to_env = task_to_env self._num_generations_per_prompt = num_generations_per_prompt self._max_seq_len = max_seq_len self._max_rollout_turns = max_rollout_turns @@ -492,7 +492,7 @@ async def _run_rollouts( self, inputs: list[dict], timer: Timer, timer_prefix: str ) -> tuple[list[Completion], LLMMessageLogType, dict[str, Any]]: """Dispatch rows to NeMo-Gym; return completions, prompt, and metrics.""" - nemo_gym_env = self._env_handles["nemo_gym"] + nemo_gym_env = self._task_to_env["nemo_gym"] # Run generation and restore input order as results stream back. with timer.time(f"{timer_prefix}/run_rollouts"): @@ -644,7 +644,7 @@ class RolloutManager: def __init__( self, tokenizer: TokenizerType, - env_handles: dict[str, EnvironmentInterface], + task_to_env: dict[str, EnvironmentInterface], num_generations_per_prompt: int, max_seq_len: int, max_rollout_turns: int = 1, @@ -670,7 +670,7 @@ def __init__( self._impl: AsyncRolloutImpl | AsyncNemoGymRolloutImpl = rollout_cls( tokenizer=tokenizer, - env_handles=env_handles, + task_to_env=task_to_env, num_generations_per_prompt=num_generations_per_prompt, max_seq_len=max_seq_len, max_rollout_turns=max_rollout_turns, @@ -709,13 +709,17 @@ async def generate_and_push( group_id = self._tq_buffer.reserve( weight_version=start_version, target_step=target_step ) - - record = await self.run_rollout(input_sample) - end_version = self._weight_version - - await self._tq_buffer.commit( - group_id, - record, - start_weight_version=start_version, - end_weight_version=end_version, - ) + try: + record = await self.run_rollout(input_sample) + end_version = self._weight_version + await self._tq_buffer.commit( + group_id, + record, + start_weight_version=start_version, + end_weight_version=end_version, + ) + except BaseException: + # A failed rollout must not leave an unready slot that can block an + # in-order sampler. commit() rolls back any DataPlane rows it wrote. + await self._tq_buffer.remove_group(group_id) + raise diff --git a/pyrefly.toml b/pyrefly.toml index f24606c677..9a4237d347 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -183,6 +183,7 @@ project-includes = [ "nemo_rl/models/generation/vllm/worker_utils.py", "nemo_rl/models/huggingface/__init__.py", "nemo_rl/models/megatron/__init__.py", + "nemo_rl/models/megatron/draft/__init__.py", "nemo_rl/models/policy/__init__.py", "nemo_rl/models/policy/interfaces.py", "nemo_rl/models/policy/utils.py", diff --git a/tests/unit/excluded_unit_tests.sh b/tests/unit/excluded_unit_tests.sh index cdd80dddb6..32e914bbdb 100644 --- a/tests/unit/excluded_unit_tests.sh +++ b/tests/unit/excluded_unit_tests.sh @@ -218,9 +218,9 @@ EXCLUDED_UNIT_TESTS=( --deselect=tests/unit/experience/test_rollouts.py::test_max_seqlen_respected_sync --deselect=tests/unit/experience/test_rollouts.py::test_max_seqlen_respected_async --deselect=tests/unit/experience/test_rollouts.py::test_run_sliding_puzzle_vllm - --deselect=tests/unit/experience/test_rollouts.py::test_async_rollout_manager - --deselect=tests/unit/experience/test_rollouts.py::test_async_rollout_manager_truncation - --deselect=tests/unit/experience/test_rollouts.py::test_async_nemo_gym_rollout_manager + --deselect=tests/unit/experience/test_rollout_manager.py::test_async_rollout_manager + --deselect=tests/unit/experience/test_rollout_manager.py::test_async_rollout_manager_truncation + --deselect=tests/unit/experience/test_rollout_manager.py::test_async_nemo_gym_rollout_manager ########################################################################### # ENVIRONMENTS diff --git a/tests/unit/experience/test_rollout_manager.py b/tests/unit/experience/test_rollout_manager.py index 0aa24a4364..b43cbb7157 100644 --- a/tests/unit/experience/test_rollout_manager.py +++ b/tests/unit/experience/test_rollout_manager.py @@ -73,6 +73,7 @@ class _FakeBuffer: def __init__(self) -> None: self.reserve_calls: list[int] = [] # weight_versions passed to reserve self.commit_calls: list[tuple[str, object, int, int]] = [] + self.remove_calls: list[str] = [] # reserve(weight_version=X) -> group_id; commit fills the slot. self._slots: list[str] = [] @@ -101,6 +102,12 @@ async def commit( ) return record + async def remove_group(self, group_id: str, *, remove_in_dp: bool = False) -> int: + del remove_in_dp + self.remove_calls.append(group_id) + self._slots.remove(group_id) + return 1 + class _FakeImpl: """Stand-in for AsyncRolloutImpl that returns a sentinel record.""" @@ -127,6 +134,21 @@ def _make_manager(buffer: _FakeBuffer, impl: _FakeImpl) -> RolloutManager: class TestGenerateAndPushFlow: + def test_rollout_failure_removes_reserved_group(self): + async def _fail_rollout(_sample): + raise RuntimeError("injected rollout failure") + + buf = _FakeBuffer() + mgr = _make_manager(buf, _FakeImpl(on_run=_fail_rollout)) + + with pytest.raises(RuntimeError, match="injected rollout failure"): + _run(mgr.generate_and_push({"prompt": "p"})) + + assert len(buf.reserve_calls) == 1 + assert len(buf.remove_calls) == 1 + assert buf._slots == [] + assert buf.commit_calls == [] + def test_reserves_then_runs_then_commits(self): events: list[str] = [] buf = _FakeBuffer() @@ -263,7 +285,7 @@ def test_rollout_manager_raises_without_impl_params(): """RolloutManager raises AssertionError when required params are missing.""" common = { "tokenizer": None, - "env_handles": {}, + "task_to_env": {}, "num_generations_per_prompt": 1, "max_seq_len": 1, } @@ -352,7 +374,7 @@ def test_async_rollout_manager( - rollout_metrics has the expected keys with correct types - completions hold independent (not aliased) message_log objects """ - vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + vllm_generation, tokenizer, task_to_env, _, _ = multi_step_setup_vllm_async input_sample = single_multi_step_calculator_input_sample num_generations = 2 max_seq_len = 1024 @@ -361,7 +383,7 @@ def test_async_rollout_manager( manager = RolloutManager( use_nemo_gym=False, tokenizer=tokenizer, - env_handles=env_handles, + task_to_env=task_to_env, num_generations_per_prompt=num_generations, max_seq_len=max_seq_len, max_rollout_turns=max_rollout_turns, @@ -411,7 +433,7 @@ def test_async_rollout_manager_truncation( single_multi_step_calculator_input_sample, ): """Small max_seq_len forces truncation and truncation_rate=1.0.""" - vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + vllm_generation, tokenizer, task_to_env, _, _ = multi_step_setup_vllm_async input_sample = single_multi_step_calculator_input_sample num_generations = 2 max_seq_len = 290 @@ -420,7 +442,7 @@ def test_async_rollout_manager_truncation( manager = RolloutManager( use_nemo_gym=False, tokenizer=tokenizer, - env_handles=env_handles, + task_to_env=task_to_env, num_generations_per_prompt=num_generations, max_seq_len=max_seq_len, max_rollout_turns=max_rollout_turns, @@ -450,7 +472,7 @@ def test_async_rollout_manager_matches_original( TODO: remove this test together with run_async_multi_turn_rollout when the legacy path is deleted. """ - vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + vllm_generation, tokenizer, task_to_env, _, _ = multi_step_setup_vllm_async input_sample = single_multi_step_calculator_input_sample num_generations = 2 max_seq_len = 1024 @@ -477,7 +499,7 @@ def test_async_rollout_manager_matches_original( policy_generation=vllm_generation, input_batch=batch, tokenizer=tokenizer, - task_to_env=env_handles, + task_to_env=task_to_env, max_seq_len=max_seq_len, max_rollout_turns=max_rollout_turns, ) @@ -485,7 +507,7 @@ def test_async_rollout_manager_matches_original( manager = RolloutManager( use_nemo_gym=False, tokenizer=tokenizer, - env_handles=env_handles, + task_to_env=task_to_env, num_generations_per_prompt=num_generations, max_seq_len=max_seq_len, max_rollout_turns=max_rollout_turns, @@ -619,7 +641,7 @@ def test_async_nemo_gym_rollout_manager( manager = RolloutManager( use_nemo_gym=True, tokenizer=nemo_gym_tokenizer, - env_handles={"nemo_gym": nemo_gym}, + task_to_env={"nemo_gym": nemo_gym}, num_generations_per_prompt=num_generations, max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], generation_config=nemo_gym_vllm_generation.cfg, @@ -730,7 +752,7 @@ def test_async_nemo_gym_rollout_manager_matches_original( manager = RolloutManager( use_nemo_gym=True, tokenizer=nemo_gym_tokenizer, - env_handles={"nemo_gym": nemo_gym}, + task_to_env={"nemo_gym": nemo_gym}, num_generations_per_prompt=num_generations, max_seq_len=nemo_gym_vllm_generation.cfg["vllm_cfg"]["max_model_len"], generation_config=nemo_gym_vllm_generation.cfg, diff --git a/tests/unit/single_controller/test_rollout_pump.py b/tests/unit/single_controller/test_rollout_pump.py index 87054c31fd..b2e82685f2 100644 --- a/tests/unit/single_controller/test_rollout_pump.py +++ b/tests/unit/single_controller/test_rollout_pump.py @@ -16,9 +16,12 @@ from __future__ import annotations +import asyncio import time +from types import SimpleNamespace from typing import Any +import pytest import ray import torch from tensordict import TensorDict @@ -45,7 +48,7 @@ ) _PARTITION_ID = "rollout_data" -# TQReplayBuffer.add tensorizes each PromptGroupRecord and writes +# TQReplayBuffer.commit tensorizes each PromptGroupRecord and writes # ``generations_per_prompt`` training rows directly to TQ. _BULK_FIELDS = [ "input_ids", @@ -151,18 +154,218 @@ def _padded(td: TensorDict) -> TensorDict: return TensorDict(out, batch_size=td.batch_size) +@pytest.mark.parametrize( + ("force_in_order", "expected_target_steps"), + [ + (False, [None, None]), + (True, [0, 1]), + ], +) +def test_rollout_pump_stamps_target_steps( + force_in_order: bool, + expected_target_steps: list[int | None], +) -> None: + class _RecordingBuffer: + def __init__(self) -> None: + self.target_step_list: list[int | None] = [] + + def reserve(self, *, target_step: int | None) -> None: + self.target_step_list.append(target_step) + + class _RecordingRolloutManager: + def __init__(self, buffer: _RecordingBuffer) -> None: + self._buffer = buffer + + async def generate_and_push( + self, prompt: Any, *, target_step: int | None = None + ) -> None: + del prompt + self._buffer.reserve(target_step=target_step) + + buffer = _RecordingBuffer() + controller_cls = SingleControllerActor.__ray_metadata__.modified_class + ctrl = object.__new__(controller_cls) + ctrl._cfg = SimpleNamespace( + max_inflight_prompts=2, + over_sampling=False, + max_weight_staleness_versions=1, + force_in_order=force_in_order, + diagnostics=False, + max_num_epochs=1, + ) + ctrl._rollout_manager = _RecordingRolloutManager(buffer) + prompt_batch = BatchedDataDict( + {"message_log": [[{"role": "user", "content": "prompt"}]]} + ) + ctrl._dataloader = [prompt_batch, prompt_batch] + ctrl._rollout_permitted = asyncio.Event() + ctrl._rollout_permitted.set() + ctrl._buffer_capacity = asyncio.Semaphore(2) + ctrl._inflight_rollouts = 0 + ctrl._dispatched_rollouts = set() + ctrl._max_rollout_version = -1 + ctrl._trainer_version = 0 + ctrl._current_epoch = 0 + + asyncio.run(ctrl._rollout_pump()) + + assert buffer.target_step_list == expected_target_steps + + +def test_rollout_pump_failure_cancels_sibling_and_releases_capacity() -> None: + class _FailingRolloutManager: + def __init__(self) -> None: + self._started = 0 + self._both_started = asyncio.Event() + self.sibling_cancelled = False + + async def generate_and_push( + self, prompt: Any, *, target_step: int | None = None + ) -> None: + del target_step + self._started += 1 + if self._started == 2: + self._both_started.set() + await self._both_started.wait() + + content = prompt["message_log"][0]["content"] + if content == "fail": + raise RuntimeError("injected rollout failure") + + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + self.sibling_cancelled = True + raise + + async def _main() -> None: + manager = _FailingRolloutManager() + controller_cls = SingleControllerActor.__ray_metadata__.modified_class + ctrl = object.__new__(controller_cls) + ctrl._cfg = SimpleNamespace( + max_inflight_prompts=2, + over_sampling=True, + max_weight_staleness_versions=1, + force_in_order=False, + diagnostics=False, + max_num_epochs=1, + ) + ctrl._rollout_manager = manager + ctrl._dataloader = [ + BatchedDataDict( + { + "message_log": [ + [{"role": "user", "content": "fail"}], + [{"role": "user", "content": "sibling"}], + ] + } + ) + ] + ctrl._rollout_permitted = asyncio.Event() + ctrl._rollout_permitted.set() + ctrl._buffer_capacity = asyncio.Semaphore(2) + ctrl._inflight_rollouts = 0 + ctrl._dispatched_rollouts = set() + ctrl._max_rollout_version = -1 + ctrl._trainer_version = 0 + ctrl._current_epoch = 0 + + with pytest.raises(ExceptionGroup) as exc_info: + await asyncio.wait_for(ctrl._rollout_pump(), timeout=1.0) + + assert exc_info.value.subgroup(RuntimeError) is not None + assert manager.sibling_cancelled + assert ctrl._inflight_rollouts == 0 + assert ctrl._buffer_capacity._value == 2 + assert ctrl._dispatched_rollouts == set() + + asyncio.run(_main()) + + +def test_rollout_pump_releases_permits_when_child_never_starts(monkeypatch) -> None: + class _NeverCalledRolloutManager: + async def generate_and_push( + self, prompt: Any, *, target_step: int | None = None + ) -> None: + del prompt, target_step + raise AssertionError("cancelled child unexpectedly started") + + class _CancelBeforeStartTaskGroup: + def __init__(self) -> None: + self._tasks: list[asyncio.Task[None]] = [] + + async def __aenter__(self) -> _CancelBeforeStartTaskGroup: + return self + + async def __aexit__(self, *args: Any) -> bool: + await asyncio.gather(*self._tasks, return_exceptions=True) + return False + + def create_task(self, coro: Any) -> asyncio.Task[None]: + task = asyncio.create_task(coro) + task.cancel() + self._tasks.append(task) + return task + + real_semaphore = asyncio.Semaphore + created_semaphores: list[asyncio.Semaphore] = [] + + def _recording_semaphore(value: int) -> asyncio.Semaphore: + semaphore = real_semaphore(value) + created_semaphores.append(semaphore) + return semaphore + + monkeypatch.setattr(asyncio, "Semaphore", _recording_semaphore) + monkeypatch.setattr(asyncio, "TaskGroup", _CancelBeforeStartTaskGroup) + + async def _main() -> None: + controller_cls = SingleControllerActor.__ray_metadata__.modified_class + ctrl = object.__new__(controller_cls) + ctrl._cfg = SimpleNamespace( + max_inflight_prompts=1, + over_sampling=True, + max_weight_staleness_versions=1, + force_in_order=False, + diagnostics=False, + max_num_epochs=1, + ) + ctrl._rollout_manager = _NeverCalledRolloutManager() + ctrl._dataloader = [ + BatchedDataDict({"message_log": [[{"role": "user", "content": "prompt"}]]}) + ] + ctrl._rollout_permitted = asyncio.Event() + ctrl._rollout_permitted.set() + ctrl._buffer_capacity = real_semaphore(1) + ctrl._inflight_rollouts = 0 + ctrl._dispatched_rollouts = set() + ctrl._max_rollout_version = -1 + ctrl._trainer_version = 0 + ctrl._current_epoch = 0 + + await ctrl._rollout_pump() + await asyncio.sleep(0) + + assert ctrl._buffer_capacity._value == 1 + assert created_semaphores[0]._value == 1 + assert ctrl._inflight_rollouts == 0 + assert ctrl._dispatched_rollouts == set() + + asyncio.run(_main()) + + +@pytest.mark.vllm def test_rollout_pump_writes_expected_tq_data( multi_step_setup_vllm_async, # noqa: F811 single_multi_step_calculator_input_sample, # noqa: F811 tmp_path, ): """SC._rollout_pump writes max_rollout_prompts * num_generations rows to TQ with the expected fields and tags.""" - vllm_generation, tokenizer, env_handles, _, _ = multi_step_setup_vllm_async + vllm_generation, tokenizer, task_to_env, _, _ = multi_step_setup_vllm_async input_sample = single_multi_step_calculator_input_sample num_generations = 2 max_rollout_prompts = 2 - # TQReplayBuffer.add writes ``num_generations`` training rows per prompt. + # TQReplayBuffer.commit writes ``num_generations`` training rows per prompt. expected_samples = max_rollout_prompts * num_generations max_seq_len = 1024 max_rollout_turns = input_sample["extra_env_info"]["max_steps"] + 1 @@ -206,7 +409,7 @@ def test_rollout_pump_writes_expected_tq_data( ) rollout_manager = RolloutManager( tokenizer=tokenizer, - env_handles=env_handles, + task_to_env=task_to_env, num_generations_per_prompt=num_generations, max_seq_len=max_seq_len, max_rollout_turns=max_rollout_turns, diff --git a/tests/unit/single_controller/test_tq_replay_buffer.py b/tests/unit/single_controller/test_tq_replay_buffer.py index 12e8249a9f..76cb919849 100644 --- a/tests/unit/single_controller/test_tq_replay_buffer.py +++ b/tests/unit/single_controller/test_tq_replay_buffer.py @@ -102,6 +102,20 @@ def depth(self) -> int: return len(self._rows) +class FailAfterPutDataPlaneClient(FakeDataPlaneClient): + """Write all rows, then fail to simulate a partial-success RPC.""" + + def put_samples( + self, + sample_ids: list[str], + partition_id: str, + fields: Any = None, + tags: list[dict[str, Any]] | None = None, + ) -> KVBatchMeta: + super().put_samples(sample_ids, partition_id, fields, tags) + raise RuntimeError("injected put failure") + + def _run(coro): return asyncio.run(coro) @@ -141,6 +155,29 @@ def _add_group( class TestTQReplayBufferReserveCommit: + def test_commit_clears_rows_when_put_raises_after_writing(self): + dp = FailAfterPutDataPlaneClient() + buf = _make_buffer(dp) + group_id = buf.reserve(weight_version=3) + + with pytest.raises(RuntimeError, match="injected put failure"): + _run( + buf.commit( + group_id, + _make_record(), + start_weight_version=3, + end_weight_version=3, + ) + ) + + assert dp.depth() == 0 + assert dp.clear_calls == [dp.put_calls[0]["sample_ids"]] + # commit() rolls back DataPlane rows; generate_and_push() owns removal + # of the reserved buffer slot. + assert buf.size() == 1 + assert buf.ready_list == [False] + assert buf.meta_list == [None] + def test_reserve_appends_placeholder_unready(self): dp = FakeDataPlaneClient() buf = _make_buffer(dp) From 48184ad00ccd68474ac2dd4118052309876d156f Mon Sep 17 00:00:00 2001 From: ruit Date: Tue, 21 Jul 2026 22:07:58 -0700 Subject: [PATCH 10/11] fix(sc): preserve cancellation during replay buffer rollback Signed-off-by: ruit --- nemo_rl/algorithms/async_utils/replay_buffer.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/nemo_rl/algorithms/async_utils/replay_buffer.py b/nemo_rl/algorithms/async_utils/replay_buffer.py index 61b1f11f8a..1bf9c5ae93 100644 --- a/nemo_rl/algorithms/async_utils/replay_buffer.py +++ b/nemo_rl/algorithms/async_utils/replay_buffer.py @@ -755,6 +755,8 @@ async def commit( partition_id=self._partition_id, ) except BaseException as rollback_error: + if isinstance(commit_error, asyncio.CancelledError): + raise commit_error from rollback_error raise BaseExceptionGroup( f"commit and rollback both failed for group_id={group_id!r}", [commit_error, rollback_error], From 62a759a065b3a53447b9fe9e99422fe4edce9974 Mon Sep 17 00:00:00 2001 From: ruit Date: Wed, 22 Jul 2026 02:07:46 -0700 Subject: [PATCH 11/11] test: update rollout manager stream fixture Signed-off-by: ruit --- tests/unit/experience/test_rollouts.py | 30 +++++++++++++++++++++++--- 1 file changed, 27 insertions(+), 3 deletions(-) diff --git a/tests/unit/experience/test_rollouts.py b/tests/unit/experience/test_rollouts.py index 04d42c006f..038384a595 100644 --- a/tests/unit/experience/test_rollouts.py +++ b/tests/unit/experience/test_rollouts.py @@ -1343,8 +1343,30 @@ class _Stream: def __init__(self): self.values = iter( [ - _ReadyRef((1, {"value": "second"}, None)), - _ReadyRef((0, {"value": "first"}, {"remote_time": 2.0})), + _ReadyRef( + ( + 1, + { + "value": "second", + "input_message_log": [ + {"role": "user", "token_ids": [1]} + ], + }, + None, + ) + ), + _ReadyRef( + ( + 0, + { + "value": "first", + "input_message_log": [ + {"role": "user", "token_ids": [1]} + ], + }, + {"remote_time": 2.0}, + ) + ), ] ) @@ -1377,7 +1399,7 @@ def remote(self, inputs, tokenizer, timer_prefix): "agent": agent, } - completions, metrics = asyncio.run( + completions, prompt_message_log, metrics = asyncio.run( manager._run_rollouts( inputs=[ {"agent_ref": {"name": "agent"}}, @@ -1389,6 +1411,8 @@ def remote(self, inputs, tokenizer, timer_prefix): ) assert completions == ["first", "second"] + assert prompt_message_log[0]["role"] == "user" + torch.testing.assert_close(prompt_message_log[0]["token_ids"], torch.tensor([1])) assert metrics == { "completion_count": 2, "agent": "agent",