From 2bb044ee2918b392842bdb87c0051e2f07778e62 Mon Sep 17 00:00:00 2001 From: Linglin Jing <50938792+jinglinglingling@users.noreply.github.com> Date: Fri, 11 Sep 2026 19:09:33 -0700 Subject: [PATCH 1/3] fix(vlm): preserve native media in async rollouts Signed-off-by: Linglin Jing <50938792+jinglinglingling@users.noreply.github.com> --- nemo_rl/experience/rollout_manager.py | 20 +++++++++++++++++++ tests/unit/experience/test_rollout_manager.py | 14 +++++++++++-- 2 files changed, 32 insertions(+), 2 deletions(-) diff --git a/nemo_rl/experience/rollout_manager.py b/nemo_rl/experience/rollout_manager.py index 21f8cc68fc0..61d17c40f5e 100644 --- a/nemo_rl/experience/rollout_manager.py +++ b/nemo_rl/experience/rollout_manager.py @@ -37,6 +37,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 @@ -570,6 +571,11 @@ 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"]) + native_generation_data = { + key: input_sample[key] + for key in NATIVE_MULTIMODAL_KEYS + if key in input_sample + } 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"] @@ -597,6 +603,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 @@ -610,6 +621,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( @@ -736,6 +748,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. @@ -760,6 +774,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 e6c191df657..8bfa27136c9 100644 --- a/tests/unit/experience/test_rollout_manager.py +++ b/tests/unit/experience/test_rollout_manager.py @@ -101,7 +101,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 +131,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 +148,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( From 188c84137b351656ccb08dd57ef187585185f8ea Mon Sep 17 00:00:00 2001 From: rohitrango Date: Mon, 14 Sep 2026 10:41:52 -0700 Subject: [PATCH 2/3] fix: satisfy pyrefly for native multimodal fields Signed-off-by: rohitrango --- nemo_rl/experience/rollout_manager.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/nemo_rl/experience/rollout_manager.py b/nemo_rl/experience/rollout_manager.py index e137aa339c2..4de54017f09 100644 --- a/nemo_rl/experience/rollout_manager.py +++ b/nemo_rl/experience/rollout_manager.py @@ -573,10 +573,11 @@ 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[key] + key: input_sample_data[key] for key in NATIVE_MULTIMODAL_KEYS - if key in input_sample + 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) From 013d4a4fabc4ff6b393fca7e44543ae4c2e03a8b Mon Sep 17 00:00:00 2001 From: Linglin Jing <50938792+jinglinglingling@users.noreply.github.com> Date: Mon, 14 Sep 2026 18:27:33 -0700 Subject: [PATCH 3/3] test: cover native media across rollout turns Signed-off-by: Linglin Jing <50938792+jinglinglingling@users.noreply.github.com> --- tests/unit/experience/test_rollout_manager.py | 84 +++++++++++++++++++ 1 file changed, 84 insertions(+) diff --git a/tests/unit/experience/test_rollout_manager.py b/tests/unit/experience/test_rollout_manager.py index dd9e253c121..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, @@ -177,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": "