Skip to content
Closed
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
32 changes: 30 additions & 2 deletions miles_plugins/omni/omni_generate_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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 = {
Expand Down
12 changes: 9 additions & 3 deletions miles_plugins/omni/rollout_contract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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


Expand Down
56 changes: 52 additions & 4 deletions tests/fast/test_omni_generate_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down
17 changes: 17 additions & 0 deletions tests/fast/test_omni_rollout_contract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ------------------------------------------------------------------


Expand Down