diff --git a/nemo_rl/experience/rollout_manager.py b/nemo_rl/experience/rollout_manager.py index 5c448caee8d..4de54017f09 100644 --- a/nemo_rl/experience/rollout_manager.py +++ b/nemo_rl/experience/rollout_manager.py @@ -39,6 +39,7 @@ ) from nemo_rl.data.interfaces import DatumSpec, LLMMessageLogType from nemo_rl.data.llm_message_utils import batched_message_log_to_flat_message +from nemo_rl.data.multimodal_utils import NATIVE_MULTIMODAL_KEYS from nemo_rl.data_plane.schema import MASK_SAMPLE from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.environments.interfaces import EnvironmentInterface @@ -572,6 +573,12 @@ async def _run_single_rollout( ) -> tuple[Completion, dict]: """Run one multi-turn rollout for a single generation index.""" current_message_log = copy.deepcopy(input_sample["message_log"]) + input_sample_data: Mapping[str, Any] = input_sample + native_generation_data = { + key: input_sample_data[key] + for key in NATIVE_MULTIMODAL_KEYS + if key in input_sample_data + } current_extra_env_info = copy.deepcopy(input_sample["extra_env_info"]) current_stop_strings = input_sample.get("stop_strings", None) task_name = input_sample["task_name"] @@ -599,6 +606,11 @@ async def _run_single_rollout( break turn_count += 1 + turn_native_generation_data = dict(native_generation_data) + # Raw processor content describes only the original conversation. + # Later turns keep the media but use the updated pre-tokenized prefix. + if turn_count > 1 and "vllm_content" in turn_native_generation_data: + turn_native_generation_data["vllm_content"] = None # Generate response for this sample using async generation. # A failure here must not be absorbed: returning a partial completion @@ -612,6 +624,7 @@ async def _run_single_rollout( ) = await self._generate_response( current_message_log, current_stop_strings, + native_generation_data=turn_native_generation_data, ) except Exception as e: raise _classify_generation_failure( @@ -738,6 +751,8 @@ async def _generate_response( self, message_log: list[dict], stop_strings: list[str] | None, + *, + native_generation_data: dict[str, Any] | None = None, ) -> tuple[dict, torch.Tensor, dict[str, Any]]: """Generate a single-turn response for one sample. @@ -762,6 +777,12 @@ async def _generate_response( generation_input_data.update( flat_messages.get_multimodal_dict(as_tensors=False) ) + if native_generation_data: + # This method handles one sample; vLLM's formatter expects batched + # native content/media side channels. + generation_input_data.update( + {key: [value] for key, value in native_generation_data.items()} + ) # Generate response # TODO: update generate_async to return a single item directly diff --git a/tests/unit/experience/test_rollout_manager.py b/tests/unit/experience/test_rollout_manager.py index c0a4bb5bd42..22e5c9bdc47 100644 --- a/tests/unit/experience/test_rollout_manager.py +++ b/tests/unit/experience/test_rollout_manager.py @@ -45,6 +45,7 @@ from nemo_rl.data.multimodal_utils import PackedTensor from nemo_rl.data.processors import nemo_gym_data_processor from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.environments.interfaces import EnvironmentReturn from nemo_rl.experience.failures import GenerationUnavailable from nemo_rl.experience.interfaces import ( NEMO_GYM_GROUP_ATTEMPT_KEY, @@ -101,7 +102,7 @@ async def apply(): return _run(apply()) -def test_generate_response_forwards_message_log_media_to_generation() -> None: +def test_generate_response_forwards_all_vllm_media_to_generation() -> None: captured: dict[str, BatchedDataDict] = {} class _Generation: @@ -131,6 +132,7 @@ async def generate_async(self, data): manager._deadline_registry = None pixel_values = PackedTensor(torch.ones(2, 3, 4, 4), dim_to_pack=0) imgs_sizes = PackedTensor(torch.tensor([[4, 4], [4, 4]]), dim_to_pack=0) + native_image = object() message_log = [ { "role": "user", @@ -147,13 +149,22 @@ async def generate_async(self, data): ] assistant_message, input_lengths, _ = _run( - manager._generate_response(message_log, [""]) + manager._generate_response( + message_log, + [""], + native_generation_data={ + "vllm_content": "rendered prompt", + "vllm_images": [native_image], + }, + ) ) generation_data = captured["data"] assert generation_data["input_ids"].tolist() == [[1, 2, 3, 4, 5]] assert generation_data["input_lengths"].tolist() == [5] assert generation_data["stop_strings"] == [[""]] + assert generation_data["vllm_content"] == ["rendered prompt"] + assert generation_data["vllm_images"] == [[native_image]] assert isinstance(generation_data["pixel_values"], PackedTensor) assert isinstance(generation_data["imgs_sizes"], PackedTensor) assert torch.equal( @@ -167,6 +178,89 @@ async def generate_async(self, data): assert assistant_message["token_ids"].tolist() == [42] +def test_run_single_rollout_preserves_native_media_across_turns(monkeypatch) -> None: + calls: list[dict] = [] + image = object() + audio = object() + video = object() + + async def generate_response( + _message_log, + _stop_strings, + *, + native_generation_data=None, + ): + calls.append(dict(native_generation_data or {})) + return ( + { + "role": "assistant", + "content": "answer", + "token_ids": torch.tensor([42]), + "generation_logprobs": torch.tensor([0.0]), + }, + torch.tensor(1), + {}, + ) + + def calculate_rewards(_batch, _task_to_env): + return EnvironmentReturn( + observations=[{"role": "user", "content": "next"}], + metadata=[None], + next_stop_strings=[None], + rewards=torch.tensor([0.0]), + terminateds=torch.tensor([len(calls) == 2]), + answers=[None], + ) + + class _Tokenizer: + def __call__(self, *_args, **_kwargs): + return SimpleNamespace(input_ids=torch.tensor([[7]])) + + monkeypatch.setattr( + "nemo_rl.experience.rollout_manager.calculate_rewards", calculate_rewards + ) + manager = object.__new__(AsyncRolloutImpl) + manager._generate_response = generate_response + manager._tokenizer = _Tokenizer() + manager._task_to_env = {} + manager._max_seq_len = 32 + manager._max_rollout_turns = 2 + manager._timeouts = SimpleNamespace(env_s=10.0) + + _run( + manager._run_single_rollout( + { + "idx": 0, + "message_log": [ + { + "role": "user", + "content": "", + "token_ids": torch.tensor([1]), + } + ], + "extra_env_info": None, + "task_name": "vlm", + "vllm_content": "