Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions nemo_rl/experience/rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 = {
Comment thread
rohitrango marked this conversation as resolved.
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"]
Expand Down Expand Up @@ -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
Expand All @@ -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(
Expand Down Expand Up @@ -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.

Expand All @@ -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
Expand Down
98 changes: 96 additions & 2 deletions tests/unit/experience/test_rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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",
Expand All @@ -147,13 +149,22 @@ async def generate_async(self, data):
]

assistant_message, input_lengths, _ = _run(
manager._generate_response(message_log, ["<stop>"])
manager._generate_response(
Comment thread
jinglinglingling marked this conversation as resolved.
message_log,
["<stop>"],
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"] == [["<stop>"]]
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(
Expand All @@ -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": "<image><audio><video>",
"vllm_images": [image],
"vllm_audios": [audio],
"vllm_videos": [video],
},
traj_idx=0,
)
)

assert len(calls) == 2
assert calls[0]["vllm_content"] == "<image><audio><video>"
assert calls[1]["vllm_content"] is None
for key, expected in (
("vllm_images", image),
("vllm_audios", audio),
("vllm_videos", video),
):
assert calls[0][key][0] is expected
assert calls[1][key][0] is expected


class _FakeBuffer:
"""Minimal TQReplayBuffer stand-in that records reserve/commit calls."""

Expand Down
Loading