diff --git a/.buildkite/common/ci_source_file_dependencies.yml b/.buildkite/common/ci_source_file_dependencies.yml index d4b1e559800..52cca36914a 100644 --- a/.buildkite/common/ci_source_file_dependencies.yml +++ b/.buildkite/common/ci_source_file_dependencies.yml @@ -303,7 +303,7 @@ source_file_dependencies: - tests/examples/test_minicpmo_realtime_web_static.py - tests/examples/test_minicpmo_realtime_duplex_simple_demo.py - omni_minicpmo_4_5_duplex_function: + omni_minicpmo_4_5_duplex_function: &minicpmo_4_5_duplex_function - *minicpmo_4_5 - *minicpmo_4_5_duplex - tests/e2e/online_serving/test_minicpmo_4_5_duplex.py @@ -328,6 +328,14 @@ source_file_dependencies: - tests/model_executor/models/minicpmo_4_5/duplex/ - tests/protocol/ + omni_minicpmo_4_5_duplex_mrv2_function: + - *minicpmo_4_5_duplex_function + - vllm_omni/worker_v2/ + - vllm_omni/outputs/ + - vllm_omni/deploy/minicpmo_4_5_duplex_mrv2.yaml + - tests/e2e/online_serving/test_minicpmo_4_5_duplex_mrv2.py + - tests/worker_v2/ + omni_minicpmo_4_5_duplex_perf: - *minicpmo_4_5 - *minicpmo_4_5_duplex diff --git a/.buildkite/cuda/test-ready.yml b/.buildkite/cuda/test-ready.yml index 90fe30fa2ad..190bd868fcf 100644 --- a/.buildkite/cuda/test-ready.yml +++ b/.buildkite/cuda/test-ready.yml @@ -207,6 +207,13 @@ steps: " mirror_hardwares: h100_1 + - label: "Omni · MiniCPM-o 4.5 MRv2 Duplex Test" + source_file_dependencies: omni_minicpmo_4_5_duplex_mrv2_function + timeout_in_minutes: 40 + commands: + - timeout 35m pytest -s -v tests/e2e/online_serving/test_minicpmo_4_5_duplex_mrv2.py -m 'advanced_model and cuda' --run-level advanced_model + mirror_hardwares: h100_1 + - label: "Omni · MiniCPM-o 4.5 Perf Test" timeout_in_minutes: 60 source_file_dependencies: omni_minicpmo_4_5_perf diff --git a/docs/design/minicpm_o45_mrv2_performance.md b/docs/design/minicpm_o45_mrv2_performance.md index 13b4f1a682b..87771b9f619 100644 --- a/docs/design/minicpm_o45_mrv2_performance.md +++ b/docs/design/minicpm_o45_mrv2_performance.md @@ -17,8 +17,9 @@ ordinary TF32 for CFM DiT GEMMs within Code2Wav forward/capture only `input_precision="tf32"` on tiled attention). This is not compensated TF32x3. The previous process matmul policy is restored afterwards. cuDNN's TF32 policy is independent. HiFT stays IEEE FP32. TF32 changes rounding. -Duplex keeps the mainline V1 session path. Turn results do not establish duplex -performance or interruption correctness. +Duplex defaults to the V1 session path; `minicpmo_4_5_duplex_mrv2.yaml` is the +opt-in CUDA overlay that runs all three duplex stages on MRv2. Turn results do +not establish duplex performance or interruption correctness. The shared Code2Wav backend defaults to the existing fused CFM body on CUDA when TF32 is allowed, subject to architecture and FP32 attention-cache @@ -50,3 +51,15 @@ Disabling TF32 with `token2wav_allow_tf32: false` or `cfm_fused_body: true` overrides that choice and may use TF32 attention. Slot pooling and row-offset merging remain opt-in. Fusion changes floating-point rounding and does not promise identical waveforms. + +## Full-duplex MRv2 + +`minicpmo_4_5_duplex_mrv2.yaml` opts all three CUDA stages into MRv2 and +otherwise inherits the default duplex profile; non-CUDA platforms keep V1. +Stage 0 keeps `async_chunk: false` because the Thinker has no async producer: +as on V1, the orchestrator hands each finished segment to `llm2tts`. +The Talker reuses the streaming prompt recipe (full attention extends its KV +prefix until capacity; sliding recompute rebuilds the previous condition, +confirmed codec ids and current condition) and keeps the 16-frame codec +penalty history across conditions. A segment flushes only on an EOS sampled +in the current output, never on `turn_end` alone. diff --git a/tests/config/test_minicpmo_4_5_duplex_mrv2_deploy.py b/tests/config/test_minicpmo_4_5_duplex_mrv2_deploy.py new file mode 100644 index 00000000000..c8f7e82cfff --- /dev/null +++ b/tests/config/test_minicpmo_4_5_duplex_mrv2_deploy.py @@ -0,0 +1,49 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +"""Resolve the MiniCPM-o three-stage duplex MRv2 deployment contract.""" + +import pytest + +from tests.helpers.stage_config import get_deploy_config_path +from vllm_omni.config.stage_config import ( + _apply_platform_overrides, + load_deploy_config, + merge_pipeline_deploy, +) +from vllm_omni.model_executor.models.minicpmo_4_5.pipeline import MINICPMO_4_5_PIPELINE +from vllm_omni.platforms import current_omni_platform + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + +_DEPLOY = "minicpmo_4_5_duplex_mrv2.yaml" + + +def _resolve_cuda_stages(monkeypatch): + # Resolve CUDA overrides even when this test runs on a non-CUDA host. + monkeypatch.setattr(current_omni_platform, "device_name", "cuda") + config = _apply_platform_overrides(load_deploy_config(get_deploy_config_path(_DEPLOY)), platform="cuda") + return config, merge_pipeline_deploy(MINICPMO_4_5_PIPELINE, config) + + +def test_duplex_mrv2_cuda_profile(monkeypatch) -> None: + config, stages = _resolve_cuda_stages(monkeypatch) + assert config.session_mode == "duplex" + # All three stages use MRv2; there is no hidden Talker fallback. + assert [s.yaml_engine_args["use_v2_model_runner"] for s in stages] == [True, True, True] + # Stage 0 duplex preprocessing stays on the synchronous path while the + # downstream stages keep streaming chunk transfer. + assert [s.yaml_engine_args["async_chunk"] for s in stages] == [False, True, True] + # Preserve asynchronous AR scheduling, including the Talker. + assert all(s.yaml_engine_args["async_scheduling"] for s in stages[:2]) + assert stages[2].yaml_engine_args.get("async_scheduling") is not False + # Retain the base profile capacity and Talker KV budget. + assert [s.yaml_engine_args["max_num_seqs"] for s in stages] == [16, 16, 16] + assert stages[1].yaml_engine_args["kv_cache_memory_bytes"] == 4 * 1024**3 + + +@pytest.mark.parametrize("platform", ["npu", "xpu", "rocm", "musa"]) +def test_duplex_mrv2_profile_keeps_non_cuda_on_v1(monkeypatch, platform) -> None: + monkeypatch.setattr(current_omni_platform, "device_name", platform) + config = _apply_platform_overrides(load_deploy_config(get_deploy_config_path(_DEPLOY)), platform=platform) + stages = merge_pipeline_deploy(MINICPMO_4_5_PIPELINE, config) + assert [s.yaml_engine_args["use_v2_model_runner"] for s in stages] == [False, False, False] diff --git a/tests/config/test_omni_config.py b/tests/config/test_omni_config.py index 31c190ccef6..b8578753134 100644 --- a/tests/config/test_omni_config.py +++ b/tests/config/test_omni_config.py @@ -2365,3 +2365,23 @@ def test_async_chunk_auto_disabled_without_processor(): # since doing so will just raise a ValueError in validation. merge_pipeline_deploy(pipeline, deploy) assert not deploy.async_chunk + + +def test_mrv2_duplex_session_validation(): + from vllm_omni.config.stage_config import validate_native_mrv2_session + + ps = StagePipelineConfig( + stage_id=1, + model_stage="tts", + input_sources=(0,), + supports_native_mrv2_data_plane=True, + supports_duplex_mrv2=True, + ) + # Allowed when supports_duplex_mrv2 is True + deploy_duplex = DeployConfig(session_mode="duplex") + validate_native_mrv2_session(deploy_duplex, ps, "v2") + + # Rejected when supports_duplex_mrv2 is False + ps_no_duplex = replace(ps, supports_duplex_mrv2=False) + with pytest.raises(ValueError, match="supports session_mode 'turn' only"): + validate_native_mrv2_session(deploy_duplex, ps_no_duplex, "v2") diff --git a/tests/core/sched/test_omni_ar_scheduler_streaming.py b/tests/core/sched/test_omni_ar_scheduler_streaming.py index 785014ff5ec..56bb2b280a6 100644 --- a/tests/core/sched/test_omni_ar_scheduler_streaming.py +++ b/tests/core/sched/test_omni_ar_scheduler_streaming.py @@ -1587,3 +1587,52 @@ def fake_schedule(self, _throttle_prefills=False): assert observed_limits == [7 if native else 8] assert sched.max_num_active_reqs == 8 + + +@pytest.mark.parametrize( + ("attention_type", "prompt_len", "recompute"), + [("full_attention", 70, True), ("full_attention", 20, False), ("sliding_recompute", 20, True)], +) +def test_mrv2_talker_reuses_confirmed_prompt_window(attention_type, prompt_len, recompute) -> None: + sched = _make_scheduler(stage_id=1, session_mode="duplex") + sched._native_data_plane = True + sched.vllm_config.model_config.custom_process_next_stage_input_func = ( + "vllm_omni.model_executor.stage_input_processors.minicpmo_4_5_omni.tts2code2wav_async_chunk" + ) + sched.max_model_len = 100 + sched.vllm_config.model_config.hf_config_name = "tts_config" + sched.vllm_config.model_config.max_model_len = 100 + sched.vllm_config.model_config.hf_config = SimpleNamespace( + model_type="minicpmtts", + max_position_embeddings=100, + attention_type=attention_type, + ) + sched.finish_requests = MagicMock() + session = _make_request() + session.external_req_id = "native-talker" + session.prompt_token_ids = [0] * prompt_len + session._all_token_ids.clear() + session._all_token_ids.extend(session.prompt_token_ids) + session.append_output_token_ids([7, 8, 9, 6561]) + session.num_prompt_tokens = prompt_len + session.num_computed_tokens = prompt_len + 3 # EOS was sampled, not fed back into the model. + session.status = RequestStatus.WAITING_FOR_STREAMING_REQ + sched.num_waiting_for_streaming_input = 1 + session.model_intermediate_buffer = { + "native_duplex": True, + "meta": { + "next_stage_prompt_len": 20, + "streaming_condition_seq": 0, + }, + } + update = _make_talker_update(20, reserve=10, condition_seq=1) + sched._update_request_as_session(session, update) + assert session.num_computed_tokens == (0 if recompute else prompt_len + 3) + assert session.num_prompt_tokens == 43 + assert update.model_intermediate_buffer["ids"]["streaming_prompt_previous_codes"] == [7, 8, 9] + assert update.model_intermediate_buffer["meta"]["streaming_prompt_recompute"] is recompute + sched.finish_requests.assert_not_called() + if recompute: + sched._free_request_blocks.assert_called_once_with(session) + else: + sched._free_request_blocks.assert_not_called() diff --git a/tests/e2e/online_serving/test_minicpmo_4_5_duplex_mrv2.py b/tests/e2e/online_serving/test_minicpmo_4_5_duplex_mrv2.py new file mode 100644 index 00000000000..0fcc3f4956c --- /dev/null +++ b/tests/e2e/online_serving/test_minicpmo_4_5_duplex_mrv2.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +"""Real-weight MRv2 duplex concurrency guard, additional to the V1 CI jobs.""" + +import asyncio +from pathlib import Path + +import pytest + +from tests.e2e.online_serving.helpers.minicpmo_4_5_duplex import ( + MODEL, + multi_session_args, + realtime_url, + resolve_ref_audio, + validated_input_wav, +) +from tests.e2e.online_serving.run_minicpmo_realtime_duplex_multi_session import run_multi_session +from tests.e2e.online_serving.test_minicpmo_4_5_duplex import _run_seeded_text_to_audio +from tests.helpers.mark import hardware_test +from tests.helpers.runtime import OmniServerParams +from tests.helpers.stage_config import get_deploy_config_path + +pytestmark = [pytest.mark.omni, pytest.mark.advanced_model] + +_SERVER = OmniServerParams( + model=MODEL, + stage_config_path=get_deploy_config_path("minicpmo_4_5_duplex_mrv2.yaml"), + use_stage_cli=False, + server_args=["--trust-remote-code"], +) + + +@hardware_test(res={"cuda": "H100"}, num_cards=1) +@pytest.mark.parametrize("omni_server", [pytest.param(_SERVER, id="mrv2-real-weights")], indirect=True) +@pytest.mark.parametrize("sessions", [1, 2, 4]) +def test_mrv2_duplex_overlapping_turns(omni_server, tmp_path: Path, sessions: int): + args = multi_session_args( + omni_server=omni_server, + input_wav=validated_input_wav(), + ref_audio=resolve_ref_audio(), + output_dir=tmp_path / f"concurrency_{sessions}", + response_required=True, + ) + args.sessions = sessions + args.turns = 2 + args.turn_duration_ms = [args.first_turn_ms] * args.turns + args.synchronized_start = True + args.disconnect_session_index = None + args.takeover_session_index = None + result = asyncio.run(run_multi_session(args)) + assert result["ok"] is True, result + assert result["session_count"] == sessions + assert result["identity_isolation_ok"] is True + assert not result["failures"] + for session in result["sessions"]: + assert session["audio_delta_count"] > 0, session + assert session["done_count"] == args.turns, session + assert session["error_count"] == 0, session + + +@hardware_test(res={"cuda": "H100"}, num_cards=1) +@pytest.mark.parametrize("omni_server", [pytest.param(_SERVER, id="mrv2-real-weights")], indirect=True) +def test_mrv2_duplex_new_session_during_long_response(omni_server): + # Unequal prompts and a delayed second session prevent identical clients + # from advancing in lockstep and hiding shared condition-state bugs. + ref_audio = resolve_ref_audio() + + async def speak(text, delay): + await asyncio.sleep(delay) + return await _run_seeded_text_to_audio( + url=realtime_url(omni_server), + model=omni_server.model, + ref_audio=ref_audio, + text=text, + silence_seconds=30.0, + ) + + async def overlap(): + return await asyncio.gather( + speak("What is the capital of China? Answer in about 40 words.", 0), + speak("What is the capital of France? Answer in one short sentence.", 2), + ) + + first, second = asyncio.run(overlap()) + for result in (first, second): + assert "response.done" in result["event_types"], result + assert result["audio_bytes"] > 0, result + assert str(result["transcript"]).strip(), result + assert first["audio_bytes"] > 96_000, first + assert str(first["transcript"]).strip() != str(second["transcript"]).strip() + + +@hardware_test(res={"cuda": "H100"}, num_cards=1) +@pytest.mark.parametrize("omni_server", [pytest.param(_SERVER, id="mrv2-real-weights")], indirect=True) +def test_mrv2_duplex_server_answers_image_chat_from_the_image(omni_server, openai_client): + """The duplex Thinker also serves /v1/chat/completions: its image features must reach the prompt.""" + import base64 + import io + + from PIL import Image + + buffer = io.BytesIO() + Image.new("RGB", (224, 224), (255, 0, 0)).save(buffer, format="JPEG") + image_url = "data:image/jpeg;base64," + base64.b64encode(buffer.getvalue()).decode("ascii") + request_config = { + "model": omni_server.model, + "messages": [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": image_url}}, + {"type": "text", "text": "What color is this image? Answer with one word."}, + ], + } + ], + "stream": True, + "modalities": ["text"], + "key_words": {"text": ["red"]}, + "extra_body": {"chat_template_kwargs": {"enable_thinking": False}}, + } + openai_client.send_omni_request(request_config) diff --git a/tests/model_executor/models/minicpmo_4_5/duplex/test_window_wiring.py b/tests/model_executor/models/minicpmo_4_5/duplex/test_window_wiring.py index aebf1331255..17f77f7a7ac 100644 --- a/tests/model_executor/models/minicpmo_4_5/duplex/test_window_wiring.py +++ b/tests/model_executor/models/minicpmo_4_5/duplex/test_window_wiring.py @@ -1479,6 +1479,72 @@ def _evict_window_units_for_reanchor(self, state, reanchor): assert torch.equal(k_pool, k_pool_after_first) +@pytest.mark.parametrize("explicit_old_computed", [False, True]) +def test_mrv2_reanchor_uses_request_slot_without_persistent_input_batch(explicit_old_computed): + reanchor = {"delta": 16, "moved_from": 48, "sink_blocks": 2} + if explicit_old_computed: + reanchor["old_computed_tokens"] = 64 + pool = torch.randn(8, 1, 16, 16) + original = pool.clone() + inv_freq = _get_inv_freq(head_dim=16) + block_ids = [0, 1, 2, 3] + info = {"duplex": {"stage0_reanchor": dict(reanchor)}} + runner = SimpleNamespace( + req_states=SimpleNamespace(req_id_to_index={"r": 2}, num_computed_tokens_np=np.array([0, 0, 48])), + model_state=SimpleNamespace(intermediate_buffer=SimpleNamespace(buffers=[{}, {}, info])), + block_tables=SimpleNamespace( + num_blocks=SimpleNamespace(np=np.array([[0, 0, 4]])), + block_tables=[SimpleNamespace(gpu=torch.tensor([[0] * 4, [0] * 4, block_ids]))], + ), + device=torch.device("cpu"), + cache_config=SimpleNamespace(block_size=16), + kv_caches=[pool], + _duplex_inv_freq=inv_freq, + ) + scheduler = SimpleNamespace(num_scheduled_tokens={"r": 1}, scheduled_new_reqs=[]) + MiniCPMO45DuplexWorkerHelper.maybe_apply_reanchor(runner, scheduler_output=scheduler) + # Only the retained tail in compacted block 2 is shifted; sink and + # unrelated pages stay untouched and scheduler counts are not decremented twice. + torch.testing.assert_close(pool[2], rotate_keys(original[2], 16, inv_freq), rtol=1e-5, atol=1e-6) + assert torch.equal(pool[:2], original[:2]) + assert torch.equal(pool[3:], original[3:]) + assert runner.req_states.num_computed_tokens_np.tolist() == [0, 0, 48] + assert "stage0_reanchor" not in info["duplex"] + + +def test_mrv2_reanchor_reads_each_group_block_table_once(monkeypatch): + pools = [torch.randn(8, 1, 16, 16) for _ in range(3)] + originals = [pool.clone() for pool in pools] + inv_freq = _get_inv_freq(head_dim=16) + info = {"duplex": {"stage0_reanchor": {"delta": 16, "moved_from": 48, "sink_blocks": 2}}} + runner = SimpleNamespace( + req_states=SimpleNamespace(req_id_to_index={"r": 0}, num_computed_tokens_np=np.array([48])), + model_state=SimpleNamespace(intermediate_buffer=SimpleNamespace(buffers=[info])), + block_tables=SimpleNamespace( + num_blocks=SimpleNamespace(np=np.array([[4]])), + block_tables=[SimpleNamespace(gpu=torch.tensor([[0, 1, 2, 3]]))], + ), + device=torch.device("cpu"), + cache_config=SimpleNamespace(block_size=16), + kv_caches=pools, + _duplex_inv_freq=inv_freq, + ) + resolve = MiniCPMO45DuplexWorkerHelper.resolve_group_block_ids + calls = [] + + def counted(runner_, req_id, req_idx, group_idx=0): + calls.append(group_idx) + return resolve(runner_, req_id, req_idx, group_idx=group_idx) + + monkeypatch.setattr(MiniCPMO45DuplexWorkerHelper, "resolve_group_block_ids", staticmethod(counted)) + MiniCPMO45DuplexWorkerHelper.maybe_apply_reanchor( + runner, scheduler_output=SimpleNamespace(num_scheduled_tokens={"r": 1}, scheduled_new_reqs=[]) + ) + assert calls == [0] + for pool, original in zip(pools, originals): + torch.testing.assert_close(pool[2], rotate_keys(original[2], 16, inv_freq), rtol=1e-5, atol=1e-6) + + def test_history_slicing_multi_row_embeddings_and_non_aligned_prefix(): """Verify that history eviction properly handles: 1. Non-aligned prefix: tokens in [prefix_tokens, sink_end) are kept in the sink! diff --git a/tests/model_executor/models/minicpmo_4_5/test_duplex_mrv2_sampling.py b/tests/model_executor/models/minicpmo_4_5/test_duplex_mrv2_sampling.py new file mode 100644 index 00000000000..c30a99d0f94 --- /dev/null +++ b/tests/model_executor/models/minicpmo_4_5/test_duplex_mrv2_sampling.py @@ -0,0 +1,287 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +from types import SimpleNamespace + +import numpy as np +import pytest +import torch +from vllm.sampling_params import SamplingParams + +from tests.helpers.mark import hardware_test +from vllm_omni.model_executor.duplex_sampling import DuplexSamplingHelper +from vllm_omni.model_executor.models.minicpmo_4_5.duplex.mrv2_sampling import ( + MiniCPMO45DuplexSampler, + _v1_shaped_runner, +) + +pytestmark = [pytest.mark.core_model] + + +@pytest.mark.cpu +def test_mrv2_rows_read_sampling_params_from_intermediate_buffers(): + helper = DuplexSamplingHelper() + infos = { + "a": {"duplex": {"data_plane": True}, "sampling_params": SamplingParams(temperature=0.2, top_k=7, top_p=0.5)}, + "b": {"duplex": {"data_plane": True}, "sampling_params": SamplingParams(temperature=0.8, top_k=15, top_p=0.9)}, + } + runner = _v1_shaped_runner(SimpleNamespace(req_ids=["b", "a"]), infos) + for request_id in infos: + helper.refresh_active_request(runner, request_id) + rows = helper.rows(runner) + assert [(row.request_id, row.temperature, row.top_k, row.top_p) for row in rows] == [ + ("b", 0.8, 15, 0.9), + ("a", 0.2, 7, 0.5), + ] + assert all(row.max_tokens == 16 for row in rows) + + +@pytest.mark.cpu +def test_mrv2_thinker_payload_uses_output_channel_contract(): + sampler = MiniCPMO45DuplexSampler(object(), SimpleNamespace()) + tokens = torch.tensor([[7]]) + standard = (SimpleNamespace(sampled_token_ids=tokens), torch.ones(1), torch.zeros(1)) + output = sampler.sample_step(None, None, None, None, lambda *_args: standard) + payload = {"latent": torch.ones(1, 2)} + assert output.sampler_output is standard[0] + assert output.include_hidden_states is False + assert output.finalize_multimodal(payload, [1]) is payload + + +@hardware_test(res={"cuda": "H100"}, num_cards=1) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("partial_prefill", [True, False]) +def test_mrv2_sampler_preserves_history_seed_counts_and_prefill_eligibility(mocker, partial_prefill): + device = "cuda" + params = SamplingParams(temperature=0.2, top_k=7, top_p=0.5, seed=13) + model = SimpleNamespace( + _mrv2_sampling_infos={ + "a": {"duplex": {"data_plane": True, "session_id": "session"}, "sampling_params": params} + }, + prepare_duplex_sampling=mocker.Mock(), + sample=mocker.Mock( + return_value=SimpleNamespace(sampled_token_ids=torch.tensor([[9]], device=device), logprobs_tensors=None) + ), + ) + # Request a occupies slot 2; neither slot number nor padded tokens are + # sampling row indices. The uncomputed tail must never enter its history. + tokens = torch.tensor([[0] * 8, [0] * 8, [1, 2, 3, 4, 6, 8, 77, 88]], device=device) + states = SimpleNamespace( + all_token_ids=SimpleNamespace(gpu=tokens), + prompt_len=SimpleNamespace(np=np.array([0, 0, 4])), + prefill_len=SimpleNamespace(gpu=torch.tensor([0, 0, 7 if partial_prefill else 4], device=device)), + ) + batch = SimpleNamespace( + req_ids=["a"], + num_reqs=1, + num_draft_tokens=0, + idx_mapping_np=np.array([2]), + idx_mapping=torch.tensor([2], device=device), + seq_lens=torch.tensor([6], device=device), + cu_num_logits=torch.tensor([0, 1], device=device), + is_prefilling_np=np.array([partial_prefill]), + num_computed_prefill_tokens_np=np.array([5]), + num_computed_tokens_np=np.array([5]), + num_scheduled_tokens=np.array([1]), + prefill_len_np=np.array([7 if partial_prefill else 4]), + ) + base_output = SimpleNamespace(num_sampled=torch.tensor([0], device=device)) + base = mocker.Mock(return_value=base_output) + base.req_states = states + sampler = MiniCPMO45DuplexSampler(base, model) + output = sampler(torch.zeros(1, 40, device=device), batch) + if partial_prefill: + assert output is base_output + model.sample.assert_not_called() + assert model.prepare_duplex_sampling.call_args.args[2] == () + assert sampler.generators == {} + else: + assert output.num_sampled.tolist() == [1] + assert output.num_rejected.tolist() == [0] + assert output.sampled_token_ids.tolist() == [[9]] + md = model.sample.call_args.args[1] + assert md.output_token_ids == [[6, 8]] + assert md.generators[0].initial_seed() == 13 + assert md.top_k.tolist() == [7] + generator = md.generators[0] + sampler(torch.zeros(1, 40, device=device), batch) + assert model.sample.call_args.args[1].generators[0] is generator + sampler.forget_requests(["a"]) + assert sampler.generators == {} + + +@pytest.mark.cpu +def test_thinker_history_copies_only_new_ids_and_resets_at_condition_change(): + reads = [] + values = torch.arange(24).reshape(3, 8) + + class Ledger: + def __getitem__(self, index): + reads.append((index[0], index[1].start, index[1].stop)) + return values[index] + + states = SimpleNamespace( + prompt_len=SimpleNamespace(np=np.array([0, 0, 2])), + all_token_ids=SimpleNamespace(gpu=Ledger()), + ) + model = SimpleNamespace(_mrv2_sampling_infos={}) + sampler = MiniCPMO45DuplexSampler(SimpleNamespace(req_states=states), model) + row = SimpleNamespace(row_idx=0, request_id="a", seq=1) + batch = SimpleNamespace( + num_reqs=1, + idx_mapping_np=np.array([2]), + num_computed_tokens_np=np.array([3]), + num_scheduled_tokens=np.array([1]), + ) + infos = {"a": {"sampling_params": SamplingParams()}} + first = sampler._metadata(batch, [row], infos, "cpu") + first.output_token_ids[0].append(999) # A consumer cannot mutate the cache. + batch.num_computed_tokens_np[0] = 4 + second = sampler._metadata(batch, [row], infos, "cpu") + assert second.output_token_ids == [[18, 19, 20]] + assert reads == [(2, 2, 4), (2, 4, 5)] + row.seq = 2 + sampler._metadata(batch, [row], infos, "cpu") + assert reads[-1] == (2, 2, 5) + batch.num_computed_tokens_np[0] = 2 + assert sampler._metadata(batch, [row], infos, "cpu").output_token_ids == [[18]] + sampler.forget_requests(["a"]) + assert sampler._histories == {} + + +@pytest.mark.cpu +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +def test_thinker_histories_and_rng_follow_requests_across_reorder_and_slot_reuse(device): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("CUDA required") + request_ids = [f"request-{slot}" for slot in range(4)] + tokens = torch.arange(32, device=device).reshape(4, 8) + states = SimpleNamespace( + prompt_len=SimpleNamespace(np=np.full(4, 2)), + all_token_ids=SimpleNamespace(gpu=tokens), + ) + infos = {req_id: {"sampling_params": SamplingParams(seed=slot)} for slot, req_id in enumerate(request_ids)} + sampler = MiniCPMO45DuplexSampler(SimpleNamespace(req_states=states), SimpleNamespace(_mrv2_sampling_infos=infos)) + reference_rng = { + req_id: torch.Generator(device=device).manual_seed(slot) for slot, req_id in enumerate(request_ids) + } + for step in range(4): + if step == 2: + sampler.forget_requests([request_ids[0]]) + assert request_ids[0] not in sampler.generators + assert request_ids[0] not in sampler._histories + request_ids[0] = "replacement" + tokens[0].add_(100) + infos["replacement"] = {"sampling_params": SamplingParams(seed=31)} + reference_rng["replacement"] = torch.Generator(device=device).manual_seed(31) + slots = np.roll(np.arange(4)[::-1], step) + # Each request has a different accepted length; a new condition can + # replace tokens without changing either the slot or the prompt length. + lengths = [3 + (int(slot) + step) % 4 for slot in slots] + if step == 3: + tokens.add_(1000) + rows = [ + SimpleNamespace(row_idx=row, request_id=request_ids[slot], seq=int(step == 3)) + for row, slot in enumerate(slots) + ] + batch = SimpleNamespace( + num_reqs=4, + idx_mapping_np=slots, + num_computed_tokens_np=np.array(lengths) - 1, + num_scheduled_tokens=np.ones(4, dtype=int), + ) + metadata = sampler._metadata(batch, rows, infos, device) + for row, slot in enumerate(slots): + assert metadata.output_token_ids[row] == tokens[slot, 2 : lengths[row]].tolist() + torch.testing.assert_close( + torch.rand(4, device=device, generator=metadata.generators[row]), + torch.rand(4, device=device, generator=reference_rng[request_ids[slot]]), + ) + sampler.forget_requests(request_ids) + assert sampler._histories == sampler.generators == infos == {} + + +@pytest.mark.cpu +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +def test_thinker_decode_reuses_deferred_samples_without_device_ledger_reads(device): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("CUDA required") + + class NoLedgerReads: + def __getitem__(self, index): + raise AssertionError("steady-state sampling must not read back the device token ledger") + + states = SimpleNamespace( + prompt_len=SimpleNamespace(np=np.array([2, 2])), + all_token_ids=SimpleNamespace(gpu=NoLedgerReads()), + ) + infos = {req_id: {"sampling_params": SamplingParams(seed=seed)} for seed, req_id in enumerate(("a", "b"))} + sampler = MiniCPMO45DuplexSampler(SimpleNamespace(req_states=states), SimpleNamespace(_mrv2_sampling_infos=infos)) + expected: dict[str, list[int]] = {"a": [], "b": []} + for step in range(8): + slots = [0, 1] if step % 2 == 0 else [1, 0] + rows = [SimpleNamespace(row_idx=i, request_id=("a", "b")[slot], seq=0) for i, slot in enumerate(slots)] + batch = SimpleNamespace( + num_reqs=2, + idx_mapping_np=np.array(slots), + num_computed_tokens_np=np.full(2, 1 + step), + num_scheduled_tokens=np.ones(2, dtype=int), + ) + metadata = sampler._metadata(batch, rows, infos, device) + assert metadata.output_token_ids == [expected[row.request_id] for row in rows] + sampled = [step * 10 + slot for slot in slots] + sampler._defer_history(rows, torch.tensor(sampled, device=device)) + for row, token in zip(rows, sampled): + expected[row.request_id].append(token) + # Finishing a request before the pending copy is consumed cannot restore + # its history or append that sample to another request reusing its slot. + sampler.forget_requests(["a"]) + sampler._commit_history() + assert "a" not in sampler._histories + assert sampler._histories["b"][1] == expected["b"] + + +@pytest.mark.cpu +def test_thinker_reuses_policy_snapshot_without_another_device_copy(mocker): + pending = SimpleNamespace(host=torch.tensor([[7, 3, 0], [9, 4, 1]]), event=mocker.Mock(), row_idxs=[0, 1]) + model = SimpleNamespace(_minicpmo45_duplex_pending_samples=pending) + sampler = MiniCPMO45DuplexSampler(SimpleNamespace(), model) + sampler._histories = {"a": ((0, 2, 0), []), "b": ((1, 2, 0), [])} + rows = [SimpleNamespace(request_id=req_id) for req_id in ("a", "b")] + # No device tensor is needed: the policy has already captured the result. + sampler._defer_history(rows, None) + sampler._commit_history() + assert sampler._histories["a"][1] == [7] + assert sampler._histories["b"][1] == [9] + pending.event.synchronize.assert_called_once() + + +@pytest.mark.cpu +@pytest.mark.parametrize("all_policy_rows", [True, False]) +def test_thinker_policy_reads_logits_in_place_when_every_row_is_duplex(mocker, all_policy_rows): + params = SamplingParams(temperature=0.0) + infos = {"a": {"duplex": {"data_plane": True, "session_id": "s"}, "sampling_params": params}} + infos["b"] = dict(infos["a"]) if all_policy_rows else {"sampling_params": params} + model = SimpleNamespace(_mrv2_sampling_infos=infos, prepare_duplex_sampling=mocker.Mock(), sample=mocker.Mock()) + model.sample.return_value = None # the stock sampler takes over + states = SimpleNamespace(prompt_len=SimpleNamespace(np=np.array([4, 4]))) + batch = SimpleNamespace( + req_ids=["a", "b"], + num_reqs=2, + num_draft_tokens=0, + idx_mapping_np=np.array([0, 1]), + is_prefilling_np=np.array([False, False]), + num_computed_tokens_np=np.array([3, 3]), + num_scheduled_tokens=np.array([1, 1]), + ) + base = mocker.Mock(return_value="stock") + base.req_states = states + sampler = MiniCPMO45DuplexSampler(base, model) + logits = torch.randn(2, 8) + assert sampler(logits, batch) == "stock" + policy_logits = model.prepare_duplex_sampling.call_args.args[0] + if all_policy_rows: + assert policy_logits is logits + else: + assert policy_logits is not logits + torch.testing.assert_close(policy_logits, logits[:1]) + assert base.call_args.args[0] is logits diff --git a/tests/model_executor/models/minicpmo_4_5/test_talker_batching.py b/tests/model_executor/models/minicpmo_4_5/test_talker_batching.py index 8be176852be..69bcb48ad5e 100644 --- a/tests/model_executor/models/minicpmo_4_5/test_talker_batching.py +++ b/tests/model_executor/models/minicpmo_4_5/test_talker_batching.py @@ -1150,7 +1150,8 @@ def test_native_duplex_rollover_matches_official_sliding_recompute(mocker) -> No assert build_condition.call_count == 3 -def test_native_duplex_condition_advance_without_rollover_updates_window_state(mocker) -> None: +@pytest.mark.parametrize("mrv2", [False, True]) +def test_native_duplex_condition_advance_without_rollover_updates_window_state(mocker, mrv2) -> None: talker = _make_talker() talker.emb_text = nn.Embedding(1, 2) talker.emb_code = nn.ModuleList([nn.Embedding(8, 2)]) @@ -1175,7 +1176,10 @@ def test_native_duplex_condition_advance_without_rollover_updates_window_state(m initial_state = talker._request_condition_states["req-history"] assert initial_state["condition_seq"] == 0 assert torch.equal(initial_state["condition"], first_condition) - talker._request_audio_states["req-history"]["recent_codes"] = [1, 2, 3] + if not mrv2: + talker._request_audio_states["req-history"]["recent_codes"] = [1, 2, 3] + else: + common["ids"] = {"streaming_prompt_previous_codes": [1, 2, 3]} _, embeds, _ = talker.preprocess( torch.zeros(4, dtype=torch.long), @@ -1204,6 +1208,7 @@ def test_native_duplex_condition_advance_without_rollover_updates_window_state(m assert torch.equal(retry_embeds, second_condition) assert talker._request_condition_states["req-history"]["base_recent_codes"] == (1, 2, 3) + common.pop("ids", None) previous_codes = [4, 5] expected_rollover = torch.cat( [ @@ -1431,3 +1436,47 @@ def test_sample_keeps_general_penalties_without_codec_history(mocker): ) assert captured["metadata"].no_penalties is False assert captured["metadata"].repetition_penalties.tolist() == [pytest.approx(1.05)] + + +@pytest.mark.parametrize("legacy_none", [False, True]) +@pytest.mark.parametrize("sessions", [2, 4, 8, 16]) +def test_mrv2_overlapping_native_conditions_use_distinct_request_ids(mocker, legacy_none, sessions): + talker = _make_talker() + talker.emb_text = nn.Embedding(1, 2) + mocker.patch.object(talker, "_build_condition_embeddings", return_value=torch.ones(2, 2)) + + def condition(req_id, seq, turn_start=False): + return talker.preprocess( + torch.zeros(2, dtype=torch.long), + None, + **({"request_id": None} if legacy_none else {}), + req_id=req_id, + native_duplex=True, + _omni_is_prefill=True, + _omni_prompt_len=2, + tts_token_ids=torch.tensor([1]), + tts_hidden_states=torch.ones(1, 2), + meta={"streaming_condition_seq": seq, "turn_start": turn_start}, + ) + + condition("session-a", 0, True) + for seq in (1, 2, 3): + condition("session-a", seq) + others = [f"session-{i}" for i in range(sessions - 1)] + for req_id in others: + condition(req_id, 0, True) + for seq in range(4, 12): + condition("session-a", seq) + for req_id in reversed(others): + condition(req_id, seq - 3) + assert talker._request_condition_states["session-a"]["condition_seq"] == 11 + assert all(talker._request_condition_states[req_id]["condition_seq"] == 8 for req_id in others) + assert set(talker._request_audio_states) == {"session-a", *others} + talker.on_requests_finished({others[0]}) + talker._flush_deferred_cleanup() + assert others[0] not in talker._request_condition_states + assert others[0] not in talker._request_audio_states + condition("new-session", 0, True) + condition("session-a", 12) + assert talker._request_condition_states["new-session"]["condition_seq"] == 0 + assert talker._request_condition_states["session-a"]["condition_seq"] == 12 diff --git a/tests/model_executor/models/minicpmo_4_5/test_talker_mrv2.py b/tests/model_executor/models/minicpmo_4_5/test_talker_mrv2.py index 950975c6444..689a8a05c7e 100644 --- a/tests/model_executor/models/minicpmo_4_5/test_talker_mrv2.py +++ b/tests/model_executor/models/minicpmo_4_5/test_talker_mrv2.py @@ -140,6 +140,97 @@ def test_sampler_adapter_keeps_upstream_counts_and_only_forces_codec_eos(mocker, talker.take_mrv2_forced_eos.assert_called_once_with(batch, base.req_states, 2) +@pytest.mark.parametrize("condition_seq", [9, None]) +def test_mrv2_talker_native_duplex_output(condition_seq) -> None: + talker = _talker() + rows = [ + dict(slot=0, prompt_len=4, computed=6, span=[17], prefill=False), + ] + batch, padded = _batch(rows, pad_to=1) + buffers = [ + { + "native_duplex": True, + "duplex": {"epoch": 2, "turn_id": 5}, + "meta": {"native_duplex_segment_text": "hello", "streaming_condition_seq": condition_seq}, + } + ] + hidden = torch.zeros((1, 4)) + out = talker.make_omni_output_mrv2( + hidden, + input_batch=batch, + req_states=_req_states({0: 4}), + model_intermediate_buffer=buffers, + ) + assert isinstance(out, OmniOutput) + meta = out.multimodal_outputs["meta"] + assert meta["native_duplex"][0].item() is True + assert meta["duplex_epoch"][0].item() == 2 + assert meta["duplex_turn_id"][0].item() == 5 + assert meta["streaming_condition_seq"][0].item() == (9 if condition_seq is not None else -1) + assert bytes(meta["llm_output_text_utf8"][0].tolist()).decode("utf-8") == "hello" + assert bytes(meta["native_duplex_segment_text"][0].tolist()).decode("utf-8") == "hello" + + +@pytest.mark.parametrize("fence", [{"epoch": -1, "turn_id": 5}, {"epoch": True, "turn_id": 5}, {"turn_id": 5}]) +def test_mrv2_talker_rejects_native_duplex_without_fence_identity(fence) -> None: + """Same contract as V1 make_omni_output: condition fencing needs non-negative int identities.""" + batch, _ = _batch([dict(slot=0, prompt_len=4, computed=6, span=[17], prefill=False)], pad_to=1) + with pytest.raises(RuntimeError, match="requires non-negative integer epoch and turn_id"): + _talker().make_omni_output_mrv2( + torch.zeros((1, 4)), + input_batch=batch, + req_states=_req_states({0: 4}), + model_intermediate_buffer=[{"native_duplex": True, "duplex": fence}], + ) + + +def test_mrv2_context_with_only_eos_slot_forces_eos_at_prefill(): + talker = _talker(max_position_embeddings=5) + batch, _ = _batch([dict(slot=2, prompt_len=4, computed=0, span=[0] * 4, prefill=True)]) + talker.make_omni_output_mrv2( + torch.zeros((4, 4)), + input_batch=batch, + req_states=_req_states({2: 4}), + model_intermediate_buffer=[{"audio_state": {"finished": False}}], + ) + assert talker.take_mrv2_forced_eos(batch, None, 1).tolist() == [True] + + +@pytest.mark.parametrize( + "step,turn_start,turn_end,masked", + [ + (step, False, True, masked) + for step, masked in [(24, False), (25, True), (29, True), (30, False), (50, True), (55, False)] + ] + + [(0, True, False, False), (0, False, False, True), (24, False, False, True), (25, False, False, False)], +) +def test_mrv2_turn_end_drain_masks_cadence_eos(mocker, step, turn_start, turn_end, masked): + from vllm_omni.model_executor.models.minicpmo_4_5.minicpmo_4_5_omni_tts import MiniCPMO45TalkerSampler + + talker = _talker() + batch, _ = _batch([dict(slot=2, prompt_len=4, computed=3 + step, span=[17], prefill=False)]) + states = _req_states({2: 4}) + talker.make_omni_output_mrv2( + torch.zeros((1, 4)), + input_batch=batch, + req_states=states, + model_intermediate_buffer=[ + { + "native_duplex": True, + "duplex": {"epoch": 3, "turn_id": 7}, + "meta": {"turn_start": turn_start, "turn_end": turn_end}, + } + ], + ) + base = mocker.Mock(side_effect=lambda logits, _: SimpleNamespace(sampled_token_ids=logits.argmax(-1)[:, None])) + base.req_states = states + sampler = MiniCPMO45TalkerSampler(base, talker) + logits = torch.zeros(1, _EOS + 1) + logits[0, _EOS] = 10 + output = sampler(logits, batch) + assert (output.sampled_token_ids.item() != _EOS) is masked + + @pytest.mark.cuda def test_mrv2_sampler_applies_codec_window_penalty_instead_of_stock_penalty(): """The real MRv2 sampler pipeline scores the V1 16-frame codec penalty. @@ -200,3 +291,70 @@ def test_mrv2_sampler_applies_codec_window_penalty_instead_of_stock_penalty(): assert processed[0, 7].item() == pytest.approx(logits[0, 7].item() * penalty**3) # The runner still finds the output bin counts on ``penalties_state``. assert sampler.penalties_state.output_bin_counts is not None + + +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +@pytest.mark.parametrize("output_count", [0, 3, 16, 20]) +def test_codec_penalty_carries_previous_segment_through_reordered_slots(device, output_count): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("CUDA required") + from vllm_omni.model_executor.models.minicpmo_4_5.minicpmo_4_5_omni_tts import ( + _apply_batched_repetition_penalty, + _apply_codec_window_penalty_gpu, + ) + + prompt_lens = torch.tensor([3, 0, 2], device=device) + lengths = prompt_lens + output_count + tokens = torch.zeros(3, 64, dtype=torch.long, device=device) + prefixes = torch.full((3, 16), -1, dtype=torch.long, device=device) + histories = [] + for slot, prefix in [(2, [5, 5, 5, 9]), (0, [7, 7])]: + output = ([1, 7, 2] * 7)[:output_count] + start = int(prompt_lens[slot]) + tokens[slot, start : start + output_count] = torch.tensor(output, device=device) + prefixes[slot, -len(prefix) :] = torch.tensor(prefix, device=device) + histories.append(torch.tensor(prefix + output)) + logits = torch.linspace(-3, 3, 32, device=device).repeat(2, 1) + expected = _apply_batched_repetition_penalty(logits.cpu(), histories, penalty=1.05, window_size=16) + _apply_codec_window_penalty_gpu( + logits, + torch.tensor([2, 0], device=device), + tokens, + lengths, + prompt_lens, + torch.full((3,), 1.05, device=device), + window_size=16, + history_prefix=prefixes, + ) + torch.testing.assert_close(logits.cpu(), expected) + + +def test_mrv2_prefill_seeds_codec_history_for_the_actual_slot(mocker): + talker = _talker() + talker._mrv2_penalty_state = mocker.Mock() + batch, _ = _batch([dict(slot=3, prompt_len=2, computed=0, span=[0, 0], prefill=True)]) + talker.make_omni_output_mrv2( + torch.zeros(2, 4), + input_batch=batch, + req_states=_req_states({3: 2}), + model_intermediate_buffer=[{"audio_state": {"recent_codes": [5, 7, 5]}}], + ) + talker._mrv2_penalty_state.set_history_prefix.assert_called_once_with([3], [[5, 7, 5]]) + + +def test_turn_decode_does_not_build_duplex_limit_tensors(monkeypatch): + import vllm_omni.model_executor.models.minicpmo_4_5.minicpmo_4_5_omni_tts as module + + def reject_copy(*args, **kwargs): + raise AssertionError("ordinary turn decode must not upload duplex limits") + + monkeypatch.setattr(module, "index_to_device", reject_copy) + talker = _talker() + batch, _ = _batch([dict(slot=0, prompt_len=4, computed=5, span=[1], prefill=False)]) + talker.make_omni_output_mrv2( + torch.zeros(1, 4), + input_batch=batch, + req_states=_req_states({0: 4}), + model_intermediate_buffer=[{}], + ) + assert talker._mrv2_masked_eos is None diff --git a/tests/model_executor/models/minicpmo_4_5/test_thinker_mrv2.py b/tests/model_executor/models/minicpmo_4_5/test_thinker_mrv2.py index d34bf614e9a..396f8017a30 100644 --- a/tests/model_executor/models/minicpmo_4_5/test_thinker_mrv2.py +++ b/tests/model_executor/models/minicpmo_4_5/test_thinker_mrv2.py @@ -4,6 +4,7 @@ from types import SimpleNamespace +import numpy as np import pytest import torch @@ -14,13 +15,13 @@ @pytest.mark.parametrize( - "profile,capacities,kv_gib", + "profile,capacities,kv_gib,runners", [ - ("minicpmo_4_5_turn_mrv2.yaml", [16, 8, 8], 2), - ("minicpmo_4_5_turn_mrv2_h200.yaml", [16, 16, 8], 4), + ("minicpmo_4_5_turn_mrv2.yaml", [16, 8, 8], 2, [True, True, True]), + ("minicpmo_4_5_turn_mrv2_h200.yaml", [16, 16, 8], 4, [True, True, True]), ], ) -def test_mrv2_profile_retains_full_thinker_handoff(profile, capacities, kv_gib, monkeypatch): +def test_mrv2_profile_retains_full_thinker_handoff(profile, capacities, kv_gib, runners, monkeypatch): from pathlib import Path from vllm_omni.config.stage_config import _apply_platform_overrides, load_deploy_config, merge_pipeline_deploy @@ -33,7 +34,7 @@ def test_mrv2_profile_retains_full_thinker_handoff(profile, capacities, kv_gib, deploy = Path(__file__).resolve().parents[4] / "vllm_omni/deploy" / profile config = _apply_platform_overrides(load_deploy_config(deploy), platform="cuda") stages = merge_pipeline_deploy(MINICPMO_4_5_PIPELINE, config) - assert [s.yaml_engine_args["use_v2_model_runner"] for s in stages] == [True, True, True] + assert [s.yaml_engine_args["use_v2_model_runner"] for s in stages] == runners assert [s.yaml_engine_args["async_chunk"] for s in stages] == [False, True, True] assert [s.yaml_engine_args["max_num_seqs"] for s in stages] == capacities assert stages[1].yaml_engine_args["kv_cache_memory_bytes"] == kv_gib * 1024**3 @@ -81,6 +82,7 @@ def test_row_ledger_uses_live_batch_after_replay(mocker): model = _model(mocker) # Chunked prefill (two tokens), decode (one token), and graph padding. batch = SimpleNamespace( + req_ids=["a", "b"], input_ids=torch.tensor([11, 12, 41, 0]), positions=torch.tensor([6, 7, 20, 0]), ) @@ -98,3 +100,148 @@ def test_row_ledger_uses_live_batch_after_replay(mocker): assert out.multimodal_outputs["latent"] is hidden torch.testing.assert_close(out.multimodal_outputs["latent_input_ids"], batch.input_ids[:, None]) torch.testing.assert_close(out.multimodal_outputs["latent_positions"], batch.positions[:, None]) + + +def test_mrv2_thinker_duplex_output_and_prompt_rows(mocker): + model = _model(mocker, session="duplex") + assert model.has_preprocess is True + # Row "a" completes its append prefill on this step; row "b" decodes. + batch = SimpleNamespace( + req_ids=["a", "b"], + num_reqs=2, + input_ids=torch.tensor([101, 102]), + positions=torch.tensor([0, 1]), + is_prefilling_np=np.array([True, False]), + num_computed_prefill_tokens_np=np.array([2, 0]), + num_scheduled_tokens=np.array([1, 1]), + prefill_len_np=np.array([3, 5]), + ) + hidden = torch.randn(2, 8) + buffers = [ + { + "duplex": { + "duplex_prompt_token_ids": [1, 2, 3], + "special_token_ids": {"tts_bos_token_id": 151703, "turn_start_token_id": 151644}, + } + }, + {}, + ] + out = model.make_omni_output_mrv2( + hidden, + input_batch=batch, + req_states=None, + model_intermediate_buffer=buffers, + ) + assert out.text_hidden_states is hidden + assert out.multimodal_outputs["latent"] is hidden + assert out.multimodal_outputs["duplex_prompt_token_ids"] == [[1, 2, 3], None] + assert "tts_bos_token_id" in out.multimodal_outputs["meta"] + assert out.multimodal_outputs["meta"]["tts_bos_token_id"][0].item() == 151703 + # Tokenizer constants fill rows that carry none (e.g. an append that built + # no unit); a None entry would stick in that request's accumulated meta. + assert out.multimodal_outputs["meta"]["tts_bos_token_id"][1].item() == 151703 + # Host-only metadata: a device tensor would cost a blocking H2D per row and key. + assert out.multimodal_outputs["meta"]["tts_bos_token_id"][0].device.type == "cpu" + + # Pure decode steps leave the accumulated prompt/meta snapshot untouched. + batch.is_prefilling_np = np.array([False, False]) + out = model.make_omni_output_mrv2( + hidden, + input_batch=batch, + req_states=None, + model_intermediate_buffer=buffers, + ) + assert "duplex_prompt_token_ids" not in out.multimodal_outputs + assert "meta" not in out.multimodal_outputs + assert out.multimodal_outputs["latent"] is hidden + + +def test_mrv2_turn_thinker_retains_native_multimodal_path(mocker): + assert _model(mocker, session="turn").has_preprocess is False + + +def test_mrv2_thinker_custom_sampler_and_lifecycle(mocker): + from vllm_omni.model_executor.models.minicpmo_4_5.duplex.mrv2_sampling import ( + MiniCPMO45DuplexSampler, + ) + from vllm_omni.worker_v2.model_states.omni_model_state import OmniModelState + + model = _model(mocker, session="duplex") + base_sampler = mocker.MagicMock() + sampler, _ = model.mrv2_custom_sampler(base_sampler) + assert isinstance(sampler, MiniCPMO45DuplexSampler) + assert model._mrv2_duplex_sampler is sampler + + state_mock = mocker.MagicMock(spec=OmniModelState) + state_mock.model = model + resolved = OmniModelState.custom_sampler(state_mock, base_sampler) + assert resolved is not None + assert isinstance(resolved[0], MiniCPMO45DuplexSampler) + + model.preprocess( + torch.tensor([1]), + torch.randn(1, 4), + req_id="req-1", + duplex={"data_plane": True, "session_id": "s1", "epoch": 0, "seq": 1}, + ) + assert "req-1" in model._mrv2_sampling_infos + + model.on_requests_finished({"req-1"}) + assert "req-1" not in model._mrv2_sampling_infos + + +def test_mrv2_cancel_before_first_prefill_completes_frees_session(mocker): + """A request cancelled mid-prefill never reaches the sampler, so preprocessing owns its session.""" + model = _model(mocker, session="duplex") + helper = SimpleNamespace( + sessions={}, + take_staged_prefill=lambda *_: None, + _decode_audio_payload=lambda _: None, + frame_kwargs=mocker.MagicMock(side_effect=ValueError("stop after the session exists")), + ) + model._minicpmo45_duplex_data_plane_helper = helper + mocker.patch.object( + model, + "_minicpmo45_duplex_session_state", + side_effect=lambda h, sid, _: h.sessions.setdefault(sid, object()), + ) + model.preprocess( + torch.tensor([1]), + torch.randn(1, 4), + req_id="req-1", + duplex={"data_plane": True, "session_id": "s1", "epoch": 0, "seq": 1, "payload": {}}, + ) + assert "s1" in helper.sessions + + model.on_requests_finished({"req-1"}) + assert helper.sessions == {} + + +def test_mrv2_duplex_thinker_batches_prefill_appends(mocker): + """MRv2 hands its prefill rows to the V1 cross-session ``preprocess_batch``.""" + model = _model(mocker, session="duplex") + batch = mocker.patch.object(model, "preprocess_batch") + infos = [ + {"req_id": "a", "duplex": {"data_plane": True, "session_id": "s1"}}, + {"req_id": "b", "duplex": {"data_plane": True, "session_id": "s2"}}, + {"req_id": "c"}, # not a duplex row + {}, + ] + model.preprocess_batch_mrv2(req_infos=infos, device=torch.device("cpu")) + kwargs = batch.call_args.kwargs + assert kwargs["req_ids"] == ["a", "b"] + assert kwargs["model_intermediate_buffer"] == {"a": infos[0], "b": infos[1]} + batch.reset_mock() + model.preprocess_batch_mrv2(req_infos=[{"req_id": "c"}], device=torch.device("cpu")) + batch.assert_not_called() + + +@pytest.mark.parametrize("v2,session,keeps", [(True, "duplex", True), (True, "turn", False), (False, "duplex", False)]) +def test_only_the_mrv2_duplex_thinker_keeps_native_multimodal_inputs(mocker, v2, session, keeps): + # Chat requests with media share the duplex Thinker; turn MRv2 has no preprocess and V1 encodes first. + assert _model(mocker, v2=v2, session=session).preprocess_keeps_mm_inputs is keeps + + +@pytest.mark.parametrize("v2,session", [(True, "turn"), (False, "duplex")]) +def test_mrv2_batch_preprocess_hook_is_duplex_mrv2_only(mocker, v2, session): + assert not hasattr(_model(mocker, v2=v2, session=session), "preprocess_batch_mrv2") diff --git a/tests/model_executor/stage_input_processors/test_minicpmo_4_5_async_chunk.py b/tests/model_executor/stage_input_processors/test_minicpmo_4_5_async_chunk.py index 4a4f72a004d..13036689ba1 100644 --- a/tests/model_executor/stage_input_processors/test_minicpmo_4_5_async_chunk.py +++ b/tests/model_executor/stage_input_processors/test_minicpmo_4_5_async_chunk.py @@ -194,6 +194,49 @@ def test_duplex_turn_end_waits_for_terminal_codec_flush() -> None: assert final.meta.turn_end is True +@pytest.mark.parametrize( + "frame_count,delta_frames,body_chunks", + [ + # Single-frame deltas across the 25-frame chunk boundary. + (1, 1, 0), + (24, 1, 0), + (25, 1, 1), + (26, 1, 1), + (60, 1, 2), + # Batched deltas: two full chunks, then a short tail held for the final flush. + (52, 25, 2), + ], +) +def test_duplex_final_segment_preserves_every_code_and_closes_once( + frame_count: int, delta_frames: int, body_chunks: int +) -> None: + manager = _manager() + request = _request("final-segment") + emitted = [] + bodies = 0 + for start in range(0, frame_count, delta_frames): + codes = range(start, min(start + delta_frames, frame_count)) + payload = tts2code2wav_async_chunk(manager, _duplex_delta(*codes, turn_end=True), request, False) + if payload is not None: + assert payload.meta.last_chunk is False + assert payload.meta.turn_end is False + bodies += 1 + emitted.extend(_codes(payload)[payload.meta.codec_left_context_frames :]) + assert bodies == body_chunks + + final = tts2code2wav_async_chunk(manager, _duplex_delta(turn_end=True), request, True) + assert final is not None + assert final.meta.last_chunk is True + assert final.meta.turn_end is True + emitted.extend(_codes(final)[final.meta.codec_left_context_frames :]) + assert emitted == list(range(frame_count)) + + duplicate = tts2code2wav_async_chunk(manager, _duplex_delta(turn_end=True), request, True) + assert duplicate is not None + assert duplicate.codes is None + assert not duplicate.meta.last_chunk + + def test_first_chunk_forwards_reference_voice_and_duplex_identity() -> None: manager = _manager() request = _request("req") @@ -554,6 +597,139 @@ def test_cancel_drops_epoch_state_and_stale_request_cannot_publish() -> None: assert _codes(payload) == [4218, 4218, 4218, *range(25)] +@pytest.mark.parametrize("aborted", [False, True]) +def test_mrv2_abort_terminal_drops_pending_codec_frames(aborted: bool) -> None: + """The MRv2 snapshot has no scheduler status; its abort mark must behave like V1's FINISHED_ABORTED.""" + from vllm_omni.worker_v2.omni_data_plane import _NativeRequestState + + manager = _manager() + state = _NativeRequestState(request_id="native", external_req_id="native", prompt_token_ids=[0] * 3) + pending = state.snapshot(include_token_history=True, sampled_token_ids=[]) + assert tts2code2wav_async_chunk(manager, _duplex_delta(*range(10)), pending) is None + state.finished = True + state.aborted = aborted + terminal = tts2code2wav_async_chunk(manager, None, state.snapshot(include_token_history=True)) + if aborted: + assert terminal is None + assert "native" not in manager.code_prompt_token_ids + else: + assert terminal is not None + assert _codes(terminal)[3:] == list(range(10)) + + +@pytest.mark.parametrize("turn_end", [False, True]) +def test_mrv2_sampled_codec_eos_flushes_resumable_segment(turn_end: bool) -> None: + from vllm_omni.worker_v2.omni_data_plane import _NativeRequestState + + state = _NativeRequestState( + request_id="native", + external_req_id="native", + prompt_token_ids=[0] * 3, + resumable=True, + sampling_params=SimpleNamespace(stop_token_ids=[6561]), + ) + state.accept_tokens([6561]) + request = state.snapshot(include_token_history=True, sampled_token_ids=[6561]) + assert not request.is_finished() # The session remains available for another unit. + payload = tts2code2wav_async_chunk(_manager(), _duplex_delta(*range(7), turn_end=turn_end), request) + assert payload is not None + assert _codes(payload) == [4218] * 3 + list(range(7)) + assert payload.meta.tts_is_last_chunk is True + assert payload.meta.last_chunk is turn_end + + +def test_old_sampled_eos_cannot_close_the_next_turn(): + from vllm_omni.worker_v2.omni_data_plane import _NativeRequestState + + manager = _manager() + state = _NativeRequestState( + request_id="native", + external_req_id="native", + prompt_token_ids=[0, 0], + resumable=True, + sampling_params=SimpleNamespace(stop_token_ids=[6561]), + ) + state.accept_tokens([6561]) + # The next payload has no newly accepted token, even though the ledger + # still ends in the previous segment's EOS. + stale = state.snapshot(include_token_history=True) + first = tts2code2wav_async_chunk(manager, _duplex_delta(turn_id=8, turn_end=True), stale) + assert first is None + state.accept_tokens([1]) + body = tts2code2wav_async_chunk( + manager, _duplex_delta(*range(30), turn_id=8, turn_end=True), state.snapshot(include_token_history=True) + ) + assert body is not None and body.meta.last_chunk is False + tail = tts2code2wav_async_chunk( + manager, _duplex_delta(turn_id=8, turn_end=True), state.snapshot(include_token_history=True), True + ) + assert tail is not None and tail.meta.last_chunk is True + assert _codes(body)[3:] + _codes(tail)[3:] == list(range(30)) + + +@pytest.mark.parametrize("lookahead", [1, 2, 4]) +def test_mrv2_async_lookahead_cannot_reopen_a_completed_condition(lookahead): + manager = _manager() + requests = [_request("a"), _request("b")] + for request in requests: + request.sampling_params = SimpleNamespace(stop_token_ids=[6561]) + + def output(request, seq, codes, *, eos=False, turn_end=False): + payload = _duplex_delta(*codes, text=f"condition-{seq}", turn_end=turn_end) + payload["meta"]["streaming_condition_seq"] = torch.tensor(seq) + request.sampled_token_ids = [6561] if eos else [] + return tts2code2wav_async_chunk(manager, payload, request, False) + + for request in requests: + first = output(request, 0, range(25), eos=True) + assert first is not None and first.meta.tts_is_last_chunk + assert _codes(first)[3:] == list(range(25)) + for request in reversed(requests): + for _ in range(lookahead): + assert output(request, 0, [999], eos=True) is None + assert output(request, 1, range(25, 35)) is None + # Even after a new condition starts, an old snapshot must not close it + # or contaminate its queued codes/text. + assert output(request, 0, [999], eos=True) is None + last = output(request, 1, range(35, 50), eos=True, turn_end=True) + assert last is not None and last.meta.last_chunk + assert _codes(last)[3:] == list(range(25, 50)) + assert last.meta.chunk_seq == 1 + assert output(request, 1, [999], eos=True, turn_end=True) is None + + +@pytest.mark.parametrize("request_terminal", [False, True]) +def test_mrv2_duplex_turn_end_keeps_live_code2wav_stream_open(request_terminal: bool) -> None: + # A turn end must not close the Code2Wav stream of a live resumable request, + # or the next turn's chunks are never received. + from vllm_omni.distributed.omni_connectors.model_runner.omni_connector_payload_transport import ( + _OmniConnectorPayloadTransportMixin as OmniConnectorPayloadTransport, + ) + from vllm_omni.worker_v2.omni_data_plane import _NativeRequestState + + state = _NativeRequestState( + request_id="native", + external_req_id="native", + prompt_token_ids=[0] * 3, + resumable=True, + sampling_params=SimpleNamespace(stop_token_ids=[6561]), + ) + state.accept_tokens([6561]) + state.finished = request_terminal + request = state.snapshot(include_token_history=True, sampled_token_ids=None if request_terminal else [6561]) + payload = tts2code2wav_async_chunk( + _manager(), _duplex_delta(*range(7), turn_end=True), request, request.is_finished() + ) + + assert payload is not None + assert payload.meta.last_chunk is True + assert payload.meta.turn_end is True + assert payload.meta.is_segment_finished.item() is True + assert payload.meta.finished.item() is request_terminal + metadata = OmniConnectorPayloadTransport._extract_scheduling_metadata({"meta": {"finished": payload.meta.finished}}) + assert metadata.get("input_terminal", False) is request_terminal + + @pytest.mark.parametrize("last_valid", [False, True]) def test_full_payload_accumulates_codec_validity_per_frame(last_valid): """A final invalid/EOS row must not invalidate the whole utterance.""" diff --git a/tests/model_executor/stage_input_processors/test_minicpmo_4_5_omni.py b/tests/model_executor/stage_input_processors/test_minicpmo_4_5_omni.py index 472e7121cb3..4512cfdcf42 100644 --- a/tests/model_executor/stage_input_processors/test_minicpmo_4_5_omni.py +++ b/tests/model_executor/stage_input_processors/test_minicpmo_4_5_omni.py @@ -102,7 +102,8 @@ def test_llm2tts_carries_request_ref_audio() -> None: assert info["ids"]["tts"] == [11, 12] -def test_native_duplex_speak_segment_reaches_split_talker() -> None: +@pytest.mark.parametrize("metadata_layout", ["scalar", "list", "tensor", "flat_tensor"]) +def test_native_duplex_speak_segment_reaches_split_talker(metadata_layout: str) -> None: prompt_ids = [101, 102] output_ids = [9304, 21, 22, 9308] latent = torch.arange(24, dtype=torch.float32).reshape(6, 4) @@ -123,6 +124,17 @@ def test_native_duplex_speak_segment_reaches_split_talker() -> None: }, }, ) + # Per-step boundary snapshots are accumulated as lists or concatenated + # tensors by the output processor before the native duplex handoff. + meta = source.outputs[0].multimodal_output["meta"] + if metadata_layout == "list": + meta.update({key: [value, value] for key, value in meta.items()}) + elif metadata_layout in ("tensor", "flat_tensor"): + meta.update({key: torch.tensor([[value], [value]]) for key, value in meta.items()}) + if metadata_layout == "flat_tensor": + mm_output = source.outputs[0].multimodal_output + mm_output.pop("meta") + mm_output.update({f"meta.{key}": value for key, value in meta.items()}) context = SimpleNamespace( bridge_states={ "duplex": { diff --git a/tests/worker_v2/test_omni_ar_model_runner.py b/tests/worker_v2/test_omni_ar_model_runner.py index d3f1b9eb416..cfe95a4411c 100644 --- a/tests/worker_v2/test_omni_ar_model_runner.py +++ b/tests/worker_v2/test_omni_ar_model_runner.py @@ -392,6 +392,45 @@ def test_async_output_slices_request_payloads_with_graph_padding( torch.testing.assert_close(payload["hidden"], hidden[offsets[i] : offsets[i + 1]]) +def test_finalized_full_payload_latent_ledger_is_accumulated_once(monkeypatch): + from tests.engine.test_output_processor_mrv2_text import ( + _Detokenizer, + _make_processor, + _rehydrated_output, + ) + from vllm_omni.outputs.output_modality import OutputModality + + monkeypatch.setattr(torch.cuda, "set_stream", lambda _stream: None) + ids = torch.tensor([[11], [12], [13]]) + positions = torch.arange(3).reshape(-1, 1) + hidden = torch.ones(3, 4) + batch = SimpleNamespace( + query_start_loc_np=np.array([0, 3]), + num_scheduled_tokens=np.array([3]), + num_reqs=1, + num_tokens_after_padding=3, + ) + # The MiniCPM-o duplex Thinker sampler publishes through finalize_multimodal + # instead of the full-payload fallback that mirrors rows into both channels. + output = _async_output( + req_ids=["r"], + multimodal_outputs={"latent": hidden, "latent_input_ids": ids, "latent_positions": positions}, + input_batch=batch, + async_chunk=False, + finalize_multimodal=lambda payload, _num_sampled: payload, + ).get_output() + assert output.multimodal_outputs is None + engine_output = _rehydrated_output(pooling_output=output.pooler_output[0]) + if output.multimodal_outputs: + engine_output.multimodal_output = output.multimodal_outputs[0] + processor, _ = _make_processor(OutputModality.LATENT, _Detokenizer()) + processed = processor.process_outputs([engine_output]) + payload = processed.request_outputs[0].outputs[0].multimodal_output + torch.testing.assert_close(payload["latent_input_ids"], ids) + torch.testing.assert_close(payload["latent_positions"], positions) + torch.testing.assert_close(payload["latent"], hidden) + + def test_async_chunk_output_stages_mm_on_copy_stream_before_get_output(monkeypatch) -> None: monkeypatch.setattr(torch.cuda, "set_stream", lambda _stream: None) calls = [] diff --git a/tests/worker_v2/test_omni_data_plane.py b/tests/worker_v2/test_omni_data_plane.py index 62898a456e2..5723dd3c329 100644 --- a/tests/worker_v2/test_omni_data_plane.py +++ b/tests/worker_v2/test_omni_data_plane.py @@ -105,6 +105,7 @@ def test_full_payload_waits_for_terminal_and_last_deferred_frame(plane): [(snapshot, payload)] = plane.record.batches[0] assert snapshot.is_finished() + assert snapshot.aborted is False assert snapshot.output_token_ids == [21, 2150] torch.testing.assert_close(payload["codes.audio"], torch.cat([first, last])) torch.testing.assert_close(payload["codes.ref"], ref) @@ -122,6 +123,7 @@ def test_full_payload_abort_discards_partial_and_late_outputs(plane): assert plane.abort_requests({"internal"}) == 1 [(snapshot, payload)] = plane.record.batches[0] assert snapshot.is_finished() and payload is None + assert snapshot.aborted is True assert _complete(plane, [{"codes.audio": torch.tensor([[3, 4]])}]) == 0 assert len(plane.record.batches) == 1 assert not plane._pending_full_payload_send @@ -675,3 +677,59 @@ def init_connectors(self, *, model_config): assert plane._get_model_config() is model_config assert _get_accept_hidden_layer_index(plane) == 24 + + +@pytest.mark.parametrize("factory", ["from_base", "from_request"]) +def test_native_request_preserves_streaming_flag_and_reference_audio(plane, factory): + from vllm.v1.core.sched.output import NewRequestData + + from vllm_omni.core.sched.output import OmniNewRequestData + from vllm_omni.model_executor.stage_input_processors.minicpmo_4_5_omni import tts2code2wav_async_chunk + + base = NewRequestData( + req_id="internal", + prompt_token_ids=[0, 0], + mm_features=[], + sampling_params=SamplingParams(stop_token_ids=[6561]), + pooling_params=None, + block_ids=([],), + num_computed_tokens=0, + lora_request=None, + ) + owner = SimpleNamespace( + **vars(base), + request_id="internal", + external_req_id="voice", + resumable=True, + model_intermediate_buffer={"codes": {"ref": [0.1, -0.1]}, "meta": {"ref_audio_sr": 16000}}, + ) + data = ( + OmniNewRequestData.from_base(base, owner) + if factory == "from_base" + else OmniNewRequestData.from_request(owner, ([],), prefill_token_ids=[0, 0]) + ) + plane.register_request(data) + state = plane._native_requests["internal"] + assert state.resumable is True + state.accept_tokens([6561]) + snapshot = state.snapshot(include_token_history=True, sampled_token_ids=[6561]) + payload = tts2code2wav_async_chunk( + SimpleNamespace(), + {"codes": {"audio": torch.arange(7)}, "meta": {"native_duplex": True, "turn_end": True}}, + snapshot, + ) + assert payload is not None and payload.meta.last_chunk is True + assert payload.meta.ref_audio_sr == 16000 + torch.testing.assert_close(payload.codes.ref, torch.tensor([0.1, -0.1])) + + +def test_native_snapshot_distinguishes_fresh_samples_from_ledger_tail(plane): + plane.register_request(_new_request()) + _complete(plane, [{"codes.audio": torch.tensor([[1]])}], token=2150) + first = plane.record.batches[-1][0][0] + assert first.sampled_token_ids == [2150] + plane.complete_outputs(req_ids=["internal"], inter_stage_outputs=[{"meta.turn_end": True}], sampled_token_ids=[[]]) + next_snapshot = plane.record.batches[-1][0][0] + assert next_snapshot.last_output_token_id == 2150 + assert next_snapshot.sampled_token_ids == [] + assert first.sampled_token_ids == [2150] diff --git a/tests/worker_v2/test_omni_gpu_model_runner.py b/tests/worker_v2/test_omni_gpu_model_runner.py index 803d7581a81..a0943c6a611 100644 --- a/tests/worker_v2/test_omni_gpu_model_runner.py +++ b/tests/worker_v2/test_omni_gpu_model_runner.py @@ -95,6 +95,28 @@ def test_full_payload_receive_is_polled_without_scheduled_tokens(mocker): plane.recv_full_payload_inputs.assert_called_once_with(scheduler_output) +def test_duplex_kv_reanchor_runs_after_block_table_writes(): + # Same model hook as the V1 AR runner: reanchor reads the committed block tables. + runner = object.__new__(OmniGPUModelRunner) + order = [] + for name in ("_prepare_native_data_plane", "finish_requests", "free_states", "add_requests", "update_requests"): + setattr(runner, name, lambda *_args: None) + runner._sync_native_data_plane_payloads = lambda *_args: None + runner.block_tables = SimpleNamespace(apply_staged_writes=lambda: order.append("block_tables")) + runner.model = SimpleNamespace( + apply_duplex_kv_reanchor=lambda r, scheduler_output: order.append(("reanchor", r, scheduler_output)) + ) + runner.aux_output_connector = None + output = object() + runner.kv_connector = SimpleNamespace(no_forward=lambda _s: output) + runner._merge_ec_connector_no_forward = lambda _s, value: value + runner._attach_native_data_plane_signals = lambda value: value + scheduler_output = SchedulerOutput.make_empty() + + assert runner.execute_model(scheduler_output) is output + assert order == ["block_tables", ("reanchor", runner, scheduler_output)] + + @pytest.mark.parametrize("output_form", ["tuple", "omni"]) @pytest.mark.parametrize("profile_only", [False, True]) def test_capture_model_unwraps_exclude_full_and_capture_mtp(output_form, profile_only): diff --git a/tests/worker_v2/test_omni_model_state.py b/tests/worker_v2/test_omni_model_state.py index 80333f2aa33..d6e3ac1e48c 100644 --- a/tests/worker_v2/test_omni_model_state.py +++ b/tests/worker_v2/test_omni_model_state.py @@ -235,6 +235,64 @@ def test_static_decode_embeddings_refresh_from_input_ids(): assert OmniModelState._preprocess_result_needs_writeback(original, original.view_as(original)) is True +@pytest.mark.parametrize( + "declared,mm_inputs,encoder_dim,reuses", + [(True, True, 4, True), (False, True, 4, False), (True, False, 4, False), (True, True, 6, False)], +) +def test_preprocess_model_reuses_encoder_buffer_only_when_declared( + monkeypatch, declared, mm_inputs, encoder_dim, reuses +): + def default_init(self, vllm_config, model, encoder_cache, device): + self.vllm_config = vllm_config + self.model_config = vllm_config.model_config + self.scheduler_config = vllm_config.scheduler_config + self.model = model + self.device = device + self.max_num_reqs = 2 + self.max_num_tokens = 8 + self.dtype = torch.float32 + self.supports_mm_inputs = encoder_cache is not None + if self.supports_mm_inputs: + self.encoder_runner = SimpleNamespace(inputs_embeds=torch.zeros(8, encoder_dim)) + + class Model(torch.nn.Module): + has_preprocess = True + preprocess_keeps_mm_inputs = declared + + def embed_input_ids(self, input_ids): + return torch.zeros(input_ids.shape[0], 4) + + monkeypatch.setattr(DefaultModelState, "__init__", default_init) + config = SimpleNamespace(model_config=SimpleNamespace(), scheduler_config=SimpleNamespace(max_num_seqs=2)) + state = OmniModelState(config, Model(), object() if mm_inputs else None, torch.device("cpu")) + assert state.preprocess_keeps_mm_inputs is reuses + assert state._static_inputs_embeds.shape == (8, 4) + encoder_runner = getattr(state, "encoder_runner", None) + assert (encoder_runner is not None and state._static_inputs_embeds is encoder_runner.inputs_embeds) is reuses + + +def test_preprocess_on_encoder_embeddings_keeps_media_rows(): + """With native multimodal inputs kept, the encoder buffer already holds this step's rows.""" + state = _make_state(has_preprocess=True) + encoder_embeds = torch.tensor([[7.0, 7.0], [1.0, 1.0]]) # row 0: image features, row 1: text + state._static_inputs_embeds = encoder_embeds + state.preprocess_keeps_mm_inputs = True + state.model.embed_input_ids = lambda input_ids: torch.zeros(input_ids.shape[0], 2) + duplex_embeds = torch.tensor([[3.0, 3.0]]) + + def preprocess(input_ids, input_embeds, **info): + # A duplex append replaces its own row; a chat row passes through. + return input_ids, duplex_embeds if info["req_id"] == "duplex" else input_embeds, {} + + state.model.preprocess = preprocess + _fill_buffers(state, "chat", "duplex") + state.run_preprocess( + _DummyInputBatch([0, 1]), + {"input_ids": torch.tensor([151667, 42], dtype=torch.long), "inputs_embeds": encoder_embeds[:2]}, + ) + assert torch.equal(encoder_embeds, torch.tensor([[7.0, 7.0], [3.0, 3.0]])) + + def test_moss_local_decode_runs_depth_predictor_and_routes_eos(mocker): """Exercise MRV2 dispatch through the real Local hook, output and logits.""" from vllm.sampling_params import SamplingParams diff --git a/vllm_omni/config/stage_config.py b/vllm_omni/config/stage_config.py index 2b59f4ed153..47acbfa82d9 100644 --- a/vllm_omni/config/stage_config.py +++ b/vllm_omni/config/stage_config.py @@ -252,6 +252,7 @@ class StagePipelineConfig: async_chunk_process_next_stage_input_func: str | None = None sync_process_input_func: str | None = None supports_native_mrv2_data_plane: bool = False + supports_duplex_mrv2: bool = False # Rewrites the Stage-0 view of a raw prompt before vLLM input processing. # The callable receives ``(prompt, sampling_params_list)``; downstream # stages continue to receive the original prompt. @@ -1012,10 +1013,12 @@ def validate_native_mrv2_session(deploy: DeployConfig, ps: StagePipelineConfig, Streaming-session prompt replacement exists only in the V1 chunk adapter, so a stage that receives from an upstream stage on MRv2 supports - turn-based sessions only. + turn-based sessions only unless it explicitly declares supports_duplex_mrv2. """ if stage_runner != "v2" or not ps.supports_native_mrv2_data_plane or not ps.input_sources: return + if deploy.session_mode == "duplex" and ps.supports_duplex_mrv2: + return if deploy.session_mode != "turn": raise ValueError( f"stage {ps.stage_id}: model_runner v2 supports session_mode 'turn' only, got " diff --git a/vllm_omni/core/sched/omni_ar_scheduler.py b/vllm_omni/core/sched/omni_ar_scheduler.py index f7c4c3a6b21..6cd34177996 100644 --- a/vllm_omni/core/sched/omni_ar_scheduler.py +++ b/vllm_omni/core/sched/omni_ar_scheduler.py @@ -1037,22 +1037,28 @@ def _update_request_as_session(self, session: Request, update: StreamingUpdate) "update_streaming_prompt_for_condition", None, ) - if stage_id != 0 and streaming_prompt_payload is not None and callable(update_streaming_prompt): + native_prompt = bool(getattr(self, "_native_data_plane", False)) + if ( + stage_id != 0 + and streaming_prompt_payload is not None + and (callable(update_streaming_prompt) or native_prompt) + ): mm_feature_base = session.num_computed_tokens try: - replaced = update_streaming_prompt( - streaming_prompt_payload, - session, - update_prompt=True, - ) + if callable(update_streaming_prompt): + replaced = update_streaming_prompt(streaming_prompt_payload, session, update_prompt=True) + else: + replaced = self._update_native_streaming_prompt(session, streaming_prompt_payload) except ValueError as exc: # This streaming update has already been dequeued. Report the # permanent contract failure so the next scheduling pass # finishes only this request instead of crashing EngineCore. - # callable(update_streaming_prompt) above implies a live - # adapter; the assert narrows it for the type checker. - assert chunk_transfer_adapter is not None - chunk_transfer_adapter.record_receive_failure(req_id, str(exc)) + if chunk_transfer_adapter is not None: + chunk_transfer_adapter.record_receive_failure(req_id, str(exc)) + else: + self._streaming_context_overflow[req_id] = (session.client_index, str(exc)) + if not session.is_finished(): + self.finish_requests((req_id,), RequestStatus.FINISHED_ERROR) return if replaced is not None: if replaced: @@ -1092,6 +1098,18 @@ def _update_request_as_session(self, session: Request, update: StreamingUpdate) if hasattr(update, "model_intermediate_buffer"): session.model_intermediate_buffer = update.model_intermediate_buffer + def _update_native_streaming_prompt(self, session: Request, payload: dict[str, Any]) -> bool | None: + """Delegate streaming prompt policy to the stage's payload processor.""" + from vllm.utils.import_utils import resolve_obj_by_qualname + + if not hasattr(self, "_native_prompt_hook"): + path = getattr(self.vllm_config.model_config, "custom_process_next_stage_input_func", None) + processor = resolve_obj_by_qualname(path) if path else None + self._native_prompt_hook = getattr(processor, "update_streaming_prompt_for_condition", None) + if self._native_prompt_hook is None: + return None + return self._native_prompt_hook(self.vllm_config.model_config, session, payload) + def _maybe_reanchor_streaming_window( self, session: Request, diff --git a/vllm_omni/core/sched/output.py b/vllm_omni/core/sched/output.py index 60195e7d339..7a5b910670c 100644 --- a/vllm_omni/core/sched/output.py +++ b/vllm_omni/core/sched/output.py @@ -31,6 +31,7 @@ class OmniNewRequestData(NewRequestData): external_req_id: str | None = None additional_information: AdditionalInformationPayload | dict[str, object] | None = None model_intermediate_buffer: dict[str, object] | None = None + resumable: bool = False @classmethod def from_base( @@ -45,6 +46,7 @@ def from_base( external_req_id=getattr(request, "external_req_id", None), additional_information=getattr(request, "additional_information", None), model_intermediate_buffer=getattr(request, "model_intermediate_buffer", None), + resumable=bool(getattr(request, "resumable", False)), ) @classmethod @@ -79,6 +81,7 @@ def from_request( prefill_token_ids=prefill_token_ids, additional_information=getattr(request, "additional_information", None), model_intermediate_buffer=getattr(request, "model_intermediate_buffer", None), + resumable=bool(getattr(request, "resumable", False)), ) diff --git a/vllm_omni/deploy/minicpmo_4_5_duplex_mrv2.yaml b/vllm_omni/deploy/minicpmo_4_5_duplex_mrv2.yaml new file mode 100644 index 00000000000..c259471dcaf --- /dev/null +++ b/vllm_omni/deploy/minicpmo_4_5_duplex_mrv2.yaml @@ -0,0 +1,23 @@ +# Opt-in CUDA Model Runner V2 profile for MiniCPM-o 4.5 full-duplex serving: +# Thinker, Talker and Code2Wav all run on MRv2. +# Enables MRv2 execution for duplex sessions, retaining native +# Stage-0 streaming audio preprocessing, KV sliding window reanchor, +# device duplex sampling hooks, and streaming Code2Wav waveform synthesis. +base_config: minicpmo_4_5.yaml + +session_mode: duplex +model_runner: v2 + +stages: + - stage_id: 0 + async_chunk: false + +platforms: + npu: + model_runner: v1 + xpu: + model_runner: v1 + rocm: + model_runner: v1 + musa: + model_runner: v1 diff --git a/vllm_omni/deploy/minicpmo_4_5_turn_mrv2.yaml b/vllm_omni/deploy/minicpmo_4_5_turn_mrv2.yaml index d0460108d17..d58b0076ff8 100644 --- a/vllm_omni/deploy/minicpmo_4_5_turn_mrv2.yaml +++ b/vllm_omni/deploy/minicpmo_4_5_turn_mrv2.yaml @@ -1,8 +1,8 @@ # Opt-in CUDA Model Runner V2 profile for MiniCPM-o 4.5 turn serving: the # Thinker (stage 0), Talker (stage 1) and Code2Wav (stage 2) run on MRv2. # The Thinker uses native multimodal encoding and tensor-only FULL graphs; -# downstream stages use the native data plane. These stages support -# turn-based sessions only (duplex needs the V1 chunk adapter). +# downstream stages use the native data plane. This profile serves +# turn-based sessions; duplex on MRv2 uses minicpmo_4_5_duplex_mrv2.yaml. # # On MRv2 the Talker keeps its sampled codec id, 16-frame penalty history and # EOS state on the device (no per-step D2H wait), and Code2Wav schedules diff --git a/vllm_omni/model_executor/models/minicpmo_4_5/duplex/mrv2_sampling.py b/vllm_omni/model_executor/models/minicpmo_4_5/duplex/mrv2_sampling.py new file mode 100644 index 00000000000..d7949ba1a04 --- /dev/null +++ b/vllm_omni/model_executor/models/minicpmo_4_5/duplex/mrv2_sampling.py @@ -0,0 +1,198 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +"""Adapt the MiniCPM duplex policy to MRv2's request slots and sampler output.""" + +from dataclasses import replace +from types import SimpleNamespace +from typing import Any + +import torch +from vllm.v1.worker.gpu.input_batch import get_num_sampled_and_rejected +from vllm.v1.worker.gpu.sample.output import SamplerOutput + +from vllm_omni.model_executor.duplex_sampling import DuplexSamplingHelper +from vllm_omni.utils.device_copy import index_to_device +from vllm_omni.worker_v2.omni_sampler import OmniSampler + + +def _keep_thinker_payload(payload: dict[str, Any], num_sampled: list[int]) -> dict[str, Any]: + return payload + + +def _v1_shaped_runner(input_batch: Any, infos: dict[str, Any]) -> SimpleNamespace: + """Present MRv2 request params in the V1 runner shape DuplexSamplingHelper reads.""" + params = [ + (infos.get(str(request_id)) or {}).get("sampling_params") or SimpleNamespace() + for request_id in input_batch.req_ids + ] + return SimpleNamespace( + input_batch=SimpleNamespace( + req_ids=input_batch.req_ids, + temperature_cpu=[getattr(sp, "temperature", None) for sp in params], + top_k_cpu=[getattr(sp, "top_k", None) for sp in params], + top_p_cpu=[getattr(sp, "top_p", None) for sp in params], + ), + model_intermediate_buffer=infos, + requests={ + str(request_id): SimpleNamespace(sampling_params=sp) for request_id, sp in zip(input_batch.req_ids, params) + }, + ) + + +class MiniCPMO45DuplexSampler(OmniSampler): + """Keep policy RNG/session state separate from the stock MRv2 sampler.""" + + def __init__(self, base_sampler: Any, model: Any) -> None: + super().__init__(base_sampler) + self.model = model + self.generators: dict[str, torch.Generator] = {} + self._helper = DuplexSamplingHelper() + self._histories: dict[str, tuple[tuple[int, int, int | None], list[int]]] = {} + self._pending_history: Any = None + model._mrv2_duplex_sampler = self + + def forget_requests(self, request_ids) -> None: + for request_id in request_ids: + self.generators.pop(request_id, None) + self._histories.pop(request_id, None) + self._helper.active_request_ids.discard(request_id) + getattr(self.model, "_mrv2_sampling_infos", {}).pop(request_id, None) + + def _commit_history(self): + pending, self._pending_history = self._pending_history, None + if pending is None: + return + entries, host, ready = pending + if ready is not None: + # Wait only for the previous sample, never for this step's forward. + ready.synchronize() + for (req_id, cached), token in zip(entries, host.tolist()): + if self._histories.get(req_id) is cached: + cached[1].append(int(token)) + + def _defer_history(self, rows, sampled_ids): + pending = getattr(self.model, "_minicpmo45_duplex_pending_samples", None) + if pending is not None and pending.row_idxs == list(range(len(rows))): + # The normal policy already snapshots final IDs with its latch + # updates. Share that owned host buffer/event; do not copy twice. + host, ready = pending.host[:, 0], pending.event + else: + # Mixed lookahead/re-emitted EOS rows or the synchronous policy + # fallback still need one owned snapshot of the final result. + tokens = sampled_ids.reshape(-1) + host = torch.empty(tokens.shape, dtype=tokens.dtype, pin_memory=tokens.device.type == "cuda") + host.copy_(tokens, non_blocking=True) + ready = torch.Event(device=tokens.device) if tokens.device.type == "cuda" else None + if ready is not None: + ready.record() + self._pending_history = ([(row.request_id, self._histories[row.request_id]) for row in rows], host, ready) + + def _metadata(self, input_batch, rows, infos, device): + """Reuse prior sampled IDs instead of synchronously reading the device ledger. + + Without speculative decoding, the CPU scheduled lengths are exact. + Only admission/replay with an uncached output prefix needs a ledger read. + The normal prefill and decode paths use host metadata and the previous + step's asynchronous sampled-ID snapshot. + """ + self._commit_history() + histories, params, generators = [], [], {} + for local_row, row in enumerate(rows): + slot = int(input_batch.idx_mapping_np[row.row_idx]) + prompt_len = int(self.req_states.prompt_len.np[slot]) + end = int(input_batch.num_computed_tokens_np[row.row_idx]) + int( + input_batch.num_scheduled_tokens[row.row_idx] + ) + key = (slot, prompt_len, row.seq) + previous_key, history = self._histories.get(row.request_id, (None, [])) + if key != previous_key: + history = [] + elif end - prompt_len < len(history): + del history[max(0, end - prompt_len) :] + start = prompt_len + len(history) + if end > start: + history.extend(self.req_states.all_token_ids.gpu[slot, start:end].tolist()) + self._histories[row.request_id] = (key, history) + histories.append(list(history)) + sp = infos[row.request_id]["sampling_params"] + params.append(sp) + if sp.seed is not None: + generator = self.generators.get(row.request_id) + if generator is None: + generator = torch.Generator(device=device).manual_seed(sp.seed) + self.generators[row.request_id] = generator + generators[local_row] = generator + return SimpleNamespace( + output_token_ids=histories, + generators=generators, + temperature=torch.tensor([sp.temperature for sp in params]), + top_k=torch.tensor([sp.top_k for sp in params]), + top_p=torch.tensor([sp.top_p for sp in params]), + all_greedy=all(sp.temperature <= 0 for sp in params), + ) + + def sample_step(self, hidden_states, input_batch, req_states, grammar_output, standard_sample): + # Publish the Thinker payload through the MRv2 output-channel contract: + # its latent row ledger is inter-stage only. The full-payload fallback + # mirrors it into multimodal_output too, so every llm2tts row would be + # accumulated twice and the per-unit ledger lookup would fail. + output = super().sample_step(hidden_states, input_batch, req_states, grammar_output, standard_sample) + return replace(output, include_hidden_states=False, finalize_multimodal=_keep_thinker_payload) + + def __call__(self, logits: torch.Tensor, input_batch: Any) -> SamplerOutput: + infos = getattr(self.model, "_mrv2_sampling_infos", {}) + helper = self._helper + runner = _v1_shaped_runner(input_batch, infos) + for request_id in input_batch.req_ids: + helper.refresh_active_request(runner, request_id) + rows = helper.rows(runner) + if rows and input_batch.num_draft_tokens: + raise NotImplementedError("MiniCPM-o MRv2 duplex sampling does not support speculative decoding") + # Partial prefills are discarded by the runner and must not mutate the + # policy latches or advance their generators. + rows = tuple( + row + for row in rows + if not input_batch.is_prefilling_np[row.row_idx] + or int(input_batch.num_computed_prefill_tokens_np[row.row_idx]) + + int(input_batch.num_scheduled_tokens[row.row_idx]) + >= int(input_batch.prefill_len_np[row.row_idx]) + ) + metadata = self._metadata(input_batch, rows, infos, logits.device) if rows else None + row_idxs = [row.row_idx for row in rows] + if logits.shape[0] == input_batch.num_reqs and row_idxs == list(range(input_batch.num_reqs)): + # Every row follows the policy: skip the index upload and the + # full-vocabulary gather. As on V1, the policy's force-listen mask + # then lands on the logits a fallback sampler would read. + selected, policy_logits = None, logits + else: + selected = index_to_device(row_idxs, logits.device) + policy_logits = logits.index_select(0, selected) + self.model.prepare_duplex_sampling( + policy_logits, metadata, tuple(replace(row, row_idx=i) for i, row in enumerate(rows)) + ) + if not rows: + return self.base_sampler(logits, input_batch) + policy_output = self.model.sample(policy_logits, metadata) + if policy_output is None: + return self.base_sampler(logits, input_batch) + self._defer_history(rows, policy_output.sampled_token_ids) + if len(rows) != input_batch.num_reqs: + # The stock sampler owns counts/logprobs for any ordinary rows. + output = self.base_sampler(logits, input_batch) + output.sampled_token_ids.index_copy_(0, selected, policy_output.sampled_token_ids.long()) + return output + counts, rejected = get_num_sampled_and_rejected( + input_batch.seq_lens.new_ones(input_batch.num_reqs), + input_batch.seq_lens, + input_batch.cu_num_logits, + input_batch.idx_mapping, + self.req_states.prefill_len.gpu, + ) + return SamplerOutput( + sampled_token_ids=policy_output.sampled_token_ids.long(), + logprobs_tensors=policy_output.logprobs_tensors, + num_nans=None, + num_sampled=counts, + num_rejected=rejected, + ) diff --git a/vllm_omni/model_executor/models/minicpmo_4_5/duplex/window_kv.py b/vllm_omni/model_executor/models/minicpmo_4_5/duplex/window_kv.py index bb7a5f4b54f..72305dc401e 100644 --- a/vllm_omni/model_executor/models/minicpmo_4_5/duplex/window_kv.py +++ b/vllm_omni/model_executor/models/minicpmo_4_5/duplex/window_kv.py @@ -843,6 +843,15 @@ def resolve_group_block_ids( actual_group = group_idx if group_idx < len(raw_blocks) else 0 return [int(b) for b in raw_blocks[actual_group]] return [int(b) for b in raw_blocks] + # MRv2 owns request slots and block tables on the runner, rather than + # on its transient InputBatch. Read after staged writes are applied. + req_states = getattr(runner, "req_states", None) + block_tables = getattr(runner, "block_tables", None) + if req_states is not None and block_tables is not None: + slot = req_states.req_id_to_index.get(req_id) + if slot is not None: + num_blocks = int(block_tables.num_blocks.np[group_idx, slot]) + return block_tables.block_tables[group_idx].gpu[slot, :num_blocks].tolist() # Fallback to input_batch block_table bt = getattr(getattr(runner, "input_batch", None), "block_table", None) if bt is not None: @@ -859,16 +868,26 @@ def resolve_group_block_ids( @classmethod def maybe_apply_reanchor(cls, runner: Any, scheduler_output: Any = None) -> None: """Apply in-place KV reanchor and rotation on worker before model forward.""" - if not hasattr(runner, "input_batch") or runner.input_batch is None: - return - num_reqs = getattr(runner.input_batch, "num_reqs", len(runner.input_batch.req_ids)) - req_ids = runner.input_batch.req_ids[:num_reqs] + req_states = getattr(runner, "req_states", None) + if req_states is not None and scheduler_output is not None: + req_ids = list(scheduler_output.num_scheduled_tokens) + else: + if not hasattr(runner, "input_batch") or runner.input_batch is None: + return + num_reqs = getattr(runner.input_batch, "num_reqs", len(runner.input_batch.req_ids)) + req_ids = runner.input_batch.req_ids[:num_reqs] applied_reanchors = getattr(runner, "_applied_stage0_reanchor_ids", None) if applied_reanchors is None: applied_reanchors = runner._applied_stage0_reanchor_ids = set() for req_idx, req_id in enumerate(req_ids): - info = runner.model_intermediate_buffer.get(req_id) + if req_states is not None: + slot = req_states.req_id_to_index.get(req_id) + if slot is None: + continue + info = runner.model_state.intermediate_buffer.buffers[slot] + else: + info = runner.model_intermediate_buffer.get(req_id) if not isinstance(info, dict): continue duplex = info.get("duplex") @@ -906,12 +925,14 @@ def maybe_apply_reanchor(cls, runner: Any, scheduler_output: Any = None) -> None # Scheduler is authoritative for logical state (block_table and computed tokens). # The parent runner's _update_states() already installed the post-compaction block IDs # and decremented num_computed_tokens_cpu. We do NOT double-compact or double-decrement here. - old_computed = int( - reanchor.get( - "old_computed_tokens", - int(runner.input_batch.num_computed_tokens_cpu[req_idx]) + plan.delta, - ) - ) + old_computed = reanchor.get("old_computed_tokens") + if old_computed is None: + if req_states is not None: + slot = req_states.req_id_to_index[req_id] + old_computed = int(req_states.num_computed_tokens_np[slot]) + plan.delta + else: + old_computed = int(runner.input_batch.num_computed_tokens_cpu[req_idx]) + plan.delta + old_computed = int(old_computed) req_state = runner.requests.get(req_id) if hasattr(runner, "requests") else None mrope_pos = getattr(req_state, "mrope_positions", None) if req_state is not None else None @@ -938,9 +959,14 @@ def maybe_apply_reanchor(cls, runner: Any, scheduler_output: Any = None) -> None inv_freq = cls.get_rope_inv_freq(runner) kv_groups = getattr(runner, "kv_cache_group_ids", None) block_size = int(getattr(getattr(runner, "cache_config", None), "block_size", 16) or 16) + # Layers share their KV group's table; on MRv2 each lookup is a device read. + group_block_ids: dict[int, list[int]] = {} for layer_idx, kv_cache in enumerate(runner.kv_caches): group_idx = kv_groups[layer_idx] if (kv_groups and layer_idx < len(kv_groups)) else 0 - layer_block_ids = cls.resolve_group_block_ids(runner, req_id, req_idx, group_idx=group_idx) + layer_block_ids = group_block_ids.get(group_idx) + if layer_block_ids is None: + layer_block_ids = cls.resolve_group_block_ids(runner, req_id, req_idx, group_idx=group_idx) + group_block_ids[group_idx] = layer_block_ids rotate_cached_keys( kv_cache, block_ids=layer_block_ids, diff --git a/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni.py b/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni.py index 25254480624..5acc4aa97e7 100644 --- a/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni.py +++ b/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni.py @@ -52,6 +52,49 @@ logger = init_logger(__name__) +def _duplex_row_outputs(request_infos: list[Any]) -> dict[str, Any]: + """Build the per-row duplex handoff metadata shared by the V1 and MRv2 Thinker outputs.""" + duplex_rows = [] + for req_info in request_infos: + duplex_info = req_info.get("duplex") if isinstance(req_info, dict) else None + duplex_rows.append(duplex_info if isinstance(duplex_info, dict) else {}) + + outputs: dict[str, Any] = {} + prompt_rows = [] + for duplex_info in duplex_rows: + prompt_token_ids = duplex_info.get("duplex_prompt_token_ids") + # This is a complete per-handoff snapshot, not a generated + # tensor delta. Keep it as row-local metadata so output + # accumulation replaces the previous value instead of + # attempting to concatenate variable-length prompts. + prompt_rows.append(list(prompt_token_ids) if isinstance(prompt_token_ids, list) else None) + if any(row is not None for row in prompt_rows): + outputs["duplex_prompt_token_ids"] = prompt_rows + + special_rows = [ + info if isinstance(info := duplex_info.get("special_token_ids"), dict) else {} for duplex_info in duplex_rows + ] + special_values: dict[str, int] = {} + for special in special_rows: + for key, value in special.items(): + if isinstance(key, str) and isinstance(value, int) and value >= 0: + special_values.setdefault(key, value) + if special_values: + # These are tokenizer constants. A row whose append built no unit has + # none on MRv2, and a None entry would stay in that request's + # accumulated metadata over the values later appends publish. + # Host tensors: every consumer reads them on the host, and a pageable + # host->device copy per row and key would wait for the whole forward. + outputs["meta"] = { + key: [ + torch.tensor([value if isinstance(value, int) and value >= 0 else constant], dtype=torch.long) + for value in (special.get(key) for special in special_rows) + ] + for key, constant in sorted(special_values.items()) + } + return outputs + + @dataclass(slots=True) class _MiniCPMO45PendingSamples: """A deferred Stage-0 duplex step awaiting its host commit.""" @@ -123,12 +166,6 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): self.model_stage = vllm_config.model_config.model_stage self.model_sampler_wants_sampling_params = self.model_stage == "tts" self._use_v2_model_runner = bool(getattr(vllm_config.model_config, "use_v2_model_runner", False)) - if ( - self.model_stage == "llm" - and self._use_v2_model_runner - and getattr(vllm_config.model_config, "session_mode", "turn") != "turn" - ): - raise NotImplementedError("MiniCPM-o duplex Thinker requires model_runner: v1") if ( self.model_stage == "llm" and self._use_v2_model_runner @@ -230,14 +267,21 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): # preprocess for duplex audio, while the Talker converts the # tts_token_ids/tts_hidden_states handoff into its conditioning # embeddings and initializes request-local codec generation state. - # Turn-mode Thinker inputs use the native multimodal encoder/cache on - # MRv2. Marking it as a custom-preprocess model disables that encoder - # path in the runner. Duplex still uses the V1 preprocess hook. - self.has_preprocess = self.model_stage == "tts" or not self._use_v2_model_runner + is_duplex = getattr(vllm_config.model_config, "session_mode", "turn") == "duplex" + self.has_preprocess = ( + self.model_stage == "tts" or not self._use_v2_model_runner or (self.model_stage == "llm" and is_duplex) + ) + # The duplex Thinker also serves media chat requests, whose image, + # audio and video features come from the native MRv2 encoder. + self.preprocess_keeps_mm_inputs = self.model_stage == "llm" and is_duplex and self._use_v2_model_runner # Neither AR stage has a postprocess, so step outputs can use the # runner's async snapshot instead of a blocking per-step D2H. self.use_async_omni_output = self.model_stage in {"llm", "tts"} + if self.model_stage == "llm" and is_duplex and self._use_v2_model_runner: + # Only the duplex Thinker exposes the MRv2 batch-preprocess hook. + self.preprocess_batch_mrv2 = self._preprocess_batch_mrv2 + if self.model_stage == "llm" and getattr(vllm_config.model_config, "session_mode", "turn") == "duplex": # Build the Stage-0 duplex runtime (remote-code processor and # tokenizer) with the model. Built lazily, it costs several seconds @@ -283,11 +327,9 @@ def prepare_duplex_sampling( self._minicpmo45_duplex_row_sessions = { row.row_idx: row.session_id for row in rows if row.session_id is not None } - request_sessions = getattr(self, "_minicpmo45_duplex_request_sessions", None) - if not isinstance(request_sessions, dict): - request_sessions = {} - self._minicpmo45_duplex_request_sessions = request_sessions - request_sessions.update({row.request_id: row.session_id for row in rows if row.session_id is not None}) + self._minicpmo45_duplex_request_session_map().update( + {row.request_id: row.session_id for row in rows if row.session_id is not None} + ) self._minicpmo45_duplex_row_payloads = {row.row_idx: row.payload for row in rows if row.payload is not None} self._minicpmo45_duplex_row_max_tokens = { row.row_idx: row.max_tokens for row in rows if row.max_tokens is not None @@ -427,13 +469,24 @@ def preprocess( embeds = input_embeds if input_embeds is not None else self.get_input_embeddings(input_ids) return input_ids, embeds, {} + if getattr(self, "_use_v2_model_runner", False) and self.model_stage == "llm": + req_id = kwargs.get("req_id") + if req_id is not None: + if not hasattr(self, "_mrv2_sampling_infos"): + self._mrv2_sampling_infos = {} + self._mrv2_sampling_infos[str(req_id)] = kwargs + duplex = kwargs.get("duplex") if not isinstance(duplex, dict) or duplex.get("data_plane") is not True: embeds = input_embeds if input_embeds is not None else self.get_input_embeddings(input_ids) return input_ids, embeds, {} - prompt_len_meta = kwargs.get("duplex_prompt_len") - token_offset_meta = kwargs.get("duplex_token_offset", 0) + prompt_len_meta = kwargs.get("_omni_prompt_len") + if prompt_len_meta is None: + prompt_len_meta = kwargs.get("duplex_prompt_len") + token_offset_meta = kwargs.get("_omni_num_computed_tokens") + if token_offset_meta is None: + token_offset_meta = kwargs.get("duplex_token_offset", 0) if ( isinstance(prompt_len_meta, int) and isinstance(token_offset_meta, int) @@ -457,6 +510,12 @@ def preprocess( # A deferred sample of this session updates the latches its append reads. self._commit_minicpmo45_duplex_pending_samples(session_ids={session_id}) state = self._minicpmo45_duplex_session_state(helper, session_id, duplex) + req_id = kwargs.get("req_id") + if req_id is not None: + # Own the session from its first append: MRv2 samples no partial + # prefill row, so a request cancelled mid-prefill never reaches + # prepare_duplex_sampling and on_requests_finished must still free it. + self._minicpmo45_duplex_request_session_map()[str(req_id)] = session_id prefill_kwargs = self._minicpmo45_duplex_prefill_kwargs(duplex, payload) seq = prefill_kwargs["seq"] result = helper.take_staged_prefill(state, prefill_kwargs["epoch"], seq) @@ -488,7 +547,9 @@ def preprocess( ) full_req_embeds = result["inputs_embeds"].to(device=input_ids.device, dtype=target_dtype) full_input_token_ids = list(result.get("input_token_ids") or []) - prompt_len = kwargs.get("duplex_prompt_len") + prompt_len = kwargs.get("_omni_prompt_len") + if prompt_len is None: + prompt_len = kwargs.get("duplex_prompt_len") try: prompt_len = int(prompt_len) if prompt_len is not None else int(full_req_embeds.shape[0]) except (TypeError, ValueError): @@ -523,7 +584,9 @@ def preprocess( ) span_len = int(input_ids.shape[0]) - token_offset = kwargs.get("duplex_token_offset", 0) + token_offset = kwargs.get("_omni_num_computed_tokens") + if token_offset is None: + token_offset = kwargs.get("duplex_token_offset", 0) try: token_offset = max(0, int(token_offset)) except (TypeError, ValueError): @@ -609,6 +672,29 @@ def preprocess_batch( if len(staged) >= 2: helper.stage_prefill_batch(staged) + def _preprocess_batch_mrv2(self, *, req_infos: list[dict[str, Any]], device: torch.device) -> None: + """MRv2 ``preprocess_batch_mrv2`` hook (prefill rows only): ``preprocess_batch``'s cross-session batching. + + Without it MRv2 builds every duplex append in ``preprocess`` one session + at a time: one batch-1 streaming-encoder pass per unit and no shared + vision-tower call for camera frames. + """ + buffers = { + str(info["req_id"]): info + for info in req_infos + if isinstance(info, dict) and info.get("req_id") is not None and isinstance(info.get("duplex"), dict) + } + if buffers: + self.preprocess_batch(req_ids=list(buffers), model_intermediate_buffer=buffers, device=device) + + def _minicpmo45_duplex_request_session_map(self) -> dict[str, str]: + """Request id -> Stage-0 session id, read by ``on_requests_finished`` to free the session.""" + request_sessions = getattr(self, "_minicpmo45_duplex_request_sessions", None) + if not isinstance(request_sessions, dict): + request_sessions = {} + self._minicpmo45_duplex_request_sessions = request_sessions + return request_sessions + def _minicpmo45_duplex_session_state(self, helper, session_id: str, duplex: dict[str, Any]): """The Stage-0 state of ``session_id``, created with its session context on first use.""" state = helper.sessions.get(session_id) @@ -763,51 +849,7 @@ def forward( runtime_info = kwargs.get("runtime_additional_information") if runtime_info and isinstance(runtime_info, list) and len(runtime_info) > 0: - duplex_rows = [] - for req_info in runtime_info: - duplex_info = req_info.get("duplex") if isinstance(req_info, dict) else None - duplex_rows.append(duplex_info if isinstance(duplex_info, dict) else {}) - - prompt_rows = [] - for duplex_info in duplex_rows: - prompt_token_ids = duplex_info.get("duplex_prompt_token_ids") - # This is a complete per-handoff snapshot, not a generated - # tensor delta. Keep it as row-local metadata so output - # accumulation replaces the previous value instead of - # attempting to concatenate variable-length prompts. - prompt_rows.append(list(prompt_token_ids) if isinstance(prompt_token_ids, list) else None) - if any(row is not None for row in prompt_rows): - multimodal_outputs["duplex_prompt_token_ids"] = prompt_rows - - special_keys = { - key - for duplex_info in duplex_rows - for key, value in ( - duplex_info.get("special_token_ids", {}).items() - if isinstance(duplex_info.get("special_token_ids"), dict) - else () - ) - if isinstance(key, str) and isinstance(value, int) and value >= 0 - } - if special_keys: - # Host tensors: every consumer reads them on the host, and a pageable - # host->device copy per row and key would wait for the whole forward. - multimodal_outputs["meta"] = { - key: [ - torch.tensor([int(value)], dtype=torch.long) - if isinstance(value, int) and value >= 0 - else None - for duplex_info in duplex_rows - for value in [ - ( - duplex_info.get("special_token_ids", {}).get(key) - if isinstance(duplex_info.get("special_token_ids"), dict) - else None - ) - ] - ] - for key in sorted(special_keys) - } + multimodal_outputs.update(_duplex_row_outputs(runtime_info)) return OmniOutput( text_hidden_states=text_hidden_states, multimodal_outputs=multimodal_outputs, @@ -835,18 +877,48 @@ def make_omni_output_mrv2(self, model_outputs, *, input_batch, req_states, model req_states=req_states, model_intermediate_buffer=model_intermediate_buffer, ) - if any(isinstance(info, dict) and info.get("duplex") for info in model_intermediate_buffer): - raise NotImplementedError("MiniCPM-o duplex Thinker requires model_runner: v1") num_tokens = model_outputs.shape[0] + multimodal_outputs: dict[str, Any] = { + "latent": model_outputs, + "latent_input_ids": input_batch.input_ids[:num_tokens].reshape(-1, 1), + "latent_positions": input_batch.positions[:num_tokens].reshape(-1, 1), + } + if model_intermediate_buffer and any( + isinstance(info, dict) and info.get("duplex") for info in model_intermediate_buffer + ): + # The prompt snapshot and special ids change only with an append, + # and the output accumulator keeps the last value of a missing key. + # Publish them on steps that complete an append's prefill instead + # of copying every row's whole prompt on every decode step. + num_reqs = int(input_batch.num_reqs) + completes_prefill = input_batch.is_prefilling_np[:num_reqs] & ( + input_batch.num_computed_prefill_tokens_np[:num_reqs] + input_batch.num_scheduled_tokens[:num_reqs] + >= input_batch.prefill_len_np[:num_reqs] + ) + if completes_prefill.any(): + multimodal_outputs.update(_duplex_row_outputs(model_intermediate_buffer)) + return OmniOutput( text_hidden_states=model_outputs, - multimodal_outputs={ - "latent": model_outputs, - "latent_input_ids": input_batch.input_ids[:num_tokens].reshape(-1, 1), - "latent_positions": input_batch.positions[:num_tokens].reshape(-1, 1), - }, + multimodal_outputs=multimodal_outputs, ) + @property + def mm_outputs_fresh_per_step(self) -> bool: + """MRv2 Talker outputs are allocated per step by ``make_omni_output_mrv2``.""" + return self.model_stage == "tts" and self._use_v2_model_runner + + def mrv2_custom_sampler(self, sampler: Any) -> tuple[Any, None]: + if self.model_stage == "tts": + return self.talker.mrv2_custom_sampler(sampler) + if self.model_stage == "llm" and getattr(self.vllm_config.model_config, "session_mode", "turn") == "duplex": + from vllm_omni.model_executor.models.minicpmo_4_5.duplex.mrv2_sampling import ( + MiniCPMO45DuplexSampler, + ) + + return MiniCPMO45DuplexSampler(sampler, self), None + return sampler, None + def make_omni_output(self, model_outputs, **kwargs): if self.model_stage != "tts": return model_outputs @@ -874,6 +946,9 @@ def on_requests_finished(self, finished_req_ids: set[str] | list[str]) -> None: finished = set(finished_req_ids) completed_segments = {segment for segment in forced_segments if segment[0] in finished} forced_segments.difference_update(completed_segments) + mrv2_sampler = getattr(self, "_mrv2_duplex_sampler", None) + if mrv2_sampler is not None and hasattr(mrv2_sampler, "forget_requests"): + mrv2_sampler.forget_requests(finished_req_ids) if hasattr(self.model, "on_requests_finished"): self.model.on_requests_finished(finished_req_ids) diff --git a/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni_tts.py b/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni_tts.py index f8fa783bc79..f3e7e590229 100644 --- a/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni_tts.py +++ b/vllm_omni/model_executor/models/minicpmo_4_5/minicpmo_4_5_omni_tts.py @@ -107,6 +107,45 @@ def blank_scheduler_prompt_for_penalties( return torch.full_like(prompt_token_ids, int(vocab_size)) +def _native_duplex_row_meta(info_dict: Mapping[str, Any]) -> tuple[bool, int, int, str, bool]: + """One Talker row's native-duplex fence identity, segment text and turn end, shared by V1 and MRv2 outputs.""" + native_duplex = info_dict.get("native_duplex") is True + duplex_info = info_dict.get("duplex") + if not isinstance(duplex_info, dict): + duplex_info = {} + epoch = duplex_info.get("epoch", -1) + turn_id = duplex_info.get("turn_id", -1) + if native_duplex and not all( + isinstance(value, int) and not isinstance(value, bool) and value >= 0 for value in (epoch, turn_id) + ): + raise RuntimeError( + "MiniCPM-o native duplex Talker requires non-negative integer " + f"epoch and turn_id, got epoch={epoch!r}, turn_id={turn_id!r}" + ) + meta_info = info_dict.get("meta") + if not isinstance(meta_info, dict): + meta_info = {} + segment_text = meta_info.get("native_duplex_segment_text", "") if native_duplex else "" + if not isinstance(segment_text, str): + segment_text = "" + turn_eos_id = meta_info.get("turn_eos_token_id") + ids_info = info_dict.get("ids") + tts_ids = ids_info.get("tts") if native_duplex and isinstance(ids_info, dict) else None + if isinstance(tts_ids, torch.Tensor): + contains_turn_eos = isinstance(turn_eos_id, int) and bool(torch.any(tts_ids.reshape(-1) == turn_eos_id).item()) + elif isinstance(tts_ids, (list, tuple)): + contains_turn_eos = isinstance(turn_eos_id, int) and turn_eos_id in tts_ids + else: + contains_turn_eos = False + return ( + native_duplex, + epoch if isinstance(epoch, int) else -1, + turn_id if isinstance(turn_id, int) else -1, + segment_text, + native_duplex and contains_turn_eos, + ) + + def _restore_weight_norm_weight(weight_g: torch.Tensor, weight_v: torch.Tensor) -> torch.Tensor: """Materialize ``weight_norm(..., dim=0)`` checkpoint parameters.""" return torch._weight_norm(weight_v, weight_g, dim=0) @@ -172,6 +211,7 @@ def _apply_codec_window_penalty_gpu( penalty: torch.Tensor, *, window_size: int, + history_prefix: torch.Tensor | None = None, ) -> None: """In place: ``_apply_batched_repetition_penalty`` on the runner's device token history. @@ -194,6 +234,12 @@ def _apply_codec_window_penalty_gpu( valid = positions >= start.unsqueeze(1) rows = all_token_ids.index_select(0, slots) history = rows.gather(1, positions.clamp_min(0)).long() + if history_prefix is not None: + prompt = prompt_len.index_select(0, slots).long().unsqueeze(1) + prefix_positions = (positions - prompt + window_size).clamp(0, window_size - 1) + prefix = history_prefix.index_select(0, slots).gather(1, prefix_positions) + history = torch.where(positions < prompt, prefix, history) + valid = valid | (positions < prompt) # Invalid positions count into a spare column that is dropped below. history = torch.where(valid & (history >= 0) & (history < vocab_size), history, vocab_size) frequencies = torch.zeros((num_rows, vocab_size + 1), dtype=torch.long, device=logits.device) @@ -230,12 +276,16 @@ def __init__(self, base: Any, *, window_size: int) -> None: self.repetition_penalty.copy_to_uva() self.use_window = np.zeros(max_num_reqs, dtype=bool) self.use_penalty = np.zeros(max_num_reqs, dtype=bool) + self.history_prefix = torch.full( + (max_num_reqs, self.window_size), -1, dtype=torch.long, device=self.req_states.device + ) @property def output_bin_counts(self) -> torch.Tensor: return self.base.output_bin_counts def add_request(self, req_idx: int, sampling_params: Any) -> bool: + self.history_prefix[req_idx].fill_(-1) repetition = float(getattr(sampling_params, "repetition_penalty", 1.0)) self.repetition_penalty.np[req_idx] = repetition self.use_window[req_idx] = repetition != 1.0 @@ -254,6 +304,12 @@ def apply_staged_writes(self) -> None: self.repetition_penalty.copy_to_uva() self.base.apply_staged_writes() + def set_history_prefix(self, slots: list[int], histories: list[list[int]]) -> None: + rows = [[-1] * (self.window_size - len(h[-self.window_size :])) + h[-self.window_size :] for h in histories] + self.history_prefix.index_copy_( + 0, index_to_device(slots, self.req_states.device), index_to_device(rows, self.req_states.device) + ) + def apply(self, logits: torch.Tensor, ctx: LogitsContext) -> torch.Tensor: if np.any(self.use_window[ctx.idx_mapping_np]): _apply_codec_window_penalty_gpu( @@ -264,6 +320,7 @@ def apply(self, logits: torch.Tensor, ctx: LogitsContext) -> torch.Tensor: self.req_states.prompt_len.gpu, self.repetition_penalty.gpu, window_size=self.window_size, + history_prefix=self.history_prefix, ) return self.base.apply(logits, ctx) @@ -288,6 +345,7 @@ def _install_mrv2_talker_sampler(sampler: Any, talker: "MiniCPMO45OmniTTSForCond # ``penalties_state`` is still read by the runner for output bin counts. processors[slot] = window sampler.penalties_state = window + talker._mrv2_penalty_state = window talker._mrv2_empty_speech = torch.zeros( int(sampler.req_states.max_num_reqs), dtype=torch.bool, device=sampler.req_states.device ) @@ -307,6 +365,10 @@ def __init__(self, base_sampler, talker): def __call__(self, logits: torch.Tensor, input_batch: Any) -> Any: forced = self.talker.take_mrv2_forced_eos(input_batch, self.req_states, logits.shape[0]) + masked = getattr(self.talker, "_mrv2_masked_eos", None) + self.talker._mrv2_masked_eos = None + if masked is not None and masked.shape[0] == logits.shape[0]: + logits[:, int(self.talker._codec_eos_id)].masked_fill_(masked, float("-inf")) output = self.base_sampler(logits, input_batch) if forced is not None: sampled = output.sampled_token_ids @@ -381,6 +443,7 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): self._mrv2_empty_speech: torch.Tensor | None = None # Rows the sampler must force to codec EOS, computed with this step's output. self._mrv2_forced_eos: torch.Tensor | None = None + self._mrv2_masked_eos: torch.Tensor | None = None self._mrv2_decode_rows_logged = False def _init_native_talker(self, prefix: str) -> None: @@ -521,7 +584,11 @@ def _build_streaming_recompute_embeddings( ) audio_state = self._request_audio_states.get(request_id) recent_codes = audio_state.get("recent_codes") if isinstance(audio_state, dict) else None - if isinstance(recent_codes, torch.Tensor): + ids = info_dict.get("ids") + confirmed_codes = ids.get("streaming_prompt_previous_codes") if isinstance(ids, Mapping) else None + if isinstance(confirmed_codes, list): + base_recent_codes = (*state["base_recent_codes"], *confirmed_codes)[-_CODEC_PENALTY_WINDOW:] + elif isinstance(recent_codes, torch.Tensor): base_recent_codes = recent_codes.detach().clone() elif isinstance(recent_codes, list): base_recent_codes = tuple(int(code_id) for code_id in recent_codes[-_CODEC_PENALTY_WINDOW:]) @@ -603,7 +670,8 @@ def preprocess( is_prefill = bool(info_dict.get("_omni_is_prefill", False)) state = info_dict.get("audio_state") first_call = not isinstance(state, dict) - request_id = str(info_dict.get("request_id", "0")) + request_id = info_dict.get("request_id") + request_id = str(info_dict.get("req_id", "0") if request_id is None else request_id) if is_prefill or first_call: token_ids, hidden_states = get_tts_handoff(info_dict) @@ -897,14 +965,17 @@ def make_omni_output_mrv2( # empty condition (preprocess marks such a request finished). slots: list[int] = [] flags: list[bool] = [] + histories: list[list[int]] = [] for row in np.flatnonzero(input_batch.is_prefilling_np[:num_reqs]).tolist(): info = model_intermediate_buffer[row] - if info.get("native_duplex") is True: - raise NotImplementedError("MiniCPM-o native duplex Talker requires model_runner: v1") - state = info.get("audio_state") + state = info.get("audio_state") if isinstance(info, dict) else None flags.append(bool(state.get("finished")) if isinstance(state, dict) else False) + histories.append(list(state.get("recent_codes", [])) if isinstance(state, dict) else []) slots.append(int(input_batch.idx_mapping_np[row])) if slots: + penalties = getattr(self, "_mrv2_penalty_state", None) + if penalties is not None: + penalties.set_history_prefix(slots, histories) empty_speech.index_copy_( 0, index_to_device(slots, device), @@ -924,19 +995,82 @@ def make_omni_output_mrv2( valid = decode & ~eos_input & ~empty # MiniCPMTTS.generate's max_new_token, clamped to the Talker context # (see ``preprocess``); the sample after the last allowed code is EOS. + # For native duplex chunks, limit is _DUPLEX_CODEC_FRAMES_PER_CHUNK (25), + # unless turn-end drain is active. + native_duplex_batch = any( + isinstance(info, dict) and info.get("native_duplex") is True for info in model_intermediate_buffer + ) remaining = int(self._tts_config.max_position_embeddings) - prompt_len - limit = torch.clamp(remaining, min=1, max=_OFFLINE_CODEC_MAX_NEW_TOKENS) - 1 - self._mrv2_forced_eos = empty | eos_input | (step >= limit) + if not native_duplex_batch: + limit = torch.clamp(remaining, min=1, max=_OFFLINE_CODEC_MAX_NEW_TOKENS) - 1 + self._mrv2_forced_eos = empty | eos_input | (step >= limit) + self._mrv2_masked_eos = None + else: + limits, minimums, draining = [], [], [] + for i in range(num_reqs): + info = model_intermediate_buffer[i] if i < len(model_intermediate_buffer) else {} + info_dict = info if isinstance(info, dict) else {} + if info_dict.get("native_duplex") is True: + meta = info_dict.get("meta") if isinstance(info_dict.get("meta"), dict) else {} + max_tok, min_tok = _native_duplex_chunk_budget(meta) + limits.append(max_tok - 1) + minimums.append(min_tok) + draining.append(bool(meta.get("turn_end"))) + else: + limits.append(_OFFLINE_CODEC_MAX_NEW_TOKENS - 1) + minimums.append(0) + draining.append(False) + device_limits = index_to_device(limits, device, dtype=torch.long) + limit = torch.minimum(torch.clamp(remaining - 1, min=0), device_limits) + self._mrv2_forced_eos = empty | eos_input | (step >= limit) + cadence = (step >= _DUPLEX_CODEC_FRAMES_PER_CHUNK) & ( + step % _DUPLEX_CODEC_FRAMES_PER_CHUNK < _DUPLEX_TURN_END_BOUNDARY_MASK_STEPS + ) + self._mrv2_masked_eos = ~self._mrv2_forced_eos & ( + (step < index_to_device(minimums, device, dtype=torch.long)) + | (index_to_device(draining, device, dtype=torch.bool) & cadence) + ) frame_valid = torch.zeros(num_tokens, dtype=torch.bool, device=device) frame_valid.index_copy_(0, last_rows, valid) if not self._mrv2_decode_rows_logged and not input_batch.has_prefill: self._mrv2_decode_rows_logged = True logger.info("MiniCPM-o Talker: MRv2 device-side codec output active (no host read of sampled ids)") + + meta_outputs: dict[str, Any] = {"codec_frame_valid": frame_valid} + if native_duplex_batch: + native_duplex_flags: list[torch.Tensor] = [] + duplex_epochs: list[torch.Tensor] = [] + duplex_turn_ids: list[torch.Tensor] = [] + condition_seqs: list[torch.Tensor] = [] + segment_texts_utf8: list[torch.Tensor] = [] + turn_end_flags: list[torch.Tensor] = [] + for info in model_intermediate_buffer: + info_dict = info if isinstance(info, dict) else {} + native_duplex, epoch, turn_id, segment_text, turn_end = _native_duplex_row_meta(info_dict) + meta_info = info_dict.get("meta") if isinstance(info_dict.get("meta"), dict) else {} + condition_seq = meta_info.get("streaming_condition_seq", -1) + condition_seqs.append( + torch.tensor(condition_seq if isinstance(condition_seq, int) else -1, dtype=torch.long) + ) + native_duplex_flags.append(torch.tensor(native_duplex, dtype=torch.bool)) + duplex_epochs.append(torch.tensor(epoch, dtype=torch.long)) + duplex_turn_ids.append(torch.tensor(turn_id, dtype=torch.long)) + segment_texts_utf8.append(torch.tensor(list(segment_text.encode("utf-8")), dtype=torch.uint8)) + turn_end_flags.append(torch.tensor(turn_end, dtype=torch.bool)) + meta_outputs["native_duplex"] = native_duplex_flags + meta_outputs["duplex_epoch"] = duplex_epochs + meta_outputs["duplex_turn_id"] = duplex_turn_ids + meta_outputs["streaming_condition_seq"] = condition_seqs + # Key matching Stage 2 and tts2code2wav_async_chunk expectations: + meta_outputs["llm_output_text_utf8"] = segment_texts_utf8 + meta_outputs["native_duplex_segment_text"] = segment_texts_utf8 + meta_outputs["turn_end"] = turn_end_flags + return OmniOutput( text_hidden_states=hidden, multimodal_outputs={ - "codes": {"audio": token_ids.to(dtype=torch.long).reshape(num_tokens, 1)}, - "meta": {"codec_frame_valid": frame_valid}, + "codes": {"audio": token_ids.to(dtype=torch.long, copy=True).reshape(num_tokens, 1)}, + "meta": meta_outputs, }, ) @@ -980,47 +1114,13 @@ def make_omni_output( frame_valid = [torch.empty(0, dtype=torch.bool, device="cpu") for _ in infos] for index, info in enumerate(infos): info_dict = info if isinstance(info, dict) else {} - native_duplex = info_dict.get("native_duplex") is True if emit_duplex_metadata: - duplex_info = info_dict.get("duplex") - if not isinstance(duplex_info, dict): - duplex_info = {} - epoch = duplex_info.get("epoch", -1) - turn_id = duplex_info.get("turn_id", -1) - if native_duplex and not all( - isinstance(value, int) and not isinstance(value, bool) and value >= 0 for value in (epoch, turn_id) - ): - raise RuntimeError( - "MiniCPM-o native duplex Talker requires non-negative integer " - f"epoch and turn_id, got epoch={epoch!r}, turn_id={turn_id!r}" - ) - meta_info = info_dict.get("meta") - if not isinstance(meta_info, dict): - meta_info = {} - segment_text = meta_info.get("native_duplex_segment_text", "") if native_duplex else "" - if not isinstance(segment_text, str): - segment_text = "" - turn_eos_id = meta_info.get("turn_eos_token_id") - ids_info = info_dict.get("ids") - tts_ids = ids_info.get("tts") if native_duplex and isinstance(ids_info, dict) else None - if isinstance(tts_ids, torch.Tensor): - contains_turn_eos = isinstance(turn_eos_id, int) and bool( - torch.any(tts_ids.reshape(-1) == turn_eos_id).item() - ) - elif isinstance(tts_ids, (list, tuple)): - contains_turn_eos = isinstance(turn_eos_id, int) and turn_eos_id in tts_ids - else: - contains_turn_eos = False + native_duplex, epoch, turn_id, segment_text, turn_end = _native_duplex_row_meta(info_dict) native_duplex_flags.append(torch.tensor(native_duplex, dtype=torch.bool)) - duplex_epochs.append(torch.tensor(epoch if isinstance(epoch, int) else -1, dtype=torch.long)) - duplex_turn_ids.append(torch.tensor(turn_id if isinstance(turn_id, int) else -1, dtype=torch.long)) - segment_texts_utf8.append( - torch.tensor( - list(segment_text.encode("utf-8")), - dtype=torch.uint8, - ) - ) - turn_end_flags.append(torch.tensor(native_duplex and contains_turn_eos, dtype=torch.bool)) + duplex_epochs.append(torch.tensor(epoch, dtype=torch.long)) + duplex_turn_ids.append(torch.tensor(turn_id, dtype=torch.long)) + segment_texts_utf8.append(torch.tensor(list(segment_text.encode("utf-8")), dtype=torch.uint8)) + turn_end_flags.append(torch.tensor(turn_end, dtype=torch.bool)) if not isinstance(info, dict): continue diff --git a/vllm_omni/model_executor/models/minicpmo_4_5/pipeline.py b/vllm_omni/model_executor/models/minicpmo_4_5/pipeline.py index cee8d7b9019..3273dd16f6f 100644 --- a/vllm_omni/model_executor/models/minicpmo_4_5/pipeline.py +++ b/vllm_omni/model_executor/models/minicpmo_4_5/pipeline.py @@ -48,6 +48,7 @@ requires_multimodal_data=True, engine_output_type="latent", supports_native_mrv2_data_plane=True, + supports_duplex_mrv2=True, sampling_constraints={ "detokenize": True, # The llm2tts bridge discards this boundary and every row @@ -66,9 +67,8 @@ custom_process_input_func=f"{_PROC}.llm2tts", custom_process_next_stage_input_func=f"{_PROC}.tts2code2wav_full_payload", async_chunk_process_next_stage_input_func=f"{_PROC}.tts2code2wav_async_chunk", - # Takes effect only when the deploy selects model_runner v2 for - # this stage (turn sessions only; duplex stays on V1). supports_native_mrv2_data_plane=True, + supports_duplex_mrv2=True, sampling_constraints={ "detokenize": False, # MiniCPM-o 4.5 codec EOS is tts_config.num_audio_tokens - 1. @@ -86,6 +86,7 @@ engine_output_type="audio", model_arch="MiniCPMO45Code2Wav", supports_native_mrv2_data_plane=True, + supports_duplex_mrv2=True, sync_process_input_func=f"{_PROC}.tts2code2wav_token_only", # Sends the reference audio with the async-chunk placeholder so # Code2Wav can prepare it while it waits for chunk 0. diff --git a/vllm_omni/model_executor/stage_input_processors/minicpmo_4_5_omni.py b/vllm_omni/model_executor/stage_input_processors/minicpmo_4_5_omni.py index 5488036b7c7..2f8273ad3b2 100644 --- a/vllm_omni/model_executor/stage_input_processors/minicpmo_4_5_omni.py +++ b/vllm_omni/model_executor/stage_input_processors/minicpmo_4_5_omni.py @@ -163,6 +163,11 @@ def _to_transport_list(value): def _coerce_int(value): + # Output accumulation can concatenate repeated scalar metadata snapshots. + while isinstance(value, (list, tuple)): + if len(value) == 0: + return None + value = value[0] if hasattr(value, "detach"): flat = value.detach().cpu().reshape(-1) if flat.numel() == 0: @@ -270,6 +275,10 @@ def _drop_codec_state(transfer_manager: Any, request_id: str) -> None: def _is_aborted(request: Any) -> bool: + # V1 passes the scheduler request; the MRv2 native transport passes a + # snapshot that carries no status and marks a cancelled terminal instead. + if getattr(request, "aborted", False) is True: + return True status_name = getattr(getattr(request, "status", None), "name", "") return any(marker in status_name for marker in ("ABORT", "CANCEL", "IGNORED", "ERROR")) @@ -291,6 +300,8 @@ def tts2code2wav_async_chunk( duplex_epoch = _coerce_int(output_meta.get("duplex_epoch")) duplex_turn_id = _coerce_int(output_meta.get("duplex_turn_id")) segment_text_utf8 = output_meta.get("llm_output_text_utf8") + if isinstance(segment_text_utf8, (list, tuple)) and len(segment_text_utf8) > 0: + segment_text_utf8 = segment_text_utf8[0] if not isinstance(segment_text_utf8, torch.Tensor): segment_text_utf8 = None turn_end = bool(_coerce_int(output_meta.get("turn_end"))) @@ -323,6 +334,7 @@ def tts2code2wav_async_chunk( record["cache_epoch"] = int(record["cache_epoch"]) + 1 record["chunk_seq"] = 0 record["last_terminal_turn"] = None + record.pop("last_closed_condition", None) _drop_codec_state(transfer_manager, request_id) if _is_aborted(request): @@ -330,6 +342,17 @@ def tts2code2wav_async_chunk( _drop_codec_state(transfer_manager, request_id) return None + # Async scheduling can publish lookahead outputs after this condition's + # EOS, including after the next condition starts. Fence by the identity + # captured with the output, not the request's mutable current metadata. + condition_seq = _coerce_int(output_meta.get("streaming_condition_seq")) + condition_key = None + if native_duplex and all(isinstance(value, int) and value >= 0 for value in (*duplex_turn_key, condition_seq)): + condition_key = (*duplex_turn_key, condition_seq) + closed = record.get("last_closed_condition") + if closed is not None and condition_key <= closed: + return None + if native_duplex and turn_end and record.get("last_terminal_turn") == duplex_turn_key: # Emit an empty replacement snapshot so Code2Wav cannot replay the # prior terminal audio when this control-only boundary arrives. @@ -369,6 +392,19 @@ def tts2code2wav_async_chunk( state["segment_text_recorded"] = True request_finished = getattr(request, "is_finished", None) finished = bool(is_finished or (callable(request_finished) and request_finished())) + # A native-duplex turn end closes one Code2Wav turn, not the resumable + # request. meta.finished is the whole-stream terminal on the MRv2 native + # transport (the V1 chunk adapter overwrites it with the scheduler's + # request finish), so it must follow the request, not last_chunk. + request_terminal = finished + # MRv2 materializes sampled ids before publishing this payload. A + # resumable request is still alive at codec EOS; flush the segment using + # its confirmed host-side token rather than closing on every turn_end row. + if native_duplex: + stop_ids = getattr(getattr(request, "sampling_params", None), "stop_token_ids", ()) or () + finished = finished or any(token in stop_ids for token in getattr(request, "sampled_token_ids", ())) + if finished and condition_key is not None: + record["last_closed_condition"] = condition_key chunk_frames, left_context_frames = _codec_config(transfer_manager) flush_pending = finished last_chunk = bool(flush_pending and (not native_duplex or turn_end)) @@ -443,7 +479,7 @@ def tts2code2wav_async_chunk( left_context_size=len(context), last_chunk=last_chunk, stream_finished=finished_tensor, - finished=finished_tensor, + finished=torch.tensor(request_terminal, dtype=torch.bool) if native_duplex else finished_tensor, is_segment_finished=finished_tensor, req_id=[request_id], duplex_epoch=duplex_epoch, @@ -1151,3 +1187,50 @@ def llm2tts( duplex_state["model_turn_id"] = current_model_turn_id + 1 return tts_inputs + + +def _update_native_talker_prompt(model_config: Any, session: Any, payload: dict[str, Any]) -> bool | None: + """Apply the existing Talker window recipe without a V1 transfer adapter. + + Sender-only MRv2 stages receive conditions through StreamingUpdate. + The request already owns the previous condition metadata and confirmed + codec ledger, so no parallel request registry is needed. + """ + from vllm_omni.distributed.omni_connectors.adapter import construct_next_stage_streaming_input_prompt + from vllm_omni.distributed.omni_connectors.transfer_adapter.chunk_transfer_adapter import ( + _resolve_talker_streaming_prompt_config, + ) + + limit, previous_chunks, on_capacity = _resolve_talker_streaming_prompt_config(model_config) + if previous_chunks != 1 or payload.get("native_duplex") is not True: + return None + previous_info = getattr(session, "model_intermediate_buffer", None) or {} + previous = previous_info.get("meta", {}) + meta = payload["meta"] + previous_seq, seq = previous.get("streaming_condition_seq"), meta.get("streaming_condition_seq") + if ( + not isinstance(previous_seq, int) + or isinstance(previous_seq, bool) + or not isinstance(seq, int) + or isinstance(seq, bool) + or seq != previous_seq + 1 + ): + raise ValueError("native Talker streaming_condition_seq must advance by one") + # Carry only confirmed codec ids, never the uncomputed sampled EOS. + # This also seeds the MRv2 16-code penalty after a new condition. + payload.setdefault("ids", {})["streaming_prompt_previous_codes"] = list( + session._all_token_ids[session.num_prompt_tokens : session.num_computed_tokens] + ) + return construct_next_stage_streaming_input_prompt( + payload, + session, + max_model_len=limit, + previous_condition_len=previous.get("next_stage_prompt_len"), + previous_condition_seq=previous_seq, + condition_seq=seq, + recompute_previous_chunks=previous_chunks, + recompute_on_capacity=on_capacity, + ) + + +setattr(tts2code2wav_async_chunk, "update_streaming_prompt_for_condition", _update_native_talker_prompt) diff --git a/vllm_omni/worker_v2/model_states/omni_model_state.py b/vllm_omni/worker_v2/model_states/omni_model_state.py index 0cb0f04a3e9..6622d269ac7 100644 --- a/vllm_omni/worker_v2/model_states/omni_model_state.py +++ b/vllm_omni/worker_v2/model_states/omni_model_state.py @@ -101,6 +101,7 @@ class OmniModelState(DefaultModelState): _eager_mtp = False _eager_fastpath = False _eager_rows: tuple[InputBatch, list[tuple[int, int, str, bool]], torch.Tensor] | None = None + preprocess_keeps_mm_inputs = False _first_audio_stream: torch.cuda.Stream | None = None # Set by the stage engine process when it can take one-request outputs directly. _first_audio_sender: Any = None @@ -155,9 +156,24 @@ def __init__( # Talker's codec_embedding dim may differ from hf_text_config.hidden_size; probe real dim. self._embed_dim = self._get_embed_dim(model, device) if self.has_preprocess else 0 + # A preprocess model may also take native multimodal inputs (MiniCPM-o's + # duplex Thinker serves media chat). The encoder runner rewrites every + # active row of its buffer each step, so that buffer is the static one. + encoder_runner = getattr(self, "encoder_runner", None) + encoder_embeds = getattr(encoder_runner, "inputs_embeds", None) + self.preprocess_keeps_mm_inputs = bool( + self.has_preprocess + and self.supports_mm_inputs + and getattr(model, "preprocess_keeps_mm_inputs", False) + and isinstance(encoder_embeds, torch.Tensor) + and tuple(encoder_embeds.shape) == (self.max_num_tokens, self._embed_dim) + ) + # Static inputs_embeds buffer for FULL CUDA graph — preprocess fills it in-place each step. self._static_inputs_embeds: torch.Tensor | None = None - if self._embed_dim > 0: + if self.preprocess_keeps_mm_inputs: + self._static_inputs_embeds = encoder_embeds + elif self._embed_dim > 0: self._static_inputs_embeds = torch.zeros( (self.max_num_tokens, self._embed_dim), dtype=self.dtype, @@ -722,7 +738,11 @@ def run_preprocess( if embeds is None: embeds = self.model.embed_input_ids(input_batch.input_ids[: input_batch.num_tokens]) model_inputs["inputs_embeds"] = embeds - elif self._static_inputs_embeds is not None and embeds.data_ptr() == self._static_inputs_embeds.data_ptr(): + elif ( + not self.preprocess_keeps_mm_inputs + and self._static_inputs_embeds is not None + and embeds.data_ptr() == self._static_inputs_embeds.data_ptr() + ): # FULL graph replay requires a stable inputs_embeds address. Refresh # the active rows from the current token ids before model-specific # preprocessing; otherwise decode reuses embeddings left by the diff --git a/vllm_omni/worker_v2/omni_data_plane.py b/vllm_omni/worker_v2/omni_data_plane.py index d82e913dc32..ad767e96e9c 100644 --- a/vllm_omni/worker_v2/omni_data_plane.py +++ b/vllm_omni/worker_v2/omni_data_plane.py @@ -33,12 +33,17 @@ class _NativeRequestState: external_req_id: str prompt_token_ids: list[int] additional_information: Any = None + model_intermediate_buffer: Any = None sampling_params: Any = None num_computed_tokens: int = 0 resumable: bool = False output_token_ids: list[int] = field(default_factory=list) finished: bool = False output_stopped: bool = False + # Set by ``abort_requests``: the scheduler ended the request outside its + # own output (cancel, error or a parked close), so the runner has no + # finish reason. MiniCPM-o duplex Talkers only end this way. + aborted: bool = False def accept_tokens(self, token_ids: list[int]) -> None: """Fence publications past the sampled stop, before scheduler ACK. @@ -63,7 +68,7 @@ def accept_tokens(self, token_ids: list[int]) -> None: self.output_stopped = True break - def snapshot(self, *, include_token_history: bool) -> SimpleNamespace: + def snapshot(self, *, include_token_history: bool, sampled_token_ids: list[int] | None = None) -> SimpleNamespace: prompt = list(self.prompt_token_ids) if include_token_history else [] output = list(self.output_token_ids) if include_token_history else [] finished = self.finished @@ -76,10 +81,15 @@ def snapshot(self, *, include_token_history: bool) -> SimpleNamespace: all_token_ids=prompt + output, output_token_count=len(self.output_token_ids), last_output_token_id=self.output_token_ids[-1] if self.output_token_ids else None, + # A ledger tail can belong to an earlier segment. Only these ids + # were accepted with the payload being published now. + sampled_token_ids=list(sampled_token_ids or ()), additional_information=self.additional_information, + model_intermediate_buffer=self.model_intermediate_buffer, sampling_params=self.sampling_params, num_computed_tokens=self.num_computed_tokens, resumable=self.resumable, + aborted=self.aborted, ) request.is_finished = lambda: finished return request @@ -276,6 +286,7 @@ def register_request(self, request_data: Any) -> None: external_req_id=external_req_id, prompt_token_ids=list(getattr(request_data, "prompt_token_ids", None) or []), additional_information=getattr(request_data, "additional_information", None), + model_intermediate_buffer=getattr(request_data, "model_intermediate_buffer", None), sampling_params=getattr(request_data, "sampling_params", None), num_computed_tokens=int(getattr(request_data, "num_computed_tokens", 0) or 0), resumable=bool(getattr(request_data, "resumable", False)), @@ -327,6 +338,7 @@ def abort_requests(self, req_ids: set[str]) -> int: if not active_req_ids: return 0 for req_id in active_req_ids: + self._native_requests[req_id].aborted = True self._native_outputs_in_flight.pop(req_id, None) self._native_terminal_pending.discard(req_id) # Cancellation must not flush a partially accumulated utterance. @@ -558,7 +570,10 @@ def _build_chunk_entries( ) entries.append( ( - state.snapshot(include_token_history=include_token_history), + state.snapshot( + include_token_history=include_token_history, + sampled_token_ids=sampled_by_req.get(req_id) if not state.finished else None, + ), payload, ) ) diff --git a/vllm_omni/worker_v2/omni_model_runner.py b/vllm_omni/worker_v2/omni_model_runner.py index c07c47f899f..4c602971dcb 100644 --- a/vllm_omni/worker_v2/omni_model_runner.py +++ b/vllm_omni/worker_v2/omni_model_runner.py @@ -359,8 +359,13 @@ def _init_model_state(vllm_config: Any, model: Any, encoder_cache: Any, device: self._last_multimodal_outputs = None self._configure_cudagraph_output_contract() - # Preprocess models own embedding buffers; encoder_runner sizing would mismatch. - if getattr(self.model, "has_preprocess", False) and self.supports_mm_inputs: + # Preprocess models own embedding buffers; encoder_runner sizing would + # mismatch unless the model builds its preprocess on the encoder output. + if ( + getattr(self.model, "has_preprocess", False) + and self.supports_mm_inputs + and not getattr(self.model_state, "preprocess_keeps_mm_inputs", False) + ): self.supports_mm_inputs = False self.encoder_cache = None @@ -509,6 +514,9 @@ def execute_model( self.update_requests(scheduler_output) self._sync_native_data_plane_payloads(scheduler_output) self.block_tables.apply_staged_writes() + reanchor = getattr(self.model, "apply_duplex_kv_reanchor", None) + if callable(reanchor): + reanchor(self, scheduler_output=scheduler_output) if self.aux_output_connector is not None: self.aux_output_connector.begin_step(scheduler_output.aux_output_connector_metadata) if scheduler_output.total_num_scheduled_tokens == 0: