From c8a1084f6219f59c91da25aa0ee13f4170980bf3 Mon Sep 17 00:00:00 2001 From: Jingwen Gu <75733630+JingwenGu0829@users.noreply.github.com> Date: Tue, 21 Jul 2026 22:40:02 +0000 Subject: [PATCH 01/32] Add SGLang Omni multimodal input adapter --- miles/rollout/generate_hub/sglang_omni.py | 68 ++++++++++ .../generate_utils/generate_endpoint_utils.py | 114 ++++++++++++++-- miles/utils/processing_utils.py | 28 ++++ requirements.txt | 1 + .../rollout/generate_hub/test_sglang_omni.py | 122 ++++++++++++++++++ 5 files changed, 324 insertions(+), 9 deletions(-) create mode 100644 miles/rollout/generate_hub/sglang_omni.py create mode 100644 tests/fast/rollout/generate_hub/test_sglang_omni.py diff --git a/miles/rollout/generate_hub/sglang_omni.py b/miles/rollout/generate_hub/sglang_omni.py new file mode 100644 index 00000000000..ecb8c46e93b --- /dev/null +++ b/miles/rollout/generate_hub/sglang_omni.py @@ -0,0 +1,68 @@ +"""Single-turn rollout using SGLang Omni's processed multimodal contract. + +Select with ``--custom-generate-function-path +miles.rollout.generate_hub.sglang_omni.generate``. +""" + +from miles.rollout.base_types import GenerateFnInput, GenerateFnOutput +from miles.rollout.generate_utils.generate_endpoint_utils import ( + compute_prompt_ids_from_sample, + compute_request_payload, + compute_routing_headers, + multimodal_route_headers, + update_sample_from_response, +) +from miles.utils.http_utils import post +from miles.utils.types import Sample + + +async def generate(input: GenerateFnInput) -> GenerateFnOutput: + args = input.args + sample = input.sample + sampling_params = input.sampling_params + assert sample.status in { + Sample.Status.PENDING, + Sample.Status.ABORTED, + }, f"{sample.status=}" + url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate" + + prompt_ids = compute_prompt_ids_from_sample(input.state, sample) + has_multimodal_inputs = sample.multimodal_inputs and any( + value is not None for value in sample.multimodal_inputs.values() + ) + if has_multimodal_inputs and sample.multimodal_train_inputs is None: + raise ValueError( + "SGLang Omni multimodal rollout requires processor-produced " "sample.multimodal_train_inputs" + ) + + if sample.response: + input_ids = sample.tokens + sampling_params["max_new_tokens"] -= len(sample.tokens) - len(prompt_ids) + assert sampling_params["max_new_tokens"] >= 0 + if sampling_params["max_new_tokens"] == 0: + sample.status = Sample.Status.TRUNCATED + return GenerateFnOutput(samples=sample) + else: + input_ids = prompt_ids + + payload, halt_status = compute_request_payload( + args, + input_ids=input_ids, + sampling_params=sampling_params, + multimodal_inputs=sample.multimodal_inputs, + multimodal_train_inputs=sample.multimodal_train_inputs, + ) + if payload is None: + sample.status = halt_status + return GenerateFnOutput(samples=sample) + + headers = compute_routing_headers(args, sample) or {} + headers.update(multimodal_route_headers(payload.get("multimodal_train_inputs")) or {}) + output = await post(url, payload, headers=headers or None) + await update_sample_from_response( + args, + sample, + payload=payload, + output=output, + ) + return GenerateFnOutput(samples=sample) diff --git a/miles/rollout/generate_utils/generate_endpoint_utils.py b/miles/rollout/generate_utils/generate_endpoint_utils.py index 1de0e66d079..e5b210a51c4 100644 --- a/miles/rollout/generate_utils/generate_endpoint_utils.py +++ b/miles/rollout/generate_utils/generate_endpoint_utils.py @@ -2,31 +2,124 @@ Utils to integrate SGLang's `/generate` endpoint with RL things like Sample. """ +import base64 from copy import deepcopy from typing import Any import numpy as np import pybase64 +import torch from miles.utils.lora import LORA_ADAPTER_NAME, is_lora_enabled -from miles.utils.processing_utils import encode_image_for_rollout_engine +from miles.utils.processing_utils import ( + call_processor, + encode_image_for_rollout_engine, + extract_multimodal_train_inputs, +) from miles.utils.types import Sample +_MULTIMODAL_TENSOR_DTYPES = { + "bool", + "uint8", + "int8", + "int16", + "int32", + "int64", + "float16", + "bfloat16", + "float32", + "float64", +} +_SGLANG_OMNI_MULTIMODAL_TENSOR_NAMES = frozenset( + { + "pixel_values", + "image_grid_thw", + "input_features", + "feature_attention_mask", + "audio_feature_lengths", + "pixel_values_videos", + "video_grid_thw", + "video_second_per_grid", + } +) + + +def _multimodal_modalities(tensors: dict[str, torch.Tensor]) -> list[str]: + modalities = [] + if "pixel_values" in tensors: + modalities.append("image") + if "input_features" in tensors: + modalities.append("audio") + if "pixel_values_videos" in tensors: + modalities.append("video") + return modalities + + +def serialize_multimodal_train_inputs( + multimodal_train_inputs: dict[str, Any], +) -> dict[str, Any]: + """Encode the processor tensor kwargs shared with SGLang Omni.""" + if not multimodal_train_inputs: + raise ValueError("multimodal_train_inputs must not be empty") + unknown = sorted(set(multimodal_train_inputs) - _SGLANG_OMNI_MULTIMODAL_TENSOR_NAMES) + if unknown: + raise ValueError("unsupported SGLang Omni processor tensor fields: " + ", ".join(unknown)) + + tensors: dict[str, dict[str, Any]] = {} + tensor_values: dict[str, torch.Tensor] = {} + for name, value in multimodal_train_inputs.items(): + if not isinstance(value, torch.Tensor): + raise TypeError( + "multimodal_train_inputs must contain only tensors; " f"{name!r} has type {type(value).__name__}" + ) + tensor = value.detach().to(device="cpu").contiguous() + dtype = str(tensor.dtype).removeprefix("torch.") + if dtype not in _MULTIMODAL_TENSOR_DTYPES: + raise TypeError(f"unsupported multimodal tensor dtype for {name!r}: {dtype}") + raw = tensor.reshape(-1).view(torch.uint8).numpy().tobytes() + tensors[name] = { + "dtype": dtype, + "shape": list(tensor.shape), + "data": base64.b64encode(raw).decode("ascii"), + } + tensor_values[name] = tensor + + modalities = _multimodal_modalities(tensor_values) + if not modalities: + raise ValueError("multimodal_train_inputs does not contain an image, audio, or video " "encoder tensor") + return {"version": 1, "modalities": modalities, "tensors": tensors} + + +def multimodal_route_headers( + serialized_inputs: dict[str, Any] | None, +) -> dict[str, str] | None: + """Route large JSON bundles without asking the router to scan tensor data.""" + if not serialized_inputs: + return None + capabilities = [f"{modality}_input" for modality in serialized_inputs.get("modalities", [])] + if not capabilities: + return None + return {"x-sglang-omni-route-capabilities": ",".join(capabilities)} + + # Make this an isolated function because users may want to compute their own def compute_prompt_ids_from_sample(state, sample, tools=None): prompt = sample.prompt - if state.processor and sample.multimodal_inputs and any(v is not None for v in sample.multimodal_inputs.values()): - processor_output = state.processor(text=prompt, **sample.multimodal_inputs) + if ( + state.processor + and sample.multimodal_inputs + and any(value is not None for value in sample.multimodal_inputs.values()) + ): + processor_output = call_processor(state.processor, prompt, sample.multimodal_inputs) prompt_ids = processor_output["input_ids"][0] - # TODO shall we move it to other places? then can make this function immutable - sample.multimodal_train_inputs = { - k: v for k, v in processor_output.items() if k not in ["input_ids", "attention_mask"] - } or None + sample.multimodal_train_inputs = extract_multimodal_train_inputs(processor_output) - return prompt_ids + if hasattr(prompt_ids, "tolist"): + prompt_ids = prompt_ids.tolist() + return [int(token_id) for token_id in prompt_ids] else: if not isinstance(prompt, str): prompt = state.tokenizer.apply_chat_template( @@ -56,6 +149,7 @@ def compute_request_payload( input_ids: list[int], sampling_params: dict, multimodal_inputs: dict | None = None, + multimodal_train_inputs: dict[str, Any] | None = None, ) -> tuple[dict[str, Any] | None, Sample.Status | None]: sampling_params = deepcopy(sampling_params) max_new_tokens = sampling_params.pop("max_new_tokens", args.rollout_max_response_len) @@ -73,7 +167,9 @@ def compute_request_payload( } if is_lora_enabled(args): payload["lora_path"] = LORA_ADAPTER_NAME - if image_data := (multimodal_inputs or {}).get("images"): + if multimodal_train_inputs: + payload["multimodal_train_inputs"] = serialize_multimodal_train_inputs(multimodal_train_inputs) + elif image_data := (multimodal_inputs or {}).get("images"): payload["image_data"] = [encode_image_for_rollout_engine(image) for image in image_data] return payload, None diff --git a/miles/utils/processing_utils.py b/miles/utils/processing_utils.py index 855a06a8fb2..fbcd8695ac1 100644 --- a/miles/utils/processing_utils.py +++ b/miles/utils/processing_utils.py @@ -137,6 +137,23 @@ def call_processor(processor, text, multimodal_inputs: dict | None = None): return processor(text=text, **kwargs) +def extract_multimodal_train_inputs(processor_output): + """Normalize processor kwargs to the tensor-only Megatron contract.""" + import torch + + result = {} + for name, value in processor_output.items(): + if name in ("input_ids", "attention_mask"): + continue + if not isinstance(value, torch.Tensor): + try: + value = torch.as_tensor(value) + except (TypeError, ValueError, RuntimeError) as exc: + raise TypeError(f"processor output {name!r} cannot be converted to a tensor") from exc + result[name] = value + return result or None + + def load_processor(name_or_path: str, **kwargs): try: proc = AutoProcessor.from_pretrained(name_or_path, **kwargs) @@ -153,6 +170,17 @@ def load_processor(name_or_path: str, **kwargs): def process_vision_info(prompt, processor): # TODO: temporary solution, will write image utils for miles later + if getattr(processor, "audio_token", None) is not None: + from qwen_omni_utils import process_mm_info + + image_patch_size = getattr(processor.image_processor, "patch_size", DEFAULT_PATCH_SIZE) + audios, images, videos = process_mm_info( + prompt, + use_audio_in_video=False, + image_patch_size=image_patch_size, + ) + return {"audio": audios, "images": images, "videos": videos} + from qwen_vl_utils import process_vision_info as qwen_process_vision_info if hasattr(processor.image_processor, "patch_size"): diff --git a/requirements.txt b/requirements.txt index 74b84a2977f..523877bc4d8 100644 --- a/requirements.txt +++ b/requirements.txt @@ -16,6 +16,7 @@ pybase64 pylatexenc pytest-asyncio pyyaml +qwen_omni_utils==0.0.9 # for Qwen Omni audio/video inputs qwen_vl_utils # for VLM ray[default] ring_flash_attn; platform_system == "Linux" diff --git a/tests/fast/rollout/generate_hub/test_sglang_omni.py b/tests/fast/rollout/generate_hub/test_sglang_omni.py new file mode 100644 index 00000000000..4a08d9d31ee --- /dev/null +++ b/tests/fast/rollout/generate_hub/test_sglang_omni.py @@ -0,0 +1,122 @@ +import asyncio +import base64 +from types import SimpleNamespace + +import pytest +import torch + +from miles.rollout.generate_utils.generate_endpoint_utils import serialize_multimodal_train_inputs +from miles.utils.types import Sample + + +def test_serialize_audio_video_processor_tensors(): + inputs = { + "input_features": torch.arange(6, dtype=torch.float32).reshape(1, 2, 3), + "feature_attention_mask": torch.ones((1, 2), dtype=torch.long), + "pixel_values_videos": torch.arange(12, dtype=torch.bfloat16).reshape(2, 2, 3), + "video_grid_thw": torch.tensor([[1, 2, 3]], dtype=torch.long), + "video_second_per_grid": torch.tensor([0.5], dtype=torch.float32), + } + + bundle = serialize_multimodal_train_inputs(inputs) + + assert bundle["version"] == 1 + assert bundle["modalities"] == ["audio", "video"] + assert set(bundle["tensors"]) == set(inputs) + for name, tensor in inputs.items(): + encoded = bundle["tensors"][name] + assert encoded["dtype"] == str(tensor.dtype).removeprefix("torch.") + assert encoded["shape"] == list(tensor.shape) + assert base64.b64decode(encoded["data"]) == ( + tensor.contiguous().reshape(-1).view(torch.uint8).numpy().tobytes() + ) + + +def test_qwen_omni_media_extraction_and_tensor_normalization(monkeypatch): + from miles.utils import processing_utils + + qwen_omni_utils = pytest.importorskip("qwen_omni_utils") + prompt = [{"role": "user", "content": [{"type": "audio", "audio": "a.wav"}]}] + captured = {} + + def fake_process_mm_info(conversations, **kwargs): + captured["conversations"], captured["kwargs"] = conversations, kwargs + return ["audio samples"], ["image"], ["video frames"] + + monkeypatch.setattr(qwen_omni_utils, "process_mm_info", fake_process_mm_info) + processor = SimpleNamespace(audio_token="