Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
66 changes: 59 additions & 7 deletions tests/worker/test_omni_connector_mixin.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project
"""Unit tests for OmniConnectorModelRunnerMixin.

These tests use a mock connector (in-memory dict store) and do not require
Expand Down Expand Up @@ -376,6 +376,19 @@ def test_finished_load_reqs_flow_to_chunk_ready(self):


class TestLoadCustomFuncSelection(unittest.TestCase):
def test_uses_validator_override_from_public_mixin(self):
config = SimpleNamespace(
async_chunk=True,
custom_process_next_stage_input_func=f"{__name__}._make_request",
)

with patch.object(MixinHost, "_is_connector_payload_builder", return_value=False) as validator:
selected_path, func = MixinHost._load_custom_func(config)

validator.assert_called_once_with(_make_request)
assert selected_path is None
assert func is None

def test_skips_non_payload_stage_input_processors_for_full_payload_mode(self):
incompatible_paths = [
"vllm_omni.model_executor.stage_input_processors.mimo_audio.llm2code2wav",
Expand Down Expand Up @@ -861,7 +874,7 @@ def test_rank0_only_polls_connector_for_tp_full_payload(self):
host._omni_connector.get.return_value = connector_result
tp_group = _FakeTPGroup(world_size=2, rank_in_group=0)

with patch("vllm_omni.worker.omni_connector_model_runner_mixin.get_tp_group", return_value=tp_group):
with patch.object(host, "_get_local_tp_group", return_value=tp_group):
made_progress = host._poll_single_request("r1")

self.assertTrue(made_progress)
Expand All @@ -882,7 +895,7 @@ def test_tp_follower_skips_connector_poll_for_full_payload(self):
host._get_req_chunk["r1"] = 0
tp_group = _FakeTPGroup(world_size=2, rank_in_group=1)

with patch("vllm_omni.worker.omni_connector_model_runner_mixin.get_tp_group", return_value=tp_group):
with patch.object(host, "_get_local_tp_group", return_value=tp_group):
made_progress = host._poll_single_request("r1")

self.assertFalse(made_progress)
Expand All @@ -900,7 +913,7 @@ def test_recv_full_payload_inputs_broadcasts_tp_leader_results_to_followers(self
payload = {"tok": [10], "finished": torch.tensor(True)}
tp_group = _FakeTPGroup(world_size=2, rank_in_group=1, follower_result={"r1": payload})

with patch("vllm_omni.worker.omni_connector_model_runner_mixin.get_tp_group", return_value=tp_group):
with patch.object(host, "_get_local_tp_group", return_value=tp_group):
results = host.recv_full_payload_inputs(scheduler_output=None)

self.assertEqual(results, {"r1": payload})
Expand Down Expand Up @@ -936,7 +949,7 @@ def test_rank0_only_polls_connector_for_tp_async_chunk(self):
host._omni_connector.get.return_value = (payload, 123)
tp_group = _FakeTPGroup(world_size=2, rank_in_group=0)

with patch("vllm_omni.worker.omni_connector_model_runner_mixin.get_tp_group", return_value=tp_group):
with patch.object(host, "_get_local_tp_group", return_value=tp_group):
made_progress = host._poll_single_request("r1")

self.assertTrue(made_progress)
Expand All @@ -951,7 +964,7 @@ def test_tp_follower_skips_connector_poll_for_async_chunk(self):
host = self._make_host(rank=1)
tp_group = _FakeTPGroup(world_size=2, rank_in_group=1)

with patch("vllm_omni.worker.omni_connector_model_runner_mixin.get_tp_group", return_value=tp_group):
with patch.object(host, "_get_local_tp_group", return_value=tp_group):
made_progress = host._poll_single_request("r1")

self.assertFalse(made_progress)
Expand All @@ -976,7 +989,7 @@ def test_get_output_broadcasts_tp_async_chunk_payloads_to_followers(self):
}
tp_group = _FakeTPGroup(world_size=2, rank_in_group=1, follower_result=packet)

with patch("vllm_omni.worker.omni_connector_model_runner_mixin.get_tp_group", return_value=tp_group):
with patch.object(host, "_get_local_tp_group", return_value=tp_group):
output = host.get_omni_connector_output()

self.assertEqual(output.chunk_ready_req_ids, {"r1"})
Expand Down Expand Up @@ -1049,6 +1062,45 @@ def test_cleanup_removes_kv_state(self):
class TestAsyncPayloadLifecycle(unittest.TestCase):
"""Regression tests for async payload delivery lifecycle."""

def test_accumulate_payload_concatenates_chunks(self):
host = MixinHost()
host._send_side_request_payload = {}
first = host._accumulate_payload(
"r1",
{
"embed": {"decode": torch.tensor([[1.0, 2.0]])},
"ids": {"output": [1]},
"meta": {"finished": False},
},
)
merged = host._accumulate_payload(
"r1",
{
"embed": {"decode": torch.tensor([[3.0, 4.0], [5.0, 6.0]])},
"ids": {"output": [2, 3]},
"meta": {"finished": True},
},
)
torch.testing.assert_close(merged["embed"]["decode"], torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]))
self.assertEqual(merged["ids"]["output"], [1, 2, 3])
self.assertIs(merged["meta"]["finished"], True)
self.assertEqual(first["embed"]["decode"].shape, (1, 2))
self.assertEqual(first["ids"]["output"], [1])
self.assertIs(first["meta"]["finished"], False)

def test_accumulate_payload_replaces_override_keys(self):
host = MixinHost()
host._send_side_request_payload = {}
host._accumulate_payload("r1", {"embed": {"decode": torch.ones(2, 2)}, "ids": {"output": [1, 2]}})
payload = {
"embed": {"decode": torch.zeros(1, 2)},
"ids": {"output": [3]},
"meta": {"override_keys": [["embed", "decode"], ["ids", "output"]]},
}
merged = host._accumulate_payload("r1", payload)
torch.testing.assert_close(merged["embed"]["decode"], payload["embed"]["decode"])
self.assertEqual(merged["ids"]["output"], [3])

def test_send_side_request_payload_not_cleared_before_payload_is_consumable(self):
host = MixinHost()
host.init_omni_connectors(
Expand Down
166 changes: 0 additions & 166 deletions tests/worker/test_payload_span.py

This file was deleted.

Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project
"""Internal model-runner components for Omni connector transport."""
Loading
Loading