diff --git a/recipes/cosmos3/Cosmos3-Nano.md b/recipes/cosmos3/Cosmos3-Nano.md index 2ee9c2ed7d3..0d51a1a3a5a 100644 --- a/recipes/cosmos3/Cosmos3-Nano.md +++ b/recipes/cosmos3/Cosmos3-Nano.md @@ -47,6 +47,10 @@ mode is selected per request: **`nvidia/Cosmos3-Nano-Policy-DROID`** is served the same way (`domain_name=droid_lerobot`). +- **DROID OpenPI policy server** — serve `nvidia/Cosmos3-Nano-Policy-DROID` and + connect an OpenPI-compatible websocket client to `/v1/realtime/robot/openpi`. + This path returns action chunks directly instead of an mp4. + Action requests can use `input_reference` or `video_reference` for video input. `policy` and `forward_dynamics` can also use an image reference; `inverse_dynamics` requires a video reference. @@ -245,6 +249,39 @@ VIDEO_ID=$(curl -sS -X POST http://localhost:8000/v1/videos \ # poll until status == completed, then: curl -sS "http://localhost:8000/v1/videos/$VIDEO_ID" | jq '.action | {shape, dtype, raw_action_dim, domain_id}' curl -sS -L "http://localhost:8000/v1/videos/$VIDEO_ID/content" -o cosmos3_inverse_dynamics.mp4 + +# DROID OpenPI policy server (websocket action serving). +# Requires cosmos_framework on PYTHONPATH because the pipeline reuses the +# reference RoboLab action transforms. If your checkpoint config already +# includes policy_server_config, omit the stage_overrides file and flag. +cat > cosmos3_droid_openpi_stage_overrides.json <<'JSON' +{ + "0": { + "model_config": { + "policy_server_config": { + "image_resolution": [540, 640], + "n_external_cameras": 2, + "needs_wrist_camera": true, + "needs_stereo_camera": false, + "needs_session_id": true, + "action_space": "joint_position" + } + } + } +} +JSON + +vllm serve nvidia/Cosmos3-Nano-Policy-DROID \ + --omni \ + --host 0.0.0.0 --port 8000 \ + --model-class-name Cosmos3OmniDiffusersPipeline \ + --no-guardrails \ + --stage-overrides "$(cat cosmos3_droid_openpi_stage_overrides.json)" + +# Point an OpenPI websocket client at: +# ws://localhost:8000/v1/realtime/robot/openpi +# The first server message is policy_server_config. Each infer request sends a +# msgpack-numpy observation dict and receives a writable float32 action array. ``` #### Notes @@ -272,6 +309,13 @@ curl -sS -L "http://localhost:8000/v1/videos/$VIDEO_ID/content" -o cosmos3_inver For V2V, `condition_frame_indexes_vision` selects the clean conditioned latent frame indexes (default `[0, 1]`), and `condition_video_keep` selects whether the API decodes the first or last needed reference frames (`"first"` by default). +- **DROID OpenPI observations:** include a string `prompt`, either + `observation/image` or the three-view DROID camera keys + (`observation/wrist_image_left`, `observation/exterior_image_1_left`, + `observation/exterior_image_2_left`), plus `observation/gripper_position` and + `observation/joint_position`. Optional extra params include `history_length`, + `conditioning_fps`, `action_chunk_size`, `raw_action_dim`, `deterministic_seed`, + and `session_id`. - **Known limitations:** - Guardrails-on requires `cosmos-guardrail` **and** access to the gated `nvidia/Cosmos-1.0-Guardrail` repo (accept license + `HF_TOKEN`); otherwise diff --git a/requirements/common.txt b/requirements/common.txt index fb435218a45..6ca11ac7068 100644 --- a/requirements/common.txt +++ b/requirements/common.txt @@ -22,3 +22,4 @@ prettytable>=3.8.0 aenum==3.1.16 pyzmq>=25.0.0 janus>=1.0.0 +msgpack>=1.0.0 diff --git a/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py b/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py index c51f3ff0805..f77f63d5077 100644 --- a/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py +++ b/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py @@ -5,6 +5,7 @@ import sys import types +from dataclasses import dataclass from types import SimpleNamespace from typing import Any @@ -230,6 +231,8 @@ def test_pipeline_registered_and_exported() -> None: from vllm_omni.diffusion.models.cosmos3.pipeline_cosmos3 import Cosmos3OmniDiffusersPipeline from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin from vllm_omni.diffusion.registry import ( + _DIFFUSION_ACTION_POST_PROCESS_FUNCS, + _DIFFUSION_IR_OP_PRIORITY_FUNCS, _DIFFUSION_MODELS, _DIFFUSION_POST_PROCESS_FUNCS, _DIFFUSION_PRE_PROCESS_FUNCS, @@ -245,6 +248,10 @@ def test_pipeline_registered_and_exported() -> None: ) assert _DIFFUSION_PRE_PROCESS_FUNCS["Cosmos3OmniDiffusersPipeline"] == "get_cosmos3_pre_process_func" assert _DIFFUSION_POST_PROCESS_FUNCS["Cosmos3OmniDiffusersPipeline"] == "get_cosmos3_post_process_func" + assert ( + _DIFFUSION_ACTION_POST_PROCESS_FUNCS["Cosmos3OmniDiffusersPipeline"] == "get_cosmos3_action_post_process_func" + ) + assert _DIFFUSION_IR_OP_PRIORITY_FUNCS["Cosmos3OmniDiffusersPipeline"] == "get_cosmos3_ir_op_priority_func" assert "Cosmos3OmniDiffusersPipeline" in CUSTOM_DIT_ENABLERS assert "Cosmos3OmniDiffusersPipeline" in cosmos3.__all__ @@ -413,6 +420,73 @@ def test_postprocess_handles_image_video_audio_and_validation() -> None: func({"image": video, "video": video}) +def test_action_postprocess_handles_robolab_policy_outputs() -> None: + from vllm_omni.diffusion.models.cosmos3.pipeline_cosmos3 import ( + RoboLabPolicyInputs, + get_cosmos3_action_post_process_func, + make_robolab_action_postprocess_inputs, + ) + + func = get_cosmos3_action_post_process_func(SimpleNamespace()) + inputs = RoboLabPolicyInputs( + prompt="Pick the cube.", + video_tensor=torch.zeros(1, 3, 3, 16, 16), + action_tensor=torch.zeros(2, 2), + action_condition_indexes=[0], + action_start_frame_offset=1, + raw_action_dim=2, + domain_id=7, + fps=15.0, + height=16, + width=16, + image_size=None, + num_frames=3, + num_inference_steps=4, + guidance_scale=3.0, + flow_shift=5.0, + seed=11, + history_length=1, + action_space="joint_pos", + observation={}, + ) + + action = torch.tensor([[[0.0, 0.25], [1.0, 0.75]]]) + custom_output = {"robolab_action_postprocess": make_robolab_action_postprocess_inputs(inputs)} + processed = func(action, custom_output=custom_output) + + assert processed.shape == (1, 2) + assert processed.dtype == torch.zeros((), dtype=torch.float32).numpy().dtype + torch.testing.assert_close(torch.from_numpy(processed), torch.tensor([[1.0, 0.25]])) + assert "robolab_action_postprocess" not in custom_output + + +def test_ir_op_priority_hook_preserves_platform_fields(monkeypatch: pytest.MonkeyPatch) -> None: + from vllm_omni.diffusion.models.cosmos3.pipeline_cosmos3 import get_cosmos3_ir_op_priority_func + + @dataclass + class FakeIrOpPriorityConfig: + rms_norm: list[str] + fused_add_rms_norm: list[str] + custom_op: list[str] + + fake_kernel = types.ModuleType("vllm.config.kernel") + fake_kernel.IrOpPriorityConfig = FakeIrOpPriorityConfig + monkeypatch.setitem(sys.modules, fake_kernel.__name__, fake_kernel) + + func = get_cosmos3_ir_op_priority_func(SimpleNamespace()) + default_priority = FakeIrOpPriorityConfig( + rms_norm=["vllm_c", "native"], + fused_add_rms_norm=["vllm_c", "native"], + custom_op=["platform_kernel", "native"], + ) + + merged = func(default_priority, vllm_config=SimpleNamespace()) + + assert merged.rms_norm == ["native"] + assert merged.fused_add_rms_norm == ["native"] + assert merged.custom_op == ["platform_kernel", "native"] + + def test_prompt_formatting_and_checkpoint_key_remap(make_cosmos3_pipeline) -> None: from vllm_omni.diffusion.models.cosmos3.pipeline_cosmos3 import Cosmos3OmniDiffusersPipeline @@ -762,6 +836,95 @@ def test_forward_i2v_sound_and_action_routes(self, make_cosmos3_pipeline) -> Non ) assert captured["diffuse_calls"][-1]["shared_kwargs"]["action_domain_ids"].tolist() == [7] assert output.custom_output["action"].shape == (1, 2, 2) + assert "action_only_output" not in output.custom_output + + def test_forward_dispatches_robolab_policy_flow( + self, + make_cosmos3_pipeline, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + from vllm_omni.diffusion.models.cosmos3 import pipeline_cosmos3 + + pipeline = make_cosmos3_pipeline() + pipeline.transformer = pipeline.transformer.__class__(latent_channel_size=2, action_gen=True, action_dim=4) + captured = self._install_forward_stubs(pipeline) + video_latents = torch.zeros(1, 2, 1, 1, 1) + velocity_mask = torch.ones(1, 1, 1, 1, 1) + condition_latents = torch.zeros_like(video_latents) + + inputs = pipeline_cosmos3.RoboLabPolicyInputs( + prompt="Pick the cube.", + video_tensor=torch.zeros(1, 3, 3, 16, 16), + action_tensor=torch.zeros(2, 2), + action_condition_indexes=[0], + action_start_frame_offset=1, + raw_action_dim=2, + domain_id=7, + fps=15.0, + height=16, + width=16, + image_size=None, + num_frames=3, + num_inference_steps=4, + guidance_scale=3.0, + flow_shift=5.0, + seed=11, + history_length=1, + action_space="joint_pos", + observation={}, + ) + + def fake_prepare_action_latents(**kwargs): + captured["prepare_action"] = kwargs + action_chunk_size = kwargs["action_chunk_size"] + raw_action_dim = int(kwargs["raw_action_dim"]) + return ( + torch.zeros(1, action_chunk_size, 4), + torch.ones(1, action_chunk_size, 1), + torch.zeros(1, action_chunk_size, 4), + raw_action_dim, + ) + + def fake_prepare_action_video(*args, **kwargs): + captured["prepare_action_video"] = {"args": args, "kwargs": kwargs} + return video_latents, velocity_mask, condition_latents + + monkeypatch.setattr( + pipeline_cosmos3, + "build_robolab_unipc_scheduler", + lambda num_steps, shift, device: StubScheduler(list(range(num_steps, 0, -1)), flow_shift=shift), + ) + pipeline._build_robolab_policy_inputs = lambda sp, prompt_data, request_id=None: inputs + pipeline._prepare_action_latents = fake_prepare_action_latents + pipeline._prepare_latents_action_video = fake_prepare_action_video + pipeline._decode_latents = lambda latents: (_ for _ in ()).throw( + AssertionError("RoboLab should not decode video") + ) + + output = pipeline.forward(SimpleNamespace(prompts=["ignored"], sampling_params=make_sampling_params())) + + assert captured["format"] == { + "prompt": "Pick the cube.", + "negative_prompt": "", + "num_frames": 3, + "frame_rate": 15.0, + "height": 16, + "width": 16, + "is_t2i": False, + } + assert "flow_shifts" not in captured + assert pipeline.scheduler.set_timesteps_calls == [] + assert captured["prepare_action"]["clean_action"] is inputs.action_tensor + assert captured["prepare_action"]["condition_indexes"] == [0] + assert captured["prepare_action_video"]["kwargs"] == {"image_size": None} + assert captured["diffuse_calls"][-1]["shared_kwargs"]["action_domain_ids"].tolist() == [7] + assert captured["diffuse_calls"][-1]["timesteps"].tolist() == [4, 3, 2, 1] + assert output.output == {} + assert output.custom_output["action_only_output"] is True + assert output.custom_output["action"].shape == (1, 2, 2) + assert "actions" not in output.custom_output + assert "robolab_action_postprocess" in output.custom_output + assert "robolab_policy_inputs" not in output.custom_output @pytest.mark.parametrize( ("prompt", "sampling_params", "message"), diff --git a/tests/diffusion/test_diffusion_plugin_hooks.py b/tests/diffusion/test_diffusion_plugin_hooks.py index 72426693917..6804d550d53 100644 --- a/tests/diffusion/test_diffusion_plugin_hooks.py +++ b/tests/diffusion/test_diffusion_plugin_hooks.py @@ -10,11 +10,14 @@ - Worker integration: model runner resolved via platform hook """ -from unittest.mock import patch +from types import SimpleNamespace +from unittest.mock import Mock, patch import pytest from vllm_omni.diffusion.registry import ( + _DIFFUSION_ACTION_POST_PROCESS_FUNCS, + _DIFFUSION_IR_OP_PRIORITY_FUNCS, _DIFFUSION_MODELS, _DIFFUSION_POST_PROCESS_FUNCS, _DIFFUSION_PRE_PROCESS_FUNCS, @@ -59,6 +62,8 @@ def cleanup_registry(self): original_models = _DIFFUSION_MODELS.copy() original_pre = _DIFFUSION_PRE_PROCESS_FUNCS.copy() original_post = _DIFFUSION_POST_PROCESS_FUNCS.copy() + original_action_post = _DIFFUSION_ACTION_POST_PROCESS_FUNCS.copy() + original_ir_op_priority = _DIFFUSION_IR_OP_PRIORITY_FUNCS.copy() yield _DIFFUSION_MODELS.clear() _DIFFUSION_MODELS.update(original_models) @@ -66,6 +71,10 @@ def cleanup_registry(self): _DIFFUSION_PRE_PROCESS_FUNCS.update(original_pre) _DIFFUSION_POST_PROCESS_FUNCS.clear() _DIFFUSION_POST_PROCESS_FUNCS.update(original_post) + _DIFFUSION_ACTION_POST_PROCESS_FUNCS.clear() + _DIFFUSION_ACTION_POST_PROCESS_FUNCS.update(original_action_post) + _DIFFUSION_IR_OP_PRIORITY_FUNCS.clear() + _DIFFUSION_IR_OP_PRIORITY_FUNCS.update(original_ir_op_priority) def test_register_new_model(self): """Test registering a new diffusion model with pre/post process functions.""" @@ -75,6 +84,8 @@ def test_register_new_model(self): class_name="TestPipeline", pre_process_func_name="test_pre_process", post_process_func_name="test_post_process", + action_post_process_func_name="test_action_post_process", + ir_op_priority_func_name="test_ir_op_priority", ) assert "TestPipeline" in _DIFFUSION_MODELS assert _DIFFUSION_MODELS["TestPipeline"] == ( @@ -84,6 +95,8 @@ def test_register_new_model(self): ) assert _DIFFUSION_PRE_PROCESS_FUNCS["TestPipeline"] == "test_pre_process" assert _DIFFUSION_POST_PROCESS_FUNCS["TestPipeline"] == "test_post_process" + assert _DIFFUSION_ACTION_POST_PROCESS_FUNCS["TestPipeline"] == "test_action_post_process" + assert _DIFFUSION_IR_OP_PRIORITY_FUNCS["TestPipeline"] == "test_ir_op_priority" class TestWorkerUsesHook: @@ -108,3 +121,22 @@ def test_model_runner_resolved_via_platform(self, mock_platform, mock_resolve): assert worker.model_runner is mock_runner_instance mock_platform.get_diffusion_model_runner_cls.assert_called_once() mock_resolve.assert_called_once_with("custom.path") + + @patch("vllm_omni.diffusion.worker.diffusion_worker.get_diffusion_ir_op_priority_func") + @patch("vllm_omni.diffusion.worker.diffusion_worker.current_omni_platform") + def test_ir_op_priority_hook_receives_platform_default(self, mock_platform, mock_get_hook): + """Test model IR priority hook merges from the platform default.""" + from vllm_omni.diffusion.worker.diffusion_worker import _resolve_ir_op_priority + + od_config = SimpleNamespace(model_class_name="TestPipeline") + vllm_config = SimpleNamespace() + default_priority = object() + merged_priority = object() + hook = Mock(return_value=merged_priority) + mock_platform.get_default_ir_op_priority.return_value = default_priority + mock_get_hook.return_value = hook + + assert _resolve_ir_op_priority(od_config, vllm_config) is merged_priority + mock_platform.get_default_ir_op_priority.assert_called_once_with(vllm_config) + mock_get_hook.assert_called_once_with(od_config) + hook.assert_called_once_with(default_priority, vllm_config=vllm_config) diff --git a/tests/diffusion/test_diffusion_scheduler.py b/tests/diffusion/test_diffusion_scheduler.py index 76234b7fa42..f5a4f7f7717 100644 --- a/tests/diffusion/test_diffusion_scheduler.py +++ b/tests/diffusion/test_diffusion_scheduler.py @@ -567,6 +567,78 @@ async def test_step_multi_request_reuses_multimodal_slice_logic(self, mocker: Mo torch.tensor([3.0, 4.0]), ) + @pytest.mark.asyncio + async def test_step_empty_dict_output_still_runs_postprocess(self, mocker: MockerFixture) -> None: + engine = DiffusionEngine.__new__(DiffusionEngine) + engine.od_config = SimpleNamespace( + model_class_name="mock_model", + enable_cpu_offload=False, + ) + engine.pre_process_func = None + engine.post_process_func = mocker.Mock(return_value={"video": ["processed"]}) + engine._post_process_accepts_sampling_params = False + engine._check_and_start_background_loop = mocker.AsyncMock() + engine.async_add_req_and_wait_for_response = mocker.AsyncMock( + return_value=DiffusionOutput( + output={}, + custom_output={"actions": torch.tensor([[1.0, 2.0]])}, + ) + ) + + request = OmniDiffusionRequest( + prompts=["prompt"], + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), + request_id="req-action", + ) + + mocker.patch("vllm_omni.diffusion.diffusion_engine.supports_audio_output", return_value=False) + outputs = await engine.step(request) + + engine.post_process_func.assert_called_once_with({}) + assert outputs[0].images == ["processed"] + torch.testing.assert_close(outputs[0].multimodal_output["actions"], torch.tensor([[1.0, 2.0]])) + + @pytest.mark.asyncio + async def test_step_action_only_flag_skips_postprocess(self, mocker: MockerFixture) -> None: + engine = DiffusionEngine.__new__(DiffusionEngine) + engine.od_config = SimpleNamespace( + model_class_name="mock_model", + enable_cpu_offload=False, + ) + engine.pre_process_func = None + engine.post_process_func = mocker.Mock(side_effect=AssertionError("postprocess should be skipped")) + engine.action_post_process_func = mocker.Mock(return_value=torch.tensor([[3.0, 4.0]])) + engine._post_process_accepts_sampling_params = False + engine._action_post_process_accepts_custom_output = True + engine._action_post_process_accepts_sampling_params = False + engine._check_and_start_background_loop = mocker.AsyncMock() + raw_action = torch.tensor([[1.0, 2.0]]) + engine.async_add_req_and_wait_for_response = mocker.AsyncMock( + return_value=DiffusionOutput( + output={}, + custom_output={ + "action": raw_action, + "action_only_output": True, + }, + ) + ) + + request = OmniDiffusionRequest( + prompts=["prompt"], + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), + request_id="req-action", + ) + + mocker.patch("vllm_omni.diffusion.diffusion_engine.supports_audio_output", return_value=False) + outputs = await engine.step(request) + + engine.post_process_func.assert_not_called() + engine.action_post_process_func.assert_called_once() + assert engine.action_post_process_func.call_args.args[0] is raw_action + assert "custom_output" in engine.action_post_process_func.call_args.kwargs + assert outputs[0].images == [] + torch.testing.assert_close(outputs[0].multimodal_output["actions"], torch.tensor([[3.0, 4.0]])) + class TestStepScheduler: def setup_method(self) -> None: diff --git a/tests/entrypoints/openai_api/test_openpi_connection.py b/tests/entrypoints/openai_api/test_openpi_connection.py index ba5a45c3454..83d17a28cd3 100644 --- a/tests/entrypoints/openai_api/test_openpi_connection.py +++ b/tests/entrypoints/openai_api/test_openpi_connection.py @@ -1,9 +1,7 @@ import asyncio -import builtins -import sys -import types from unittest.mock import AsyncMock, MagicMock +import numpy as np import pytest from vllm_omni.entrypoints.openpi import connection as openpi_connection @@ -52,48 +50,56 @@ def _serving_mock(): return serving -def test_pack_reports_clear_error_when_openpi_client_is_missing(monkeypatch): - real_import = builtins.__import__ - - def import_without_openpi_client(name, globals=None, locals=None, fromlist=(), level=0): - if name == "openpi_client": - raise ModuleNotFoundError("No module named 'openpi_client'", name="openpi_client") - return real_import(name, globals, locals, fromlist, level) - - monkeypatch.setattr(builtins, "__import__", import_without_openpi_client) - - with pytest.raises(ImportError) as exc_info: - openpi_connection._pack({"prompt": "pick up the object"}) +def test_pack_and_unpack_round_trip_numpy_values(): + payload = { + "image": np.arange(6, dtype=np.uint8).reshape(2, 3), + "action": np.asarray([[1.0, 2.0]], dtype=np.float32), + "scalar": np.float32(3.5), + "nested": [{"done": np.bool_(True)}], + } - message = str(exc_info.value) - assert "/v1/realtime/robot/openpi" in message - assert "pip install openpi-client" in message + decoded = openpi_connection._unpack(openpi_connection._pack(payload)) + + np.testing.assert_array_equal(decoded["image"], payload["image"]) + np.testing.assert_allclose(decoded["action"], payload["action"]) + assert decoded["image"].dtype == np.uint8 + assert decoded["action"].dtype == np.float32 + assert decoded["scalar"] == np.float32(3.5) + assert decoded["nested"][0]["done"] == np.bool_(True) + + +def test_unpack_accepts_msgpack_numpy_marker_dicts(): + action = np.asarray([[1.0, 2.0]], dtype=np.float32) + payload = { + b"actions": { + b"nd": True, + b"type": action.dtype.str, + b"kind": action.dtype.kind, + b"shape": action.shape, + b"data": action.tobytes(), + } + } + decoded = openpi_connection._unpack_numpy(payload) -def test_pack_and_unpack_delegate_to_openpi_msgpack_numpy(monkeypatch): - calls = [] + np.testing.assert_allclose(decoded[b"actions"], action) + assert decoded[b"actions"].flags.writeable is True + decoded[b"actions"][:, -1] = 0.0 + np.testing.assert_allclose(decoded[b"actions"], np.asarray([[1.0, 0.0]], dtype=np.float32)) - class FakeMsgpackNumpy: - @staticmethod - def packb(obj): - calls.append(("packb", obj)) - return b"packed" - @staticmethod - def unpackb(data): - calls.append(("unpackb", data)) - return {"unpacked": data} +def test_unpack_leaves_user_dict_without_numpy_kind_marker_unchanged(): + payload = { + "metadata": { + "nd": True, + "type": " bool: return bool(getattr(model_cls, "support_audio_output", False)) +def _func_accepts_parameter(func: object | None, parameter_name: str) -> bool: + if func is None: + return False + parameters = inspect.signature(func).parameters + return parameter_name in parameters or any( + parameter.kind == inspect.Parameter.VAR_KEYWORD for parameter in parameters.values() + ) + + def _move_tensor_tree_to_cpu(value: object) -> object: if isinstance(value, torch.Tensor): return value.cpu() if value.device.type != "cpu" else value @@ -149,12 +159,16 @@ def __init__( self.od_config = od_config self.post_process_func = get_diffusion_post_process_func(od_config) + self.action_post_process_func = get_diffusion_action_post_process_func(od_config) self.pre_process_func = get_diffusion_pre_process_func(od_config) # Cache whether the model-specific postprocess accepts request-level # sampling params so step() can support both legacy and extended hooks. - self._post_process_accepts_sampling_params = bool( - self.post_process_func is not None - and "sampling_params" in inspect.signature(self.post_process_func).parameters + self._post_process_accepts_sampling_params = _func_accepts_parameter(self.post_process_func, "sampling_params") + self._action_post_process_accepts_sampling_params = _func_accepts_parameter( + self.action_post_process_func, "sampling_params" + ) + self._action_post_process_accepts_custom_output = _func_accepts_parameter( + self.action_post_process_func, "custom_output" ) executor_class = DiffusionExecutor.get_class(od_config) @@ -261,8 +275,14 @@ async def step(self, request: OmniDiffusionRequest) -> list[OmniRequestOutput]: if self.od_config.enable_cpu_offload: output_data = _move_tensor_tree_to_cpu(output_data) + custom_output = output.custom_output or {} + action_payload = None + action_only_output = bool(custom_output.get("action_only_output")) + postprocess_start_time = time.perf_counter() - if self.post_process_func is not None: + if action_only_output: + outputs = [] + elif self.post_process_func is not None: # Some video pipelines need request-level controls during # postprocess (for example worker-side frame interpolation). if self._post_process_accepts_sampling_params: @@ -272,10 +292,8 @@ async def step(self, request: OmniDiffusionRequest) -> list[OmniRequestOutput]: else: outputs = output_data audio_payload = None - custom_output = output.custom_output or {} model_audio_sample_rate = None model_fps = None - action_payload = None if isinstance(outputs, dict): audio_payload = outputs.get("audio") action_payload = outputs.get("actions") @@ -283,6 +301,19 @@ async def step(self, request: OmniDiffusionRequest) -> list[OmniRequestOutput]: model_audio_sample_rate = outputs.get("audio_sample_rate") model_fps = outputs.get("fps") outputs = outputs.get("video", outputs) + if action_payload is None: + action_payload = custom_output.get("actions") + action_post_process_func = getattr(self, "action_post_process_func", None) + if action_payload is None and action_post_process_func is not None: + raw_action_payload = custom_output.get("action", action_payload) + if raw_action_payload is not None: + action_kwargs: dict[str, Any] = {} + if getattr(self, "_action_post_process_accepts_custom_output", False): + action_kwargs["custom_output"] = custom_output + if getattr(self, "_action_post_process_accepts_sampling_params", False): + action_kwargs["sampling_params"] = request.sampling_params + action_payload = action_post_process_func(raw_action_payload, **action_kwargs) + custom_output["actions"] = action_payload postprocess_time = time.perf_counter() - postprocess_start_time logger.debug("Post-processing completed in %.4f seconds", postprocess_time) diff --git a/vllm_omni/diffusion/models/cosmos3/__init__.py b/vllm_omni/diffusion/models/cosmos3/__init__.py index 6df062b5c0d..43be0c37aa7 100644 --- a/vllm_omni/diffusion/models/cosmos3/__init__.py +++ b/vllm_omni/diffusion/models/cosmos3/__init__.py @@ -3,6 +3,7 @@ from .pipeline_cosmos3 import ( Cosmos3OmniDiffusersPipeline, + get_cosmos3_action_post_process_func, get_cosmos3_post_process_func, get_cosmos3_pre_process_func, ) @@ -10,6 +11,7 @@ __all__ = [ "Cosmos3OmniDiffusersPipeline", + "get_cosmos3_action_post_process_func", "get_cosmos3_post_process_func", "get_cosmos3_pre_process_func", "Cosmos3VFMTransformer", diff --git a/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py b/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py index df9eb835dda..b0b1f7559c1 100644 --- a/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py +++ b/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py @@ -21,6 +21,7 @@ import os import time from collections.abc import Iterable +from dataclasses import fields from typing import Any, ClassVar import numpy as np @@ -65,16 +66,45 @@ resolve_domain_id, ) from .transformer_cosmos3 import Cosmos3VFMTransformer, resolve_sound_gen +from .utils import ( + COSMOS3_DEFAULT_CONDITION_FRAME_INDEXES_VISION, + COSMOS3_VAE_TEMPORAL_COMPRESSION, + ROBOLAB_CONCAT_VIEW_DESCRIPTION, + ROBOLAB_DEFAULT_ACTION_CHUNK_SIZE, + ROBOLAB_DEFAULT_ACTION_SPACE, + ROBOLAB_DEFAULT_CONDITIONING_FPS, + ROBOLAB_DEFAULT_DOMAIN_NAME, + ROBOLAB_DEFAULT_FLOW_SHIFT, + ROBOLAB_DEFAULT_GUIDANCE_SCALE, + ROBOLAB_DEFAULT_IMAGE_HEIGHT, + ROBOLAB_DEFAULT_IMAGE_WIDTH, + ROBOLAB_DEFAULT_NUM_INFERENCE_STEPS, + ROBOLAB_DEFAULT_RAW_ACTION_DIM, + ROBOLAB_DEFAULT_RESOLUTION, + ROBOLAB_MIDTRAIN_RAW_ACTION_DIM, + RoboLabActionPostprocessInputs, + RoboLabPolicyInputs, + build_abs_pose_from_components, + build_robolab_unipc_scheduler, + condition_pixel_frame_count, + convert_midtrain_rotation, + ensure_2d_float_array, + ensure_gripper_array, + extract_robolab_image, + extract_robolab_prompt_image, + lazy_action_transform_pipeline, + make_robolab_action_postprocess_inputs, + next_robolab_seed, + normalize_condition_frame_indexes_vision, + normalize_condition_video_keep, + normalize_robolab_action_space, + pose_abs_to_rel, + postprocess_robolab_action, + resize_rgb_uint8, +) logger = init_logger(__name__) -COSMOS3_DEFAULT_CONDITION_FRAME_INDEXES_VISION = (0, 1) -COSMOS3_DEFAULT_CONDITION_VIDEO_KEEP = "first" -# Mirrors the WAN VAE's temporal compression. Authoritative value is -# ``self.vae.config.scale_factor_temporal`` at runtime; this constant exists so -# off-line / API code that runs before the pipeline is constructed can compute -# pixel-frame budgets without instantiating the VAE. -COSMOS3_VAE_TEMPORAL_COMPRESSION = 4 COSMOS3_DEFAULT_CONDITION_PIXEL_FRAMES = ( max(COSMOS3_DEFAULT_CONDITION_FRAME_INDEXES_VISION) * COSMOS3_VAE_TEMPORAL_COMPRESSION + 1 ) @@ -107,43 +137,6 @@ COSMOS3_DEFAULT_MAX_SEQUENCE_LENGTH = 4096 -def _normalize_condition_frame_indexes_vision(value: Any) -> tuple[int, ...]: - """Normalize Cosmos3 vision-conditioning latent frame indexes.""" - if value is None: - return COSMOS3_DEFAULT_CONDITION_FRAME_INDEXES_VISION - if isinstance(value, str): - value = [item.strip() for item in value.split(",") if item.strip()] - elif isinstance(value, int): - value = [value] - - if not isinstance(value, Iterable): - raise TypeError( - "Cosmos3 condition_frame_indexes_vision must be an int, comma-separated string, " - f"or iterable of ints; got {type(value)!r}." - ) - - indexes = tuple(sorted({int(index) for index in value})) - if not indexes: - raise ValueError("Cosmos3 condition_frame_indexes_vision must contain at least one index.") - if any(index < 0 for index in indexes): - raise ValueError(f"Cosmos3 condition_frame_indexes_vision must be non-negative, got {indexes}.") - return indexes - - -def _condition_pixel_frame_count( - condition_frame_indexes_vision: Iterable[int], - temporal_compression: int = COSMOS3_VAE_TEMPORAL_COMPRESSION, -) -> int: - return max(condition_frame_indexes_vision) * int(temporal_compression) + 1 - - -def _normalize_condition_video_keep(value: Any) -> str: - keep = str(value or COSMOS3_DEFAULT_CONDITION_VIDEO_KEEP).strip().lower() - if keep not in {"first", "last"}: - raise ValueError("Cosmos3 condition_video_keep must be either 'first' or 'last'.") - return keep - - # --------------------------------------------------------------------------- # Post-process function (registered in registry.py) # --------------------------------------------------------------------------- @@ -365,16 +358,16 @@ def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: else: assert raw_video_frames is not None extra = _extra_args(request) - condition_frame_indexes_vision = _normalize_condition_frame_indexes_vision( + condition_frame_indexes_vision = normalize_condition_frame_indexes_vision( extra.get( "condition_frame_indexes_vision", prompt.get("condition_frame_indexes_vision"), ) ) - keep = _normalize_condition_video_keep( + keep = normalize_condition_video_keep( extra.get("condition_video_keep", prompt.get("condition_video_keep")) ) - max_frames = _condition_pixel_frame_count(condition_frame_indexes_vision) + max_frames = condition_pixel_frame_count(condition_frame_indexes_vision) prompt["additional_information"]["preprocessed_video"] = _preprocess_condition_video( raw_video_frames, int(target_h), @@ -491,6 +484,36 @@ def post_process_func( return post_process_func +def get_cosmos3_action_post_process_func(od_config: OmniDiffusionConfig): + del od_config + + def action_post_process_func(action: Any, custom_output: dict[str, Any] | None = None, sampling_params=None): + del sampling_params + inputs = custom_output.get("robolab_action_postprocess") if isinstance(custom_output, dict) else None + if isinstance(inputs, RoboLabActionPostprocessInputs): + processed_action = postprocess_robolab_action(action, inputs) + custom_output.pop("robolab_action_postprocess", None) + return processed_action + return action + + return action_post_process_func + + +def get_cosmos3_ir_op_priority_func(od_config: OmniDiffusionConfig): + del od_config + + def ir_op_priority_func(ir_op_priority, vllm_config=None): + del vllm_config + from vllm.config.kernel import IrOpPriorityConfig + + priority_kwargs = {field.name: list(getattr(ir_op_priority, field.name)) for field in fields(ir_op_priority)} + priority_kwargs["rms_norm"] = ["native"] + priority_kwargs["fused_add_rms_norm"] = ["native"] + return IrOpPriorityConfig(**priority_kwargs) + + return ir_op_priority_func + + # --------------------------------------------------------------------------- # Pipeline # --------------------------------------------------------------------------- @@ -549,11 +572,11 @@ def reference_video_decode_spec( return ReferenceVideoDecodeSpec(max_frames=max_frames, keep="first") return ReferenceVideoDecodeSpec(max_frames=None, keep="first") - condition_indexes = _normalize_condition_frame_indexes_vision(extra_args.get("condition_frame_indexes_vision")) - max_frames = _condition_pixel_frame_count(condition_indexes) + condition_indexes = normalize_condition_frame_indexes_vision(extra_args.get("condition_frame_indexes_vision")) + max_frames = condition_pixel_frame_count(condition_indexes) if num_frames is not None: max_frames = min(max_frames, int(num_frames)) - keep = _normalize_condition_video_keep(extra_args.get("condition_video_keep")) + keep = normalize_condition_video_keep(extra_args.get("condition_video_keep")) return ReferenceVideoDecodeSpec(max_frames=max_frames, keep=keep) def __init__( @@ -654,6 +677,7 @@ def __init__( self._guidance_scale = None self._num_timesteps = None + self._robolab_transform = None # Set True by ``enable_cache_for_cosmos3`` when cache-dit is enabled on # this pipeline. Tells the sequential-CFG loop to keep paired @@ -875,6 +899,306 @@ def _get_sp_param(sp: OmniDiffusionSamplingParams, key: str, default: Any = None return val return default + def _get_robolab_transform(self): + if self._robolab_transform is None: + action_dim = int(getattr(self.transformer, "action_dim", 64)) + self._robolab_transform = lazy_action_transform_pipeline(action_dim) + return self._robolab_transform + + def _build_robolab_policy_inputs( + self, + sp: OmniDiffusionSamplingParams, + prompt_data: Any | None = None, + request_id: str | None = None, + ) -> RoboLabPolicyInputs | None: + extra = sp.extra_args if isinstance(sp.extra_args, dict) else {} + obs = extra.get("robot_obs") + if obs is None: + obs = extra.get("observation") + if obs is None: + return None + if not isinstance(obs, dict): + raise TypeError(f"Cosmos3 RoboLab observation must be a dict, got {type(obs)!r}.") + + prompt = obs.get("prompt") + if not isinstance(prompt, str): + raise ValueError("RoboLab observation must contain string key 'prompt'.") + + def extra_param(key: str, default: Any) -> Any: + value = extra.get(key) + return default if value is None else value + + def extra_param_alias(primary_key: str, alias_key: str, default: Any) -> Any: + value = extra.get(primary_key) + if value is not None: + return value + value = extra.get(alias_key) + return default if value is None else value + + action_space = normalize_robolab_action_space(extra_param("action_space", ROBOLAB_DEFAULT_ACTION_SPACE)) + action_chunk_size = int(extra_param("action_chunk_size", ROBOLAB_DEFAULT_ACTION_CHUNK_SIZE)) + raw_action_dim_default = ( + ROBOLAB_DEFAULT_RAW_ACTION_DIM if action_space == "joint_pos" else ROBOLAB_MIDTRAIN_RAW_ACTION_DIM + ) + raw_action_dim = int(extra_param("raw_action_dim", raw_action_dim_default)) + image_h = int(extra_param("image_height", ROBOLAB_DEFAULT_IMAGE_HEIGHT)) + image_w = int(extra_param("image_width", ROBOLAB_DEFAULT_IMAGE_WIDTH)) + history_length = int(extra_param("history_length", 1)) + use_state = self._truthy(extra_param("use_state", True)) + resolution = str(extra_param("resolution", ROBOLAB_DEFAULT_RESOLUTION)) + fps = float(extra_param("conditioning_fps", ROBOLAB_DEFAULT_CONDITIONING_FPS)) + domain_name = str(extra_param("domain_name", ROBOLAB_DEFAULT_DOMAIN_NAME)) + domain_id = resolve_domain_id(domain_name=domain_name, require_explicit=True) + + if use_state and history_length < 1: + raise ValueError("RoboLab history_length must be >= 1 when use_state is true.") + if action_chunk_size <= 0: + raise ValueError(f"RoboLab action_chunk_size must be positive, got {action_chunk_size}.") + if raw_action_dim <= 0: + raise ValueError(f"RoboLab raw_action_dim must be positive, got {raw_action_dim}.") + + try: + image = extract_robolab_image(obs) + except ValueError as exc: + image = extract_robolab_prompt_image(prompt_data) + if image is None: + raise exc + if image.shape[:2] != (image_h, image_w): + image = resize_rgb_uint8(image, (image_h, image_w)) + + t_frames = action_chunk_size + 1 + video = torch.zeros((3, t_frames, image_h, image_w), dtype=torch.uint8) + video[:, 0] = torch.from_numpy(image.copy()).permute(2, 0, 1) + + use_state_rows = 1 if use_state else 0 + action = torch.zeros((action_chunk_size + use_state_rows, raw_action_dim), dtype=torch.float32) + history_action = None + num_history_rows = history_length - use_state_rows + gripper_position = 1.0 - ensure_gripper_array(obs["observation/gripper_position"]) + + if action_space == "joint_pos": + joint_position = ensure_2d_float_array(obs["observation/joint_position"], "observation/joint_position", 7) + if use_state: + action[0] = torch.from_numpy(np.concatenate((joint_position[-1], gripper_position[-1]))) + if num_history_rows > 0: + if len(joint_position) < num_history_rows + 1: + raise ValueError("Not enough joint_position rows for requested history_length.") + history_np = np.concatenate( + (joint_position[-num_history_rows - 1 : -1], gripper_position[-num_history_rows - 1 : -1]), + axis=-1, + ) + history_action = torch.from_numpy(history_np).float() + else: + eef_pos = ensure_2d_float_array(obs["observation/eef_pos"], "observation/eef_pos", 3) + eef_quat = ensure_2d_float_array(obs["observation/eef_quat"], "observation/eef_quat", 4) + if use_state: + rot6d = convert_midtrain_rotation(eef_quat[-1], "quat_xyzw", "rot6d") + action[0] = torch.from_numpy(np.concatenate((eef_pos[-1], rot6d, gripper_position[-1]))) + if num_history_rows > 0: + if len(eef_pos) < num_history_rows + 1 or len(eef_quat) < num_history_rows + 1: + raise ValueError("Not enough eef_pos/eef_quat rows for requested history_length.") + poses_abs = build_abs_pose_from_components(eef_pos, eef_quat, "quat_xyzw") + poses_rel = pose_abs_to_rel(poses_abs, rotation_format="rot6d", pose_convention="backward_framewise") + history_np = np.concatenate( + [poses_rel[-num_history_rows:], gripper_position[-num_history_rows:]], + axis=-1, + ) + history_action = torch.from_numpy(history_np).float() + + sample: dict[str, Any] = { + "ai_caption": prompt, + "video": video, + "action": action, + # Cosmos Framework consumes this as an integer conditioning bucket. + "conditioning_fps": torch.tensor(fps, dtype=torch.long), + "mode": ACTION_MODE_POLICY, + "domain_id": torch.tensor(domain_id, dtype=torch.long), + "viewpoint": "concat_view", + "additional_view_description": ROBOLAB_CONCAT_VIEW_DESCRIPTION, + } + if history_action is not None: + sample["history_action"] = history_action + + sample = self._get_robolab_transform()(sample, resolution) + sequence_plan = sample["sequence_plan"] + video_tensor = sample["video"].float() / 127.5 - 1.0 + raw_action_dim_tensor = sample.get("raw_action_dim") + if isinstance(raw_action_dim_tensor, torch.Tensor): + transformed_raw_action_dim = int(raw_action_dim_tensor.item()) + else: + transformed_raw_action_dim = raw_action_dim + + return RoboLabPolicyInputs( + prompt=sample["ai_caption"], + video_tensor=video_tensor.unsqueeze(0), + action_tensor=sample["action"].float(), + action_condition_indexes=list(getattr(sequence_plan, "condition_frame_indexes_action", []) or []), + action_start_frame_offset=int(getattr(sequence_plan, "action_start_frame_offset", 1)), + raw_action_dim=transformed_raw_action_dim, + domain_id=domain_id, + fps=fps, + height=int(video_tensor.shape[-2]), + width=int(video_tensor.shape[-1]), + image_size=sample.get("image_size"), + num_frames=int(video_tensor.shape[1]), + num_inference_steps=int( + extra_param_alias("num_inference_steps", "num_steps", ROBOLAB_DEFAULT_NUM_INFERENCE_STEPS) + ), + guidance_scale=float(extra_param_alias("guidance_scale", "guidance", ROBOLAB_DEFAULT_GUIDANCE_SCALE)), + flow_shift=float(extra_param_alias("flow_shift", "shift", ROBOLAB_DEFAULT_FLOW_SHIFT)), + seed=next_robolab_seed(extra, obs, request_id), + history_length=history_length, + action_space=action_space, + observation=obs, + ) + + @staticmethod + def _build_action_condition_mask_from_indexes( + indexes: list[int], + action_length: int, + *, + device: torch.device, + dtype: torch.dtype, + ) -> torch.Tensor: + mask = torch.zeros(1, action_length, 1, device=device, dtype=dtype) + for idx in indexes: + if idx < 0 or idx >= action_length: + raise ValueError(f"Action condition index {idx} is out of range for action length {action_length}.") + mask[:, idx, :] = 1.0 + return mask + + def _forward_robolab_policy( + self, + sp: OmniDiffusionSamplingParams, + inputs: RoboLabPolicyInputs, + pipeline_start: float, + ) -> DiffusionOutput: + if not getattr(self.transformer, "action_gen", False): + raise ValueError( + "Cosmos3 RoboLab policy serving was requested, but the transformer " + "was initialized without action modules. Check that the checkpoint " + "config enables action_gen and includes action weights." + ) + + action_mode = ACTION_MODE_POLICY + height = inputs.height + width = inputs.width + num_frames = inputs.num_frames + action_chunk_size = int(inputs.action_tensor.shape[0]) + num_inference_steps = inputs.num_inference_steps + guidance_scale = float(inputs.guidance_scale) + flow_shift_target = float(inputs.flow_shift) + domain_id = int(inputs.domain_id) + frame_rate = self._get_sp_param(sp, "resolved_frame_rate") or self._get_sp_param(sp, "frame_rate") or inputs.fps + max_sequence_length = ( + self._get_sp_param(sp, "max_sequence_length", COSMOS3_DEFAULT_MAX_SEQUENCE_LENGTH) + or COSMOS3_DEFAULT_MAX_SEQUENCE_LENGTH + ) + use_system_prompt = bool(self._get_sp_param(sp, "use_system_prompt", False)) + + self._guidance_scale = guidance_scale + self._num_timesteps = num_inference_steps + + generator = sp.generator + if generator is None: + generator = torch.Generator(device=self.device).manual_seed(int(inputs.seed)) + + cond_ids, cond_mask, uncond_ids, uncond_mask = self._format_and_tokenize_prompts( + inputs.prompt, + "", + num_frames, + frame_rate, + height, + width, + max_sequence_length, + sp, + use_system_prompt, + is_t2i=False, + ) + + action_video_tensor = inputs.video_tensor + if action_video_tensor.ndim == 4: + action_video_tensor = action_video_tensor.unsqueeze(0) + if action_video_tensor.ndim != 5: + raise ValueError( + "Cosmos3 RoboLab action video tensor must have shape [1, 3, T, H, W] " + f"or [3, T, H, W], got {tuple(action_video_tensor.shape)}." + ) + if action_video_tensor.shape[2] < num_frames: + pad = action_video_tensor[:, :, -1:].repeat(1, 1, num_frames - action_video_tensor.shape[2], 1, 1) + action_video_tensor = torch.cat([action_video_tensor, pad], dim=2) + elif action_video_tensor.shape[2] > num_frames: + action_video_tensor = action_video_tensor[:, :, :num_frames] + + action_latents, action_velocity_mask, action_condition_latents, raw_action_dim = self._prepare_action_latents( + mode=action_mode, + action_chunk_size=action_chunk_size, + raw_action_dim=int(inputs.raw_action_dim), + generator=generator, + sp=sp, + clean_action=inputs.action_tensor, + condition_indexes=inputs.action_condition_indexes, + ) + action_offset = int(inputs.action_start_frame_offset) + + latents, velocity_mask, condition_latents = self._prepare_latents_action_video( + action_video_tensor, + action_mode, + height, + width, + num_frames, + generator, + image_size=inputs.image_size, + ) + image_latent = condition_latents[:, :, 0:1] + + video_shape = (latents.shape[2], latents.shape[3], latents.shape[4]) + shared_kwargs = dict( + video_shape=video_shape, + fps=frame_rate, + noisy_frame_mask=velocity_mask, + action_domain_ids=torch.tensor([domain_id], dtype=torch.long, device=self.device), + action_noisy_mask=action_velocity_mask, + action_start_frame_offset=action_offset, + action_fps=float(self._get_sp_param(sp, "action_fps", frame_rate) or frame_rate), + ) + + scheduler = build_robolab_unipc_scheduler(num_inference_steps, flow_shift_target, self.device) + _, action_latents = self.diffuse( + latents=latents, + timesteps=scheduler.timesteps, + cond_ids=cond_ids, + cond_mask=cond_mask, + uncond_ids=uncond_ids, + uncond_mask=uncond_mask, + guidance_scale=guidance_scale, + shared_kwargs=shared_kwargs, + action_latents=action_latents, + action_velocity_mask=action_velocity_mask, + action_condition_latents=action_condition_latents, + sound_latents=None, + velocity_mask=velocity_mask, + image_latent=image_latent, + condition_latents=condition_latents, + guidance_interval=None, + raw_action_dim=raw_action_dim, + scheduler=scheduler, + ) + + if _is_rank_zero(): + logger.info("Total pipeline time: %.2fs", time.time() - pipeline_start) + + action = action_latents[:, :, :raw_action_dim].detach().cpu() + custom_action_output: dict[str, Any] = { + "action": action, + "raw_action_dim": raw_action_dim, + "action_mode": action_mode, + "domain_id": domain_id, + "action_only_output": True, + "robolab_action_postprocess": make_robolab_action_postprocess_inputs(inputs), + } + return DiffusionOutput(output={}, custom_output=custom_action_output) + @staticmethod def _truthy(value) -> bool: if isinstance(value, str): @@ -1313,7 +1637,28 @@ def _encode_conditioning_video( return latent.to(self.dtype) - def _encode_video_tensor(self, video_tensor: torch.Tensor) -> torch.Tensor: + def _latent_hw_from_image_size(self, image_size: Any | None) -> tuple[int, int] | None: + if image_size is None: + return None + if isinstance(image_size, torch.Tensor): + frame_size = image_size.detach().cpu().flatten() + else: + frame_size = torch.as_tensor(image_size).flatten() + if frame_size.numel() < 4: + return None + orig_h = int(frame_size[2].item()) + orig_w = int(frame_size[3].item()) + spatial_factor = int(self.vae_scale_factor_spatial) + return max(orig_h // spatial_factor, 1), max(orig_w // spatial_factor, 1) + + def _crop_latent_to_image_size(self, latent: torch.Tensor, image_size: Any | None) -> torch.Tensor: + latent_hw = self._latent_hw_from_image_size(image_size) + if latent_hw is None: + return latent + h_latent, w_latent = latent_hw + return latent[:, :, :, :h_latent, :w_latent].contiguous() + + def _encode_video_tensor(self, video_tensor: torch.Tensor, image_size: Any | None = None) -> torch.Tensor: """VAE-encode a preprocessed pixel video [1, 3, T, H, W].""" if video_tensor.ndim == 4: video_tensor = video_tensor.unsqueeze(0) @@ -1335,6 +1680,7 @@ def _encode_video_tensor(self, video_tensor: torch.Tensor) -> torch.Tensor: scaling_factor = getattr(self.vae.config, "scaling_factor", 1.0) latent = latent * scaling_factor + latent = self._crop_latent_to_image_size(latent, image_size) return latent.to(self.dtype) def _prepare_latents_i2v( @@ -1395,7 +1741,7 @@ def _prepare_latents_v2v( T_lat = (num_frames - 1) // self.vae_scale_factor_temporal + 1 H_lat = video_tensor.shape[-2] // self.vae_scale_factor_spatial W_lat = video_tensor.shape[-1] // self.vae_scale_factor_spatial - indexes = _normalize_condition_frame_indexes_vision(condition_frame_indexes_vision) + indexes = normalize_condition_frame_indexes_vision(condition_frame_indexes_vision) out_of_range = [index for index in indexes if index >= T_lat] if out_of_range: raise ValueError( @@ -1409,7 +1755,7 @@ def _prepare_latents_v2v( device=self.device, dtype=self.dtype, ) - condition_pixel_frames = _condition_pixel_frame_count(indexes, self.vae_scale_factor_temporal) + condition_pixel_frames = condition_pixel_frame_count(indexes, self.vae_scale_factor_temporal) condition_video = video_tensor[:, :, :condition_pixel_frames] if condition_video.shape[2] < condition_pixel_frames: pad = condition_video[:, :, -1:].repeat(1, 1, condition_pixel_frames - condition_video.shape[2], 1, 1) @@ -1445,13 +1791,18 @@ def _prepare_latents_action_video( width: int, num_frames: int, generator: torch.Generator, + image_size: Any | None = None, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Prepare video latents for action modes with mode-specific conditioning.""" del height, width C = self.transformer.latent_channel_size T_lat = (num_frames - 1) // self.vae_scale_factor_temporal + 1 - H_lat = video_tensor.shape[-2] // self.vae_scale_factor_spatial - W_lat = video_tensor.shape[-1] // self.vae_scale_factor_spatial + latent_hw = self._latent_hw_from_image_size(image_size) + if latent_hw is None: + H_lat = video_tensor.shape[-2] // self.vae_scale_factor_spatial + W_lat = video_tensor.shape[-1] // self.vae_scale_factor_spatial + else: + H_lat, W_lat = latent_hw noise = randn_tensor( (1, C, T_lat, H_lat, W_lat), @@ -1459,7 +1810,7 @@ def _prepare_latents_action_video( device=self.device, dtype=self.dtype, ) - cond_latent = self._encode_video_tensor(video_tensor) + cond_latent = self._encode_video_tensor(video_tensor, image_size=image_size) if cond_latent.shape[2:] != noise.shape[2:]: raise ValueError( "Cosmos3 action video latent shape mismatch: " @@ -1484,9 +1835,25 @@ def _prepare_action_latents( raw_action_dim: int | None, generator: torch.Generator, sp, + clean_action: torch.Tensor | None = None, + condition_indexes: list[int] | None = None, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]: action_dim = int(getattr(self.transformer, "action_dim", 64)) - if mode == ACTION_MODE_FORWARD_DYNAMICS: + if clean_action is not None: + action = clean_action.detach().to(dtype=torch.float32) + if action.ndim == 3 and action.shape[0] == 1: + action = action.squeeze(0) + if action.ndim != 2: + raise ValueError(f"Cosmos3 clean action must have shape [T, D], got {tuple(action.shape)}.") + if action.shape[0] < action_chunk_size: + pad = action[-1:].repeat(action_chunk_size - action.shape[0], 1) + action = torch.cat([action, pad], dim=0) + elif action.shape[0] > action_chunk_size: + action = action[:action_chunk_size] + if raw_action_dim is None: + raw_action_dim = int(action.shape[-1]) + clean_action = pad_action_to_dim(action, action_dim) + elif mode == ACTION_MODE_FORWARD_DYNAMICS: action = load_action_tensor(self._get_sp_param(sp, "action", None)) if action.shape[0] < action_chunk_size: pad = action[-1:].repeat(action_chunk_size - action.shape[0], 1) @@ -1508,12 +1875,20 @@ def _prepare_action_latents( raise ValueError(f"Cosmos3 raw_action_dim must be in [1, {action_dim}], got {raw_action_dim}.") clean_action = clean_action.to(device=self.device, dtype=self.dtype).unsqueeze(0) - condition_mask = build_action_condition_mask( - mode, - action_chunk_size, - device=self.device, - dtype=self.dtype, - ) + if condition_indexes is None: + condition_mask = build_action_condition_mask( + mode, + action_chunk_size, + device=self.device, + dtype=self.dtype, + ) + else: + condition_mask = self._build_action_condition_mask_from_indexes( + condition_indexes, + action_chunk_size, + device=self.device, + dtype=self.dtype, + ) noise = randn_tensor( (1, action_chunk_size, action_dim), generator=generator, @@ -1548,6 +1923,7 @@ def diffuse( condition_latents: torch.Tensor | None = None, guidance_interval: tuple[float, float] | None = None, raw_action_dim: int | None = None, + scheduler: Any | None = None, ) -> torch.Tensor | tuple[torch.Tensor, ...]: """Denoising loop with 3-mode CFG support (parallel, sequential, none). @@ -1578,6 +1954,7 @@ def diffuse( """ do_cfg = guidance_scale > 1.0 cfg_parallel = self._cfg_parallel_active() and do_cfg + step_scheduler = scheduler if scheduler is not None else self.scheduler self.transformer.reset_cache() def _cfg_active_at(t: torch.Tensor) -> bool: @@ -1655,11 +2032,11 @@ def _step( if raw_action_dim is not None and 0 < raw_action_dim < action_pred.shape[-1]: action_pred[..., raw_action_dim:] = 0 if action_latents is None and sound_latents is None: - latents = self.scheduler.step(video_pred, t, latents, return_dict=False)[0] + latents = step_scheduler.step(video_pred, t, latents, return_dict=False)[0] else: packed_noise, shapes, numels = _pack_joint(video_pred, action_pred, sound_pred) packed_latents, _, _ = _pack_joint(latents, action_latents, sound_latents) - packed_next = self.scheduler.step(packed_noise, t, packed_latents, return_dict=False)[0] + packed_next = step_scheduler.step(packed_noise, t, packed_latents, return_dict=False)[0] unpacked = _unpack_joint(packed_next, shapes, numels) latents = unpacked[0] idx = 1 @@ -1811,7 +2188,12 @@ def forward( if len(req.prompts) > 1: raise ValueError("Cosmos3OmniDiffusersPipeline currently supports a single prompt per request.") + sp = req.sampling_params prompt_data = req.prompts[0] + robolab_inputs = self._build_robolab_policy_inputs(sp, prompt_data, getattr(req, "request_id", None)) + if robolab_inputs is not None: + return self._forward_robolab_policy(sp, robolab_inputs, pipeline_start) + if isinstance(prompt_data, str): prompt = prompt_data negative_prompt = None @@ -1824,7 +2206,6 @@ def forward( image_tensor = additional_info.get("preprocessed_image") video_tensor = additional_info.get("preprocessed_video") - sp = req.sampling_params is_t2i = self._is_t2i_request(req) sound_enabled = self._is_sound_request(prompt_data, sp) action_mode = self._get_action_mode(prompt_data, sp) @@ -1948,8 +2329,7 @@ def forward( self._num_timesteps = num_inference_steps # Always resolve to a concrete target shift for this request, then - # update the scheduler. This is what guarantees mode-to-mode - # transitions restore the right schedule (no T2I to T2V leak). + # update the shared Diffusers scheduler. self._set_flow_shift(flow_shift_target) generator = sp.generator @@ -2010,12 +2390,16 @@ def forward( raw_action_dim_param = self._get_sp_param(sp, "raw_action_dim", None) raw_action_dim = int(raw_action_dim_param) if raw_action_dim_param is not None else None + clean_action = None + action_condition_indexes = None action_prepared = self._prepare_action_latents( mode=action_mode, action_chunk_size=action_chunk_size, raw_action_dim=raw_action_dim, generator=generator, sp=sp, + clean_action=clean_action, + condition_indexes=action_condition_indexes, ) action_latents, action_velocity_mask, action_condition_latents, raw_action_dim = action_prepared action_offset = action_start_frame_offset(action_mode, action_chunk_size, num_frames) @@ -2031,7 +2415,7 @@ def forward( ) image_latent = condition_latents[:, :, 0:1] elif is_v2v: - condition_frame_indexes_vision = _normalize_condition_frame_indexes_vision( + condition_frame_indexes_vision = normalize_condition_frame_indexes_vision( self._get_sp_param( sp, "condition_frame_indexes_vision", @@ -2088,9 +2472,10 @@ def forward( def _run_diffusion(start_latents): self.scheduler.set_timesteps(num_inference_steps, device=self.device) + scheduler = self.scheduler return self.diffuse( latents=start_latents, - timesteps=self.scheduler.timesteps, + timesteps=scheduler.timesteps, cond_ids=cond_ids, cond_mask=cond_mask, uncond_ids=uncond_ids, @@ -2106,6 +2491,7 @@ def _run_diffusion(start_latents): condition_latents=condition_latents, guidance_interval=guidance_interval, raw_action_dim=raw_action_dim, + scheduler=scheduler, ) if is_t2i and batch_size > 1: @@ -2153,14 +2539,12 @@ def _run_diffusion(start_latents): if action_latents is None or raw_action_dim is None or domain_id is None: raise ValueError("Cosmos3 action generation finished without action latents.") action = action_latents[:, :, :raw_action_dim].detach().cpu() - return DiffusionOutput( - output={"video": video}, - custom_output={ - "action": action, - "raw_action_dim": raw_action_dim, - "action_mode": action_mode, - "domain_id": domain_id, - }, - ) + custom_action_output: dict[str, Any] = { + "action": action, + "raw_action_dim": raw_action_dim, + "action_mode": action_mode, + "domain_id": domain_id, + } + return DiffusionOutput(output={"video": video}, custom_output=custom_action_output) return DiffusionOutput(output={"image": video} if is_t2i else {"video": video}) diff --git a/vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py b/vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py index 515e2b258cc..163158a39dc 100644 --- a/vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py +++ b/vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py @@ -33,12 +33,22 @@ from vllm_omni.diffusion.data import OmniDiffusionConfig from vllm_omni.diffusion.distributed.sp_plan import SequenceParallelInput, SequenceParallelOutput from vllm_omni.diffusion.forward_context import get_forward_context, is_forward_context_available -from vllm_omni.diffusion.layers.norm import RMSNorm +from vllm_omni.diffusion.layers.norm import RMSNorm as _VllmRMSNorm from vllm_omni.platforms import current_omni_platform logger = init_logger(__name__) +class RMSNorm(_VllmRMSNorm): + """Cosmos3-local RMSNorm that uses the FP32 native implementation.""" + + def forward_cuda(self, x: torch.Tensor) -> torch.Tensor: + return self.forward_native(x) + + def forward_hip(self, x: torch.Tensor) -> torch.Tensor: + return self.forward_native(x) + + def _get_ulysses_state() -> tuple[int, int, dist.ProcessGroup | None]: """Return (ulysses_size, ulysses_rank, ulysses_pg) from vllm-omni parallel state. @@ -539,8 +549,8 @@ def forward( v = self.to_v(hidden_states).view(B, S, self.num_kv_heads_local, self.head_dim) # Per-head QK norm - q = F.rms_norm(q, (q.shape[-1],), self.norm_q.weight, self.norm_q.variance_epsilon) - k = F.rms_norm(k, (k.shape[-1],), self.norm_k.weight, self.norm_k.variance_epsilon) + q = F.rms_norm(q, (self.head_dim,), self.norm_q.weight, eps=self.norm_q.variance_epsilon) + k = F.rms_norm(k, (self.head_dim,), self.norm_k.weight, eps=self.norm_k.variance_epsilon) # Qwen3-style RoPE q, k = _apply_rotary_pos_emb(q, k, freqs_cos, freqs_sin) @@ -704,8 +714,8 @@ def forward( v = self.to_v(hidden_states).view(B, S_gen, self.num_kv_heads_local, self.head_dim) # Per-head QK norm - q = F.rms_norm(q, (q.shape[-1],), self.norm_q.weight, self.norm_q.variance_epsilon) - k = F.rms_norm(k, (k.shape[-1],), self.norm_k.weight, self.norm_k.variance_epsilon) + q = F.rms_norm(q, (self.head_dim,), self.norm_q.weight, eps=self.norm_q.variance_epsilon) + k = F.rms_norm(k, (self.head_dim,), self.norm_k.weight, eps=self.norm_k.variance_epsilon) # Qwen3-style RoPE q, k = _apply_rotary_pos_emb(q, k, freqs_cos, freqs_sin) @@ -1446,7 +1456,8 @@ def forward( "Cosmos3 action_noisy_mask must have shape [B, T_action, 1], " f"got {tuple(action_noisy_mask.shape)}." ) - hidden_action = hidden_action + time_embed.unsqueeze(1) * action_noisy_mask.to(hidden_action.dtype) + action_noisy_mask = action_noisy_mask.to(dtype=hidden_action.dtype, device=hidden_action.device) + hidden_action = hidden_action + time_embed.unsqueeze(1) * action_noisy_mask if hidden_sound is not None: hidden_sound = hidden_sound + time_embed.unsqueeze(1) diff --git a/vllm_omni/diffusion/models/cosmos3/utils.py b/vllm_omni/diffusion/models/cosmos3/utils.py new file mode 100644 index 00000000000..bfa3fd4db18 --- /dev/null +++ b/vllm_omni/diffusion/models/cosmos3/utils.py @@ -0,0 +1,371 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from __future__ import annotations + +import zlib +from collections.abc import Iterable +from dataclasses import dataclass +from importlib import import_module +from typing import Any + +import numpy as np +import PIL.Image +import torch +import torch.nn.functional as F +from vllm.logger import init_logger + +from vllm_omni.diffusion.models.progress_bar import _is_rank_zero + +logger = init_logger(__name__) + +COSMOS3_DEFAULT_CONDITION_FRAME_INDEXES_VISION = (0, 1) +COSMOS3_DEFAULT_CONDITION_VIDEO_KEEP = "first" +# Mirrors the WAN VAE's temporal compression. Authoritative value is +# ``self.vae.config.scale_factor_temporal`` at runtime; this constant exists so +# off-line / API code that runs before the pipeline is constructed can compute +# pixel-frame budgets without instantiating the VAE. +COSMOS3_VAE_TEMPORAL_COMPRESSION = 4 + +ROBOLAB_DEFAULT_CONDITIONING_FPS = 15.0 +ROBOLAB_DEFAULT_ACTION_CHUNK_SIZE = 32 +ROBOLAB_DEFAULT_IMAGE_HEIGHT = 540 +ROBOLAB_DEFAULT_IMAGE_WIDTH = 640 +ROBOLAB_DEFAULT_RAW_ACTION_DIM = 8 +ROBOLAB_DEFAULT_DOMAIN_NAME = "droid_lerobot" +ROBOLAB_DEFAULT_RESOLUTION = "480" +ROBOLAB_DEFAULT_GUIDANCE_SCALE = 3.0 +ROBOLAB_DEFAULT_NUM_INFERENCE_STEPS = 4 +ROBOLAB_DEFAULT_FLOW_SHIFT = 5.0 +ROBOLAB_DEFAULT_SEED = 0 +ROBOLAB_DEFAULT_ACTION_SPACE = "joint_pos" +ROBOLAB_MIDTRAIN_RAW_ACTION_DIM = 10 +ROBOLAB_MIDTRAIN_POSE_ACTION_DIM = 9 # xyz position + rot6d orientation +ROBOLAB_CONCAT_VIEW_DESCRIPTION = ( + "The top row is from the wrist-mounted camera. " + "The bottom row contains two horizontally concatenated third-person perspective views of the scene from opposite " + "sides, with the robot visible." +) + + +@dataclass(frozen=True) +class RoboLabPolicyInputs: + prompt: str + video_tensor: torch.Tensor + action_tensor: torch.Tensor + action_condition_indexes: list[int] + action_start_frame_offset: int + raw_action_dim: int + domain_id: int + fps: float + height: int + width: int + image_size: Any + num_frames: int + num_inference_steps: int + guidance_scale: float + flow_shift: float + seed: int + history_length: int + action_space: str + observation: dict[str, Any] + + +@dataclass(frozen=True) +class RoboLabActionPostprocessInputs: + history_length: int + action_space: str + eef_pos: np.ndarray | None = None + eef_quat: np.ndarray | None = None + + +def make_robolab_action_postprocess_inputs(inputs: RoboLabPolicyInputs) -> RoboLabActionPostprocessInputs: + if inputs.action_space != "midtrain": + return RoboLabActionPostprocessInputs( + history_length=inputs.history_length, + action_space=inputs.action_space, + ) + + obs = inputs.observation + return RoboLabActionPostprocessInputs( + history_length=inputs.history_length, + action_space=inputs.action_space, + eef_pos=ensure_2d_float_array(obs["observation/eef_pos"], "observation/eef_pos", 3), + eef_quat=ensure_2d_float_array(obs["observation/eef_quat"], "observation/eef_quat", 4), + ) + + +def normalize_condition_frame_indexes_vision(value: Any) -> tuple[int, ...]: + """Normalize Cosmos3 vision-conditioning latent frame indexes.""" + if value is None: + return COSMOS3_DEFAULT_CONDITION_FRAME_INDEXES_VISION + if isinstance(value, str): + value = [item.strip() for item in value.split(",") if item.strip()] + elif isinstance(value, int): + value = [value] + + if not isinstance(value, Iterable): + raise TypeError( + "Cosmos3 condition_frame_indexes_vision must be an int, comma-separated string, " + f"or iterable of ints; got {type(value)!r}." + ) + + indexes = tuple(sorted({int(index) for index in value})) + if not indexes: + raise ValueError("Cosmos3 condition_frame_indexes_vision must contain at least one index.") + if any(index < 0 for index in indexes): + raise ValueError(f"Cosmos3 condition_frame_indexes_vision must be non-negative, got {indexes}.") + return indexes + + +def condition_pixel_frame_count( + condition_frame_indexes_vision: Iterable[int], + temporal_compression: int = COSMOS3_VAE_TEMPORAL_COMPRESSION, +) -> int: + return max(condition_frame_indexes_vision) * int(temporal_compression) + 1 + + +def normalize_condition_video_keep(value: Any) -> str: + keep = str(value or COSMOS3_DEFAULT_CONDITION_VIDEO_KEEP).strip().lower() + if keep not in {"first", "last"}: + raise ValueError("Cosmos3 condition_video_keep must be either 'first' or 'last'.") + return keep + + +def normalize_robolab_action_space(value: Any) -> str: + action_space = str(value or ROBOLAB_DEFAULT_ACTION_SPACE).strip().lower() + aliases = { + "jointpos": "joint_pos", + "joint_pos": "joint_pos", + "abs_ik": "midtrain", + "midtrain": "midtrain", + } + if action_space not in aliases: + raise ValueError(f"Unsupported RoboLab action_space={value!r}; expected joint_pos/jointpos or midtrain/abs_ik.") + return aliases[action_space] + + +def ensure_rgb_uint8_image(value: Any, key: str) -> np.ndarray: + image = np.asarray(value) + if image.ndim != 3 or image.shape[-1] != 3: + raise ValueError(f"{key!r} must have shape [H, W, 3], got {image.shape}.") + if image.dtype != np.uint8: + image = np.clip(image, 0, 255).astype(np.uint8) + return np.ascontiguousarray(image) + + +def ensure_2d_float_array(value: Any, key: str, width: int | None = None) -> np.ndarray: + array = np.asarray(value, dtype=np.float32) + if array.ndim == 1: + array = array[None, :] + if array.ndim != 2: + raise ValueError(f"{key!r} must have shape [T, D] or [D], got {array.shape}.") + if width is not None and array.shape[-1] != width: + raise ValueError(f"{key!r} must have width {width}, got {array.shape[-1]}.") + return np.ascontiguousarray(array) + + +def ensure_gripper_array(value: Any) -> np.ndarray: + array = np.asarray(value, dtype=np.float32) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array[:, None] + if array.ndim != 2 or array.shape[-1] != 1: + raise ValueError(f"'observation/gripper_position' must have shape [T, 1], [T], or scalar, got {array.shape}.") + return np.ascontiguousarray(array) + + +def resize_rgb_uint8(image: np.ndarray, size: tuple[int, int]) -> np.ndarray: + tensor = torch.from_numpy(image).permute(2, 0, 1).unsqueeze(0).float() + resized = F.interpolate(tensor, size=size, mode="bilinear", align_corners=False) + return np.clip(np.round(resized.squeeze(0).permute(1, 2, 0).numpy()), 0, 255).astype(np.uint8) + + +def compose_robolab_views(obs: dict[str, Any]) -> np.ndarray | None: + required_keys = ( + "observation/wrist_image_left", + "observation/exterior_image_1_left", + "observation/exterior_image_2_left", + ) + if not all(key in obs for key in required_keys): + return None + + wrist = ensure_rgb_uint8_image(obs["observation/wrist_image_left"], "observation/wrist_image_left") + left_raw = ensure_rgb_uint8_image(obs["observation/exterior_image_1_left"], "observation/exterior_image_1_left") + right_raw = ensure_rgb_uint8_image(obs["observation/exterior_image_2_left"], "observation/exterior_image_2_left") + half_h, half_w = wrist.shape[0] // 2, wrist.shape[1] // 2 + left = resize_rgb_uint8(left_raw, (half_h, half_w)) + right = resize_rgb_uint8(right_raw, (half_h, half_w)) + return np.concatenate([wrist, np.concatenate([left, right], axis=1)], axis=0) + + +def extract_robolab_image(obs: dict[str, Any]) -> np.ndarray: + if "observation/image" in obs: + return ensure_rgb_uint8_image(obs["observation/image"], "observation/image") + image = compose_robolab_views(obs) + if image is not None: + return image + raise ValueError("Observation must contain 'observation/image' or RoboLab wrist/exterior image keys.") + + +def extract_robolab_prompt_image(prompt_data: Any | None) -> np.ndarray | None: + if not isinstance(prompt_data, dict): + return None + multi_modal_data = prompt_data.get("multi_modal_data", {}) or {} + image = multi_modal_data.get("image") + if image is None: + return None + if isinstance(image, PIL.Image.Image): + return np.asarray(image.convert("RGB")) + return ensure_rgb_uint8_image(image, "multi_modal_data.image") + + +def lazy_import(module_name: str, symbol_name: str, error_message: str): + try: + module = import_module(module_name) + except ModuleNotFoundError as exc: + raise ModuleNotFoundError(error_message) from exc + return getattr(module, symbol_name) + + +def lazy_action_transform_pipeline(max_action_dim: int): + ActionTransformPipeline = lazy_import( + "cosmos_framework.data.vfm.action.transforms", + "ActionTransformPipeline", + "Cosmos3 RoboLab policy serving requires cosmos_framework on PYTHONPATH so the " + "golden ActionTransformPipeline can be reused.", + ) + return ActionTransformPipeline(max_action_dim=max_action_dim, cfg_dropout_rate=0.0) + + +def build_robolab_unipc_scheduler(num_steps: int, shift: float, device: torch.device): + FlowUniPCMultistepScheduler = lazy_import( + "cosmos_framework.model.vfm.diffusion.samplers.fm_solvers_unipc", + "FlowUniPCMultistepScheduler", + ( + "Cosmos3 RoboLab policy serving requires cosmos_framework on PYTHONPATH so the " + "golden FlowUniPCMultistepScheduler can be reused." + ), + ) + + scheduler = FlowUniPCMultistepScheduler( + num_train_timesteps=1000, + shift=1.0, + use_dynamic_shifting=False, + ) + scheduler.set_timesteps(num_steps, device=device, shift=float(shift)) + return scheduler + + +def convert_midtrain_rotation(value: Any, src: str, dst: str) -> np.ndarray: + convert_rotation = lazy_import( + "cosmos_framework.data.vfm.action.pose_utils", + "convert_rotation", + "Cosmos3 RoboLab midtrain action serving requires cosmos_framework pose_utils on PYTHONPATH.", + ) + return convert_rotation(value, src, dst) + + +def pose_abs_to_rel(*args, **kwargs) -> np.ndarray: + pose_abs_to_rel_func = lazy_import( + "cosmos_framework.data.vfm.action.pose_utils", + "pose_abs_to_rel", + "Cosmos3 RoboLab midtrain action serving requires cosmos_framework pose_utils on PYTHONPATH.", + ) + return pose_abs_to_rel_func(*args, **kwargs) + + +def pose_rel_to_abs(*args, **kwargs) -> np.ndarray: + pose_rel_to_abs_func = lazy_import( + "cosmos_framework.data.vfm.action.pose_utils", + "pose_rel_to_abs", + "Cosmos3 RoboLab midtrain action serving requires cosmos_framework pose_utils on PYTHONPATH.", + ) + return pose_rel_to_abs_func(*args, **kwargs) + + +def build_abs_pose_from_components(*args, **kwargs) -> np.ndarray: + build_abs_pose_from_components_func = lazy_import( + "cosmos_framework.data.vfm.action.pose_utils", + "build_abs_pose_from_components", + "Cosmos3 RoboLab midtrain action serving requires cosmos_framework pose_utils on PYTHONPATH.", + ) + return build_abs_pose_from_components_func(*args, **kwargs) + + +def next_robolab_seed(extra: dict[str, Any], obs: dict[str, Any], request_id: str | None) -> int: + base_seed = int(extra.get("robolab_seed") or ROBOLAB_DEFAULT_SEED) + deterministic_seed = str(extra.get("deterministic_seed", "")).strip().lower() in {"1", "true", "yes", "on"} + if deterministic_seed: + return base_seed + explicit_seed = extra.get("seed") + if explicit_seed is not None: + return int(explicit_seed) + seed_key = "|".join( + str(part) + for part in ( + base_seed, + extra.get("session_id", ""), + request_id or "", + obs.get("prompt", ""), + ) + ) + return zlib.crc32(seed_key.encode("utf-8")) & 0x7FFFFFFF + + +def log_robolab_action_summary(label: str, value: Any) -> None: + if not _is_rank_zero(): + return + if isinstance(value, torch.Tensor): + array = value.detach().float().cpu().numpy() + else: + array = np.asarray(value, dtype=np.float32) + finite = np.isfinite(array) + if finite.any(): + finite_min = float(array[finite].min()) + finite_max = float(array[finite].max()) + else: + finite_min = None + finite_max = None + if array.ndim == 0: + head = array.reshape(1).tolist() + else: + head = array.reshape(-1, array.shape[-1])[:3].tolist() + logger.info( + "RoboLab action summary %s: shape=%s nan=%d finite=%d finite_min=%s finite_max=%s head=%s", + label, + tuple(array.shape), + int(np.isnan(array).sum()), + int(finite.sum()), + finite_min, + finite_max, + head, + ) + + +def postprocess_robolab_action(action: torch.Tensor, inputs: RoboLabActionPostprocessInputs) -> np.ndarray: + action_np = action[0].float().cpu().numpy() + log_robolab_action_summary("raw_model_action", action_np) + history_length = int(inputs.history_length) + action_np = action_np[history_length:] + action_np[:, -1] = 1.0 - action_np[:, -1] + + if inputs.action_space == "midtrain": + if inputs.eef_pos is None or inputs.eef_quat is None: + raise ValueError("RoboLab midtrain action postprocess requires eef_pos and eef_quat metadata.") + initial_pose = np.eye(4, dtype=np.float32) + initial_pose[:3, :3] = convert_midtrain_rotation(inputs.eef_quat[-1], "quat_xyzw", "matrix") + initial_pose[:3, 3] = inputs.eef_pos[-1] + abs_pose = pose_rel_to_abs( + action_np[:, :ROBOLAB_MIDTRAIN_POSE_ACTION_DIM], + rotation_format="rot6d", + pose_convention="backward_framewise", + initial_pose=initial_pose, + ) + position = abs_pose[1:, :3, 3] + quat_xyzw = convert_midtrain_rotation(abs_pose[1:, :3, :3], "matrix", "quat_xyzw") + action_np = np.concatenate([position, quat_xyzw, action_np[:, ROBOLAB_MIDTRAIN_POSE_ACTION_DIM:]], axis=-1) + + log_robolab_action_summary("postprocessed_robolab_action", action_np) + return np.asarray(action_np, dtype=np.float32) diff --git a/vllm_omni/diffusion/registry.py b/vllm_omni/diffusion/registry.py index 34ac5e31b71..9cc65486a7a 100644 --- a/vllm_omni/diffusion/registry.py +++ b/vllm_omni/diffusion/registry.py @@ -512,6 +512,20 @@ def _apply_sequence_parallel_if_enabled(model, od_config: OmniDiffusionConfig) - "HiDreamImagePipeline": "get_hidream_image_post_process_func", } +_DIFFUSION_ACTION_POST_PROCESS_FUNCS = { + # arch: action_post_process_func + # `action_post_process_func` function must be placed in {mod_folder}/{mod_relname}.py, + # where mod_folder and mod_relname are defined and mapped using `_DIFFUSION_MODELS` via the `arch` key. + "Cosmos3OmniDiffusersPipeline": "get_cosmos3_action_post_process_func", +} + +_DIFFUSION_IR_OP_PRIORITY_FUNCS = { + # arch: ir_op_priority_func + # `ir_op_priority_func` function must be placed in {mod_folder}/{mod_relname}.py, + # where mod_folder and mod_relname are defined and mapped using `_DIFFUSION_MODELS` via the `arch` key. + "Cosmos3OmniDiffusersPipeline": "get_cosmos3_ir_op_priority_func", +} + _DIFFUSION_PRE_PROCESS_FUNCS = { # arch: pre_process_func # `pre_process_func` function must be placed in {mod_folder}/{mod_relname}.py, @@ -543,6 +557,8 @@ def register_diffusion_model( class_name: str, pre_process_func_name: str | None = None, post_process_func_name: str | None = None, + action_post_process_func_name: str | None = None, + ir_op_priority_func_name: str | None = None, ) -> None: """Register a diffusion model pipeline from an out-of-tree plugin. @@ -561,6 +577,12 @@ def register_diffusion_model( post_process_func_name: Optional name of the post-process function located in *module_name*. Pass ``None`` to keep the existing entry when replacing a built-in model. + action_post_process_func_name: Optional name of the action post-process + function located in *module_name*. Pass ``None`` to keep the + existing entry when replacing a built-in model. + ir_op_priority_func_name: Optional name of the IR op priority merge + function located in *module_name*. Pass ``None`` to keep the + existing entry when replacing a built-in model. """ # Register model class in DiffusionModelRegistry DiffusionModelRegistry.register_model( @@ -578,6 +600,10 @@ def register_diffusion_model( _DIFFUSION_PRE_PROCESS_FUNCS[model_arch] = pre_process_func_name if post_process_func_name is not None: _DIFFUSION_POST_PROCESS_FUNCS[model_arch] = post_process_func_name + if action_post_process_func_name is not None: + _DIFFUSION_ACTION_POST_PROCESS_FUNCS[model_arch] = action_post_process_func_name + if ir_op_priority_func_name is not None: + _DIFFUSION_IR_OP_PRIORITY_FUNCS[model_arch] = ir_op_priority_func_name logger.info( "Registered diffusion model %s -> %s.%s", @@ -608,6 +634,20 @@ def get_diffusion_post_process_func(od_config: OmniDiffusionConfig): return _load_process_func(od_config, func_name) +def get_diffusion_action_post_process_func(od_config: OmniDiffusionConfig): + if od_config.model_class_name not in _DIFFUSION_ACTION_POST_PROCESS_FUNCS: + return None + func_name = _DIFFUSION_ACTION_POST_PROCESS_FUNCS[od_config.model_class_name] + return _load_process_func(od_config, func_name) + + +def get_diffusion_ir_op_priority_func(od_config: OmniDiffusionConfig): + if od_config.model_class_name not in _DIFFUSION_IR_OP_PRIORITY_FUNCS: + return None + func_name = _DIFFUSION_IR_OP_PRIORITY_FUNCS[od_config.model_class_name] + return _load_process_func(od_config, func_name) + + def get_diffusion_pre_process_func(od_config: OmniDiffusionConfig): if od_config.model_class_name not in _DIFFUSION_PRE_PROCESS_FUNCS: return None # Return None if no pre-processing function is registered (for backward compatibility) diff --git a/vllm_omni/diffusion/worker/diffusion_worker.py b/vllm_omni/diffusion/worker/diffusion_worker.py index cd9576437fa..e3adfecc69e 100644 --- a/vllm_omni/diffusion/worker/diffusion_worker.py +++ b/vllm_omni/diffusion/worker/diffusion_worker.py @@ -43,6 +43,7 @@ from vllm_omni.diffusion.forward_context import set_forward_context from vllm_omni.diffusion.ipc import pack_diffusion_output_shm from vllm_omni.diffusion.lora.manager import DiffusionLoRAManager +from vllm_omni.diffusion.registry import get_diffusion_ir_op_priority_func from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.sched.interface import DiffusionSchedulerOutput from vllm_omni.diffusion.worker.diffusion_model_runner import DiffusionModelRunner @@ -159,6 +160,14 @@ def _create_diffusion_worker_vllm_config(device: torch.device, od_config: OmniDi return vllm_config +def _resolve_ir_op_priority(od_config: OmniDiffusionConfig, vllm_config: VllmConfig) -> Any: + ir_op_priority = current_omni_platform.get_default_ir_op_priority(vllm_config) + ir_op_priority_func = get_diffusion_ir_op_priority_func(od_config) + if ir_op_priority_func is not None: + ir_op_priority = ir_op_priority_func(ir_op_priority, vllm_config=vllm_config) + return ir_op_priority + + class DiffusionWorker: """ A worker that manages GPU infrastructure and delegates to the model runner. @@ -234,7 +243,7 @@ def init_device(self) -> None: vllm_config.quant_config = self.od_config.quantization_config # Since vLLM v0.20.0, IR wraps GPU ops. Set IR op priority preference to enforce GPU op fusion during wrapping. # Also need to log, because vLLM internally logs another line in VllmConfig.__post_init__. Avoid confusion. - vllm_config.kernel_config.ir_op_priority = current_omni_platform.get_default_ir_op_priority(vllm_config) + vllm_config.kernel_config.ir_op_priority = _resolve_ir_op_priority(self.od_config, vllm_config) logger.info( "Final IR op priority after setting vLLM-Omni overrides: %s", vllm_config.kernel_config.ir_op_priority ) diff --git a/vllm_omni/entrypoints/openai/serving_video.py b/vllm_omni/entrypoints/openai/serving_video.py index ac6038f9285..4bdbfce0217 100644 --- a/vllm_omni/entrypoints/openai/serving_video.py +++ b/vllm_omni/entrypoints/openai/serving_video.py @@ -195,18 +195,21 @@ async def _run_and_extract( # Merge extra_params into extra_args gen_params.extra_args.update(request.extra_params) - # Redact the inline ``action`` array (hundreds of floats) when - # logging so it doesn't flood the logs; everything else is logged - # verbatim. + # Redact inline arrays when logging so RoboLab policy requests do + # not flood the server log with image/state payloads. loggable = request.extra_params - action_val = loggable.get("action") - if action_val is not None: - summary = ( - f"<{type(action_val).__name__} len={len(action_val)}>" - if hasattr(action_val, "__len__") - else f"<{type(action_val).__name__}>" + redacted = {} + for key in ("action", "robot_obs", "observation"): + value = loggable.get(key) + if value is None: + continue + redacted[key] = ( + f"<{type(value).__name__} len={len(value)}>" + if hasattr(value, "__len__") + else f"<{type(value).__name__}>" ) - loggable = {**loggable, "action": summary} + if redacted: + loggable = {**loggable, **redacted} logger.info("Applied extra_params: %s", loggable) self._apply_lora(request.lora, gen_params) @@ -220,7 +223,9 @@ async def _run_and_extract( ) result = await self._run_generation(prompt, gen_params, reference_id) - videos = self._extract_video_outputs(result) + custom_output = self._extract_custom_output(result) + action_only = isinstance(custom_output, dict) and bool(custom_output.get("action_only_output")) + videos = [{"action_only_output": True}] if action_only else self._extract_video_outputs(result) audios = self._extract_audio_outputs(result, expected_count=len(videos)) actions = self._extract_action_outputs(result, expected_count=len(videos)) audio_sample_rate = self._resolve_audio_sample_rate(result) @@ -314,6 +319,11 @@ async def generate_video_bytes( if "video_codec_options" in request.extra_params: video_codec_options = request.extra_params["video_codec_options"] + action = artifacts.actions[0] + if action is not None and isinstance(artifacts.videos[0], dict): + logger.info("Action-only video request %s completed; skipping MP4 encoding.", reference_id) + return b"", artifacts.stage_durations, artifacts.peak_memory_mb, action + _t_encode_start = time.perf_counter() video_bytes = _encode_video_bytes( artifacts.videos[0], @@ -510,7 +520,8 @@ def _extract_action_outputs(cls, result: Any, expected_count: int) -> list[Video if not custom_output or "action" not in custom_output: return [None] * expected_count - action_items = cls._split_action_payload(custom_output["action"], expected_count) + action_payload = custom_output.get("actions", custom_output["action"]) + action_items = cls._split_action_payload(action_payload, expected_count) return [ cls._make_video_action(action_item, custom_output) if action_item is not None else None for action_item in action_items diff --git a/vllm_omni/entrypoints/openpi/connection.py b/vllm_omni/entrypoints/openpi/connection.py index 670c37a9233..69f722ac824 100644 --- a/vllm_omni/entrypoints/openpi/connection.py +++ b/vllm_omni/entrypoints/openpi/connection.py @@ -7,6 +7,10 @@ Connect -> server sends msgpack(PolicyServerConfig fields) Infer -> client sends msgpack(obs), server sends msgpack(ndarray) Reset -> client sends msgpack({endpoint:reset}), server sends msgpack(status) + +NumPy values use the msgpack-numpy marker mapping: + ndarray -> {nd: true, type, kind, shape, data} + scalar -> {nd: false, type, kind, data} """ from __future__ import annotations @@ -14,6 +18,8 @@ import asyncio from typing import Any +import msgspec +import numpy as np from fastapi import WebSocket from starlette.websockets import WebSocketDisconnect from vllm.logger import init_logger @@ -25,26 +31,75 @@ logger = init_logger(__name__) _DEFAULT_IDLE_TIMEOUT = 30.0 MAX_OPENPI_PAYLOAD_BYTES = 64 * 1024 * 1024 - - -def _get_msgpack_numpy() -> Any: - try: - from openpi_client import msgpack_numpy - except ImportError as exc: - raise ImportError( - "The `/v1/realtime/robot/openpi` endpoint requires the optional " - "`openpi-client` dependency. Install it with `pip install openpi-client`." - ) from exc - - return msgpack_numpy +_MISSING = object() + + +def _pack_numpy(obj: Any) -> Any: + if isinstance(obj, (np.ndarray, np.generic)) and obj.dtype.kind in ("V", "O", "c"): + raise ValueError(f"Unsupported dtype: {obj.dtype}") + if isinstance(obj, np.ndarray): + if not obj.flags.c_contiguous: + obj = np.ascontiguousarray(obj) + return { + b"nd": True, + b"data": obj.tobytes(), + b"type": obj.dtype.str, + b"kind": obj.dtype.kind, + b"shape": obj.shape, + } + if isinstance(obj, np.generic): + return { + b"nd": False, + b"data": obj.tobytes(), + b"type": obj.dtype.str, + b"kind": obj.dtype.kind, + } + raise TypeError(f"Unsupported type: {type(obj)!r}") + + +def _mapping_get(obj: dict[Any, Any], key: str, default: Any = None) -> Any: + return obj.get(key, obj.get(key.encode(), default)) + + +def _decode_marker_text(value: Any) -> str: + if isinstance(value, bytes): + return value.decode() + return str(value) + + +def _unpack_numpy(obj: Any) -> Any: + if isinstance(obj, dict): + nd = _mapping_get(obj, "nd", _MISSING) + dtype = _mapping_get(obj, "type", _MISSING) + kind = _mapping_get(obj, "kind", _MISSING) + data = _mapping_get(obj, "data", _MISSING) + if nd is not _MISSING and dtype is not _MISSING and kind is not _MISSING and data is not _MISSING: + dtype_obj = np.dtype(_decode_marker_text(dtype)) + kind_text = _decode_marker_text(kind) + if dtype_obj.kind != kind_text: + raise ValueError(f"NumPy dtype marker kind mismatch: {dtype_obj.kind!r} != {kind_text!r}") + if dtype_obj.kind in ("V", "O", "c"): + raise ValueError(f"Unsupported dtype: {dtype_obj}") + + array = np.frombuffer(data, dtype=dtype_obj).copy() + if nd: + shape = _mapping_get(obj, "shape", _MISSING) + if shape is _MISSING: + raise ValueError("NumPy ndarray marker is missing shape") + return array.reshape(tuple(shape)) + return array[0] + return {key: _unpack_numpy(value) for key, value in obj.items()} + if isinstance(obj, list): + return [_unpack_numpy(value) for value in obj] + return obj def _pack(obj: Any) -> bytes: - return _get_msgpack_numpy().packb(obj) + return msgspec.msgpack.encode(obj, enc_hook=_pack_numpy) def _unpack(data: bytes) -> Any: - return _get_msgpack_numpy().unpackb(data) + return _unpack_numpy(msgspec.msgpack.decode(data)) class RobotRealtimeConnection: