From 9a0731ec71a1092af973860a36f83ac58114def5 Mon Sep 17 00:00:00 2001 From: Jingwen Gu Date: Sat, 11 Jul 2026 23:30:13 +0000 Subject: [PATCH] Forward multimodal inputs from Omni rollouts --- miles_plugins/omni/omni_generate_fn.py | 32 +++++++++++++- miles_plugins/omni/rollout_contract.py | 12 +++-- tests/fast/test_omni_generate_fn.py | 56 ++++++++++++++++++++++-- tests/fast/test_omni_rollout_contract.py | 17 +++++++ 4 files changed, 108 insertions(+), 9 deletions(-) diff --git a/miles_plugins/omni/omni_generate_fn.py b/miles_plugins/omni/omni_generate_fn.py index b193f3f07ad..ebfe8ecf070 100644 --- a/miles_plugins/omni/omni_generate_fn.py +++ b/miles_plugins/omni/omni_generate_fn.py @@ -13,7 +13,11 @@ from miles.rollout.base_types import GenerateFnInput, GenerateFnOutput from miles.rollout.generate_utils.generate_endpoint_utils import compute_prompt_ids_from_sample from miles.utils.http_utils import post -from miles.utils.processing_utils import encode_audios_for_rollout_engine, extract_audio_inputs +from miles.utils.processing_utils import ( + encode_audios_for_rollout_engine, + encode_image_for_rollout_engine, + extract_audio_inputs, +) from miles.utils.types import Sample from .rollout_contract import apply_response_to_sample, build_generate_payload, parse_generate_response @@ -53,7 +57,9 @@ async def __call__(self, input: GenerateFnInput) -> GenerateFnOutput: metadata=_request_metadata(sample), output_modalities=sample.metadata.get("output_modalities"), return_omni_rollout=True, - audio_data=_encode_input_audio(sample), + images=_encode_input_images(sample), + audios=_encode_input_audio(sample), + videos=_encoded_input_videos(sample), ) output = await post(url, payload) @@ -98,6 +104,28 @@ def _encode_input_audio(sample: Sample) -> list[str] | None: return encode_audios_for_rollout_engine(audios) +def _encode_input_images(sample: Sample) -> list[str] | None: + """Reuse Miles' standard VLM serializer for the omni request contract.""" + images = (sample.multimodal_inputs or {}).get("images") + if not images: + return None + return [encode_image_for_rollout_engine(image) for image in images] + + +def _encoded_input_videos(sample: Sample) -> list[str] | None: + """Forward already transport-safe video references and reject other shapes.""" + videos = (sample.multimodal_inputs or {}).get("videos") + if not videos: + return None + if not all( + isinstance(video, str) + and (video.startswith("data:video/") or video.startswith("https://") or video.startswith("http://")) + for video in videos + ): + raise ValueError("Omni rollout video inputs must be data:video URIs or HTTP(S) URLs") + return list(videos) + + def _request_metadata(sample: Sample) -> dict: """Identifiers echoed back by the backend so responses can be matched to a rollout.""" fields = { diff --git a/miles_plugins/omni/rollout_contract.py b/miles_plugins/omni/rollout_contract.py index 61ca0901c24..2fe9514e078 100644 --- a/miles_plugins/omni/rollout_contract.py +++ b/miles_plugins/omni/rollout_contract.py @@ -64,7 +64,9 @@ def build_generate_payload( output_modalities: list[str] | None = None, return_logprob: bool = True, return_omni_rollout: bool = False, - audio_data: list[str] | None = None, + images: list[str] | None = None, + audios: list[str] | None = None, + videos: list[str] | None = None, ) -> dict[str, Any]: """Build an omni ``/generate`` request body from pre-tokenized inputs. @@ -83,8 +85,12 @@ def build_generate_payload( payload["metadata"] = metadata if output_modalities is not None: payload["output_modalities"] = output_modalities - if audio_data is not None: - payload["audio_data"] = audio_data + if images is not None: + payload["images"] = images + if audios is not None: + payload["audios"] = audios + if videos is not None: + payload["videos"] = videos return payload diff --git a/tests/fast/test_omni_generate_fn.py b/tests/fast/test_omni_generate_fn.py index 7ec46e4e263..8a35efa21cf 100644 --- a/tests/fast/test_omni_generate_fn.py +++ b/tests/fast/test_omni_generate_fn.py @@ -10,6 +10,8 @@ from types import SimpleNamespace import numpy as np +import pytest +from PIL import Image import miles_plugins.omni.omni_generate_fn as omni_mod from miles.rollout.base_types import GenerateFnInput @@ -88,7 +90,8 @@ async def fake_post(url, payload, **kwargs): assert payload["return_omni_rollout"] is True assert payload["sampling_params"] == {"temperature": 0.7, "seed": 9, "max_new_tokens": 64} assert payload["metadata"] == {"group_index": 2, "index": 5} - assert "audio_data" not in payload # no input audio on this sample + assert "audios" not in payload # no input audio on this sample + assert "images" not in payload # no input image on this sample assert result_sample.tokens == [1, 2, 3, 10, 11] assert result_sample.response_length == 2 @@ -142,9 +145,54 @@ async def fake_post(url, payload, **kwargs): ) asyncio.run(fn(inp)) - audio_data = captured["payload"]["audio_data"] - assert len(audio_data) == 1 - assert audio_data[0].startswith("data:audio/wav;base64,") + audios = captured["payload"]["audios"] + assert len(audios) == 1 + assert audios[0].startswith("data:audio/wav;base64,") + + +def test_omni_generate_fn_reuses_standard_image_encoder(monkeypatch): + captured = {} + + async def fake_post(url, payload, **kwargs): + captured["payload"] = payload + return _canned_response() + + monkeypatch.setattr(omni_mod, "post", fake_post) + + fn = load_generate_function(_HOOK_PATH) + sample = Sample(prompt="hi") + sample.multimodal_inputs = {"images": [Image.new("RGB", (2, 2), "red")]} + inp = GenerateFnInput( + state=_fake_state(), + sample=sample, + sampling_params={"max_new_tokens": 32}, + evaluation=False, + ) + + asyncio.run(fn(inp)) + images = captured["payload"]["images"] + assert len(images) == 1 + assert images[0].startswith("data:image/png;base64,") + + +def test_omni_generate_fn_rejects_unencoded_video(monkeypatch): + async def fail_post(url, payload, **kwargs): + raise AssertionError("invalid video must fail before HTTP transport") + + monkeypatch.setattr(omni_mod, "post", fail_post) + + fn = load_generate_function(_HOOK_PATH) + sample = Sample(prompt="hi") + sample.multimodal_inputs = {"videos": [np.zeros((2, 2, 3), dtype=np.uint8)]} + inp = GenerateFnInput( + state=_fake_state(), + sample=sample, + sampling_params={"max_new_tokens": 32}, + evaluation=False, + ) + + with pytest.raises(ValueError, match="data:video URIs"): + asyncio.run(fn(inp)) def test_omni_generate_fn_resume_keeps_loss_mask_aligned(monkeypatch): diff --git a/tests/fast/test_omni_rollout_contract.py b/tests/fast/test_omni_rollout_contract.py index df8a24ba238..d392185a033 100644 --- a/tests/fast/test_omni_rollout_contract.py +++ b/tests/fast/test_omni_rollout_contract.py @@ -74,6 +74,23 @@ def test_build_generate_payload_shape_and_metadata(): assert "metadata" not in build_generate_payload([1], {}) +def test_build_generate_payload_uses_canonical_media_names(): + payload = build_generate_payload( + [1, 2, 3], + {}, + images=["data:image/png;base64,SU1H"], + audios=["data:audio/wav;base64,QVVESU8="], + videos=["https://example.test/video.mp4"], + ) + + assert payload["images"] == ["data:image/png;base64,SU1H"] + assert payload["audios"] == ["data:audio/wav;base64,QVVESU8="] + assert payload["videos"] == ["https://example.test/video.mp4"] + assert "image_data" not in payload + assert "audio_data" not in payload + assert "video_data" not in payload + + # --- response parsing ------------------------------------------------------------------