fix(vlm): harden EPD receiver validation and liveness (#36945)

This commit is contained in:
Mick
2026-09-05 21:22:37 +08:00
committed by GitHub
parent 4b802c052b
commit 5df60a21cd
12 changed files with 1320 additions and 109 deletions
@@ -1,11 +1,25 @@
"""Unit tests for request construction in the encode-disaggregation path."""
"""Unit tests for the encode-disaggregation receiver."""
import asyncio
import threading
import time
import unittest
from array import array
from http import HTTPStatus
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.disaggregation.encoder.receiver import MMReceiverBase
from sglang.srt.disaggregation.encoder.receiver import (
MMReceiverBase,
WaitingMMRequestStatus,
WaitingRDMARequest,
WaitingZmqRequest,
WaitingZmqRequestGrpc,
_ReceiveRegistrationRunner,
)
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.io_struct import EncoderDispatchErrorReq
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -13,7 +27,239 @@ from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
def _make_registration_request(request_cls):
request = request_cls.__new__(request_cls)
request.rid = "registration-test"
request.registration_runner = _ReceiveRegistrationRunner(
"test-encoder-receive-registration"
)
request.registration_future = None
request.registration_error = None
request.registration_lock = threading.Lock()
request.status = WaitingMMRequestStatus.PENDING
request.error_msg = None
request.error_code = None
request.embedding_pool = None
request.embeddings_buffer = None
request.recv_embedding_data = None
request._pool_slot_id = None
request._mm_finalizer = None
request.recv_socket = None
request.recv_req = SimpleNamespace(rid=request.rid)
request.num_items_assigned = {Modality.IMAGE: [1]}
request.encoder_urls = ["http://encoder"]
request.host_name = "127.0.0.1"
request.receive_count = 1
request.embedding_port = 12345
return request
def _cancel_registration(request):
future = request.registration_future
request.release_resources()
deadline = time.monotonic() + 1
while not future.done() and time.monotonic() < deadline:
time.sleep(0.01)
assert future.cancelled()
class BlockingResponse:
def __init__(self, started):
self.started = started
async def __aenter__(self):
self.started.set()
await asyncio.Event().wait()
async def __aexit__(self, *args):
return False
class BlockingSession:
started = None
def __init__(self, *args, **kwargs):
pass
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
def post(self, *args, **kwargs):
return BlockingResponse(self.started)
class FailingResponse:
async def __aenter__(self):
raise ConnectionError("encoder unavailable")
async def __aexit__(self, *args):
return False
class FailingSession(BlockingSession):
def post(self, *args, **kwargs):
return FailingResponse()
class TestReceiveRegistration(CustomTestCase):
def test_http_registration_does_not_block_scheduler(self):
started = threading.Event()
BlockingSession.started = started
request = _make_registration_request(WaitingZmqRequest)
with patch(
"sglang.srt.disaggregation.encoder.receiver.aiohttp.ClientSession",
BlockingSession,
):
scheduler_call = threading.Thread(
target=request.send_encode_request, daemon=True
)
scheduler_call.start()
self.assertTrue(started.wait(timeout=1))
scheduler_call.join(timeout=0.1)
self.assertFalse(scheduler_call.is_alive())
self.assertEqual(request.status, WaitingMMRequestStatus.PENDING)
_cancel_registration(request)
def test_grpc_registration_does_not_block_scheduler(self):
started = threading.Event()
async def blocking_registration(*args, **kwargs):
started.set()
await asyncio.Event().wait()
request = _make_registration_request(WaitingZmqRequestGrpc)
with patch(
"sglang.srt.disaggregation.encoder.receiver._grpc_scheduler_receive_url",
blocking_registration,
):
scheduler_call = threading.Thread(
target=request.send_encode_request, daemon=True
)
scheduler_call.start()
self.assertTrue(started.wait(timeout=1))
scheduler_call.join(timeout=0.1)
self.assertFalse(scheduler_call.is_alive())
self.assertEqual(request.status, WaitingMMRequestStatus.PENDING)
_cancel_registration(request)
def test_failure_is_request_local(self):
request = _make_registration_request(WaitingZmqRequest)
with patch(
"sglang.srt.disaggregation.encoder.receiver.aiohttp.ClientSession",
FailingSession,
):
request.send_encode_request()
deadline = time.monotonic() + 1
while request.status == WaitingMMRequestStatus.PENDING:
self.assertLess(time.monotonic(), deadline)
request._try_recv_mm_data()
time.sleep(0.01)
self.assertEqual(request.status, WaitingMMRequestStatus.FAIL)
self.assertEqual(request.error_code, HTTPStatus.BAD_GATEWAY)
self.assertIn("encoder unavailable", request.error_msg)
class TestEncodeReceiverRequestConstruction(CustomTestCase):
def test_early_dispatch_error_waits_for_scheduler_request(self):
encode_finished = threading.Event()
scheduler_dispatch_ready = threading.Event()
reported = []
failure = EncoderDispatchErrorReq(
rid="request-1",
error_msg="encoder unavailable",
error_code=HTTPStatus.BAD_GATEWAY,
)
async def fail_encode(**kwargs):
encode_finished.set()
return failure
receiver = SimpleNamespace(encode=fail_encode)
worker = threading.Thread(
target=MMReceiverBase._run_encode_in_thread,
args=(
receiver,
failure.rid,
[],
"encode",
{},
[],
None,
scheduler_dispatch_ready,
reported.append,
),
)
worker.start()
self.assertTrue(encode_finished.wait(timeout=1))
worker.join(timeout=0.05)
self.assertTrue(worker.is_alive())
self.assertEqual(reported, [])
scheduler_dispatch_ready.set()
worker.join(timeout=1)
self.assertFalse(worker.is_alive())
self.assertEqual(reported, [failure])
def test_dispatch_error_fails_only_owning_wait(self):
class WaitingRequest:
def __init__(self, rid):
self.rid = rid
self.recv_req = SimpleNamespace(rid=rid)
self.status = WaitingMMRequestStatus.PENDING
self.error_msg = None
self.error_code = None
self.start_time = 0
def _try_recv_mm_data(self):
pass
def _fail_and_release(self, error_msg, error_code=None):
self.error_msg = error_msg
self.error_code = error_code
self.status = WaitingMMRequestStatus.FAIL
def release_resources(self):
pass
def close_recv_socket(self):
pass
owner = WaitingRequest("request-1")
other = WaitingRequest("request-2")
receiver = SimpleNamespace(
waiting_list=[owner, other],
waiting_by_rid={owner.rid: owner, other.rid: other},
scheduler_recv_socket=None,
wait_timeout=float("inf"),
tp_group=SimpleNamespace(cpu_group=object()),
_drain_scheduler_embeddings=lambda: None,
_sync_fail_info_across_tp=lambda request: None,
create_req=lambda request: request,
)
dispatch_error = EncoderDispatchErrorReq(
rid=owner.rid,
error_msg="bad media",
error_code=HTTPStatus.UNPROCESSABLE_ENTITY,
)
with patch("torch.distributed.all_reduce"):
_, abort_reqs = MMReceiverBase._process_waiting_requests(
receiver, [dispatch_error], waiting_cls=None
)
self.assertEqual(owner.status, WaitingMMRequestStatus.FAIL)
self.assertEqual(owner.error_msg, dispatch_error.error_msg)
self.assertEqual(owner.error_code, dispatch_error.error_code)
self.assertEqual(other.status, WaitingMMRequestStatus.PENDING)
self.assertEqual([req.rid for req, _, _ in abort_reqs], [owner.rid])
def test_extra_key_and_cache_salt_are_forwarded(self):
scheduler = SimpleNamespace(
model_config=SimpleNamespace(hf_eos_token_id={2}, vocab_size=128),
@@ -56,6 +302,98 @@ class TestEncodeReceiverRequestConstruction(CustomTestCase):
self.assertEqual(req.extra_key, "classification")
self.assertEqual(req.cache_salt, "tenant-a")
def test_rdma_worker_error_is_released_on_scheduler_thread(self):
scheduler_thread = threading.get_ident()
class ThreadCheckedSocket:
closed_by = None
def close(self):
self.closed_by = threading.get_ident()
recv_socket = ThreadCheckedSocket()
request = WaitingRDMARequest.__new__(WaitingRDMARequest)
request.rid = "request-1"
request.status = WaitingMMRequestStatus.PENDING
request.error_msg = None
request.error_code = None
request.recv_socket = recv_socket
request._receive_error = None
request._receive_error_lock = threading.Lock()
request._buffer_lock = threading.Lock()
request._terminal = False
request._receive_running = False
request.registration_future = None
request.embeddings_buffer = None
request._pool_slot_id = None
request.embedding_pool = None
request._mm_finalizer = None
worker = threading.Thread(
target=lambda: asyncio.run(
request._check_encoder_responses(
[ConnectionError("encoder unavailable")], "/send"
)
)
)
worker.start()
worker.join(timeout=1)
self.assertFalse(worker.is_alive())
self.assertEqual(request.status, WaitingMMRequestStatus.PENDING)
self.assertIsNone(request.recv_socket.closed_by)
request._try_recv_mm_data()
self.assertEqual(request.status, WaitingMMRequestStatus.FAIL)
self.assertIsNone(request.recv_socket)
self.assertTrue(request._terminal)
self.assertEqual(recv_socket.closed_by, scheduler_thread)
def test_tp_peer_failure_closes_local_receive_socket(self):
class WaitingRequest:
rid = "request-1"
recv_req = SimpleNamespace(rid=rid)
status = WaitingMMRequestStatus.PENDING
error_msg = "peer failed"
error_code = None
start_time = 0
released = False
closed = False
def _try_recv_mm_data(self):
pass
def release_resources(self):
self.released = True
def close_recv_socket(self):
self.closed = True
waiting_req = WaitingRequest()
receiver = SimpleNamespace(
waiting_list=[waiting_req],
waiting_by_rid={waiting_req.rid: waiting_req},
scheduler_recv_socket=None,
wait_timeout=float("inf"),
tp_group=SimpleNamespace(cpu_group=object()),
_drain_scheduler_embeddings=lambda: None,
_sync_fail_info_across_tp=lambda request: None,
create_req=lambda request: request,
)
def force_peer_failure(status, **kwargs):
status.fill_(WaitingMMRequestStatus.FAIL)
with patch("torch.distributed.all_reduce", force_peer_failure):
_, abort_reqs = MMReceiverBase._process_waiting_requests(
receiver, [], waiting_cls=None
)
self.assertTrue(waiting_req.released)
self.assertTrue(waiting_req.closed)
self.assertEqual(len(abort_reqs), 1)
if __name__ == "__main__":
unittest.main()
@@ -7,7 +7,7 @@ import time
from array import array
from concurrent.futures import ThreadPoolExecutor
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import pytest
import torch
@@ -23,8 +23,11 @@ from sglang.srt.disaggregation.encoder.preprocessor import (
)
from sglang.srt.disaggregation.encoder.receiver import (
EmbeddingData,
MMReceiverGrpc,
MMReceiverHTTP,
MultiModalEmbeddingData,
WaitingMMRequestStatus,
WaitingZmqRequest,
_encoder_media_item,
_select_mm_processor_prompt,
)
@@ -577,6 +580,102 @@ def test_epd_receiver_keeps_content_hash_aligned_with_image():
}
def test_epd_tokenizer_receiver_timeout_cancels_tasks_and_closes_socket():
async def run():
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
receiver.encode_urls = ["http://encoder"]
receiver.context = object()
receiver.host = "127.0.0.1"
receiver.recv_timeout = 0.01
receiver._extract_url_data = Mock(return_value=[{"modality": Modality.IMAGE}])
encode_cancelled = asyncio.Event()
recv_cancelled = asyncio.Event()
async def wait_until_cancelled(event, *_args, **_kwargs):
try:
await asyncio.Event().wait()
finally:
event.set()
receiver.encode = lambda *args, **kwargs: wait_until_cancelled(
encode_cancelled, *args, **kwargs
)
receiver._recv_mm_data = lambda *args, **kwargs: wait_until_cancelled(
recv_cancelled, *args, **kwargs
)
recv_socket = SimpleNamespace(close=Mock())
with patch(
"sglang.srt.disaggregation.encoder.receiver.get_zmq_socket_on_host",
return_value=(12345, recv_socket),
):
result = await receiver.recv_mm_data(
SimpleNamespace(),
mm_processor=object(),
prompt="prompt",
)
assert result is None
assert encode_cancelled.is_set()
assert recv_cancelled.is_set()
recv_socket.close.assert_called_once_with(linger=0)
asyncio.run(run())
def test_grpc_dispatch_cancellation_waits_for_blocking_calls():
async def run():
receiver = MMReceiverGrpc.__new__(MMReceiverGrpc)
receiver.host = "127.0.0.1"
calls_started = 0
calls_finished = 0
calls_lock = threading.Lock()
unblock = threading.Event()
def blocking_encode(_target, _request):
nonlocal calls_started, calls_finished
with calls_lock:
calls_started += 1
unblock.wait(timeout=2)
with calls_lock:
calls_finished += 1
with patch(
"sglang.srt.disaggregation.encoder.receiver._grpc_encode_request",
side_effect=blocking_encode,
):
task = asyncio.create_task(
receiver.encode(
req_id="req",
mm_data=[
{"modality": Modality.IMAGE, "url": "image-0"},
{"modality": Modality.IMAGE, "url": "image-1"},
],
embedding_port=1234,
endpoint_encode="encode",
num_items_assigned=[1, 1],
encode_urls=["grpc://encoder-0", "grpc://encoder-1"],
)
)
for _ in range(100):
with calls_lock:
if calls_started == 2:
break
await asyncio.sleep(0.01)
assert calls_started == 2
task.cancel()
await asyncio.sleep(0)
assert not task.done()
unblock.set()
with pytest.raises(asyncio.CancelledError):
await task
assert calls_finished == 2
asyncio.run(run())
def test_kimi_k3_epd_aggregates_original_image_sizes_in_part_order():
first = EmbeddingData(
req_id="request",
@@ -607,6 +706,109 @@ def test_kimi_k3_epd_aggregates_original_image_sizes_in_part_order():
]
@pytest.mark.parametrize(
("num_parts", "part_idx", "error"),
[
(0, 0, "num_parts must be a positive integer"),
(2, -1, "part_idx must be in"),
(2, 2, "part_idx must be in"),
],
)
def test_epd_embedding_aggregation_rejects_invalid_part_metadata(
num_parts, part_idx, error
):
part = EmbeddingData(
req_id="request",
num_parts=num_parts,
part_idx=part_idx,
grid_dim=torch.tensor([[1, 2, 2]]),
modality=Modality.IMAGE,
embedding=torch.ones(1, 2),
)
with pytest.raises(ValueError, match=error):
MultiModalEmbeddingData.from_embedding_data(part)
def test_epd_embedding_aggregation_rejects_duplicate_and_inconsistent_parts():
def make_part(num_parts, part_idx):
return EmbeddingData(
req_id="request",
num_parts=num_parts,
part_idx=part_idx,
grid_dim=torch.tensor([[1, 2, 2]]),
modality=Modality.IMAGE,
embedding=torch.ones(1, 2),
)
combined = MultiModalEmbeddingData.from_embedding_data(make_part(2, 0))
with pytest.raises(ValueError, match="duplicate embedding part 0"):
combined.add(make_part(2, 0))
with pytest.raises(ValueError, match="num_parts changed from 2 to 3"):
combined.add(make_part(3, 1))
def test_epd_scheduler_contains_invalid_embedding_part_metadata():
waiting = WaitingZmqRequest.__new__(WaitingZmqRequest)
waiting.rid = "request"
waiting.recv_req = SimpleNamespace(rid="request")
waiting.status = WaitingMMRequestStatus.PENDING
waiting.recv_embedding_data = None
waiting.model_type = None
waiting._fail_and_release = Mock()
invalid = EmbeddingData(
req_id="request_local_part_2",
num_parts=2,
part_idx=2,
grid_dim=None,
modality=Modality.IMAGE,
embedding=torch.ones(1, 2),
)
waiting.consume_parts(
[pickle.dumps(invalid.copy_without_embedding()), invalid.embedding.numpy()]
)
waiting._fail_and_release.assert_called_once()
def test_epd_tokenizer_contains_duplicate_embedding_part():
class FakeSocket:
def __init__(self, messages):
self.messages = messages
self.closed = False
async def recv_multipart(self, copy=False):
return self.messages.pop(0)
def close(self):
self.closed = True
async def run_test():
embedding = torch.tensor([[1.0, 2.0]])
part = EmbeddingData(
req_id="request_local_part_0",
num_parts=2,
part_idx=0,
grid_dim=torch.tensor([[1, 2, 2]]),
modality=Modality.IMAGE,
embedding=embedding,
)
frame = [pickle.dumps(part.copy_without_embedding()), embedding.numpy()]
socket = FakeSocket([frame, frame])
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
receiver.model_type = None
result = await receiver._recv_mm_data(
"request", socket, SimpleNamespace(), "prompt"
)
assert result is None
assert socket.closed
asyncio.run(run_test())
def test_kimi_k3_encoder_prefers_grid_thws_and_uses_temporal_pool_length():
grid_thws = torch.tensor([[3, 8, 12]])
stale_grid = torch.tensor([[1, 2, 2]])
@@ -746,6 +948,81 @@ def test_epd_scheduler_uses_token_ids_for_tokenized_mm_processors():
)
def test_epd_scheduler_ignores_foreign_error_part():
waiting = WaitingZmqRequest.__new__(WaitingZmqRequest)
waiting.rid = "current"
waiting.recv_req = SimpleNamespace(rid="current")
waiting.status = WaitingMMRequestStatus.PENDING
waiting._fail_and_release = Mock()
stale_error = EmbeddingData(
req_id="stale_local_part_0",
num_parts=1,
part_idx=0,
grid_dim=None,
modality=Modality.IMAGE,
error_msg="stale failure",
error_code=500,
)
waiting.consume_parts([pickle.dumps("not embedding data")])
waiting.consume_parts([pickle.dumps(stale_error)])
assert waiting.status == WaitingMMRequestStatus.PENDING
waiting._fail_and_release.assert_not_called()
def test_epd_tokenizer_ignores_foreign_part_before_current_embedding():
class FakeSocket:
def __init__(self, messages):
self.messages = list(messages)
self.closed = False
async def recv_multipart(self, copy=False):
return self.messages.pop(0)
def close(self):
self.closed = True
async def run_test():
stale_error = EmbeddingData(
req_id="stale_local_part_0",
num_parts=1,
part_idx=0,
grid_dim=None,
modality=Modality.IMAGE,
error_msg="stale failure",
error_code=500,
)
embedding = torch.tensor([[1.0, 2.0]])
current = EmbeddingData(
req_id="current_local_part_0",
num_parts=1,
part_idx=0,
grid_dim=None,
modality=Modality.IMAGE,
embedding=embedding,
)
socket = FakeSocket(
[
[pickle.dumps(stale_error)],
[pickle.dumps(current.copy_without_embedding()), embedding.numpy()],
]
)
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
receiver.model_type = None
processor = SimpleNamespace(
get_mm_data=lambda _prompt, embeddings, **_kwargs: embeddings,
get_validated_mm_data=lambda _prompt, embeddings, **_kwargs: embeddings,
)
result = await receiver._recv_mm_data("current", socket, processor, "prompt")
torch.testing.assert_close(result[Modality.IMAGE], embedding)
assert socket.closed
asyncio.run(run_test())
def test_epd_scheduler_routes_many_requests_over_one_receive_socket():
context = zmq.Context()
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
@@ -761,6 +1038,8 @@ def test_epd_scheduler_routes_many_requests_over_one_receive_socket():
sender = context.socket(zmq.PUSH)
try:
sender.connect(f"tcp://127.0.0.1:{port}")
sender.send_multipart([b"not a pickle"])
sender.send_multipart([pickle.dumps("not embedding data")])
for i in range(32):
mm_data = EmbeddingData(
req_id=f"rid-{i}_local_part_0",
@@ -784,6 +1063,84 @@ def test_epd_scheduler_routes_many_requests_over_one_receive_socket():
context.term()
def _receiver_for_startup_failure(rank_errors):
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
receiver.mm_processor = object()
receiver.model_type = "kimi_k3"
receiver.hostname = "127.0.0.1"
receiver.tp_size = 2
receiver.tp_group = MagicMock()
receiver.tp_group.all_gather_object.side_effect = rank_errors
receiver.scheduler_recv_socket = object()
receiver.scheduler_context = object()
receiver.scheduler_embedding_port = 1234
receiver.encode_urls = ["http://encoder"]
receiver.waiting_by_rid = {}
receiver.waiting_list = []
receiver.create_req = MagicMock(return_value=object())
return receiver
def test_epd_receiver_startup_rejects_remote_rank_failure():
receiver = _receiver_for_startup_failure(
lambda local_error: [local_error, "RuntimeError: bind failed"]
)
waiting_req = MagicMock()
waiting_req.rid = "request-id"
waiting_cls = MagicMock(return_value=waiting_req)
class TokenizedRequest:
rid = "request-id"
need_wait_for_mm_inputs = True
encoder_urls = ["http://encoder"]
with patch(
"sglang.srt.disaggregation.encoder.receiver.TokenizedGenerateReqInput",
TokenizedRequest,
):
ready, aborts = receiver._process_waiting_requests(
[TokenizedRequest()], waiting_cls
)
assert ready == []
assert len(aborts) == 1
assert "rank 1: RuntimeError: bind failed" in aborts[0][1]
assert aborts[0][2] == 500
waiting_req.send_encode_request.assert_called_once_with()
waiting_req.release_resources.assert_called_once_with()
waiting_req.close_recv_socket.assert_called_once_with()
assert receiver.waiting_list == []
assert receiver.waiting_by_rid == {}
def test_epd_receiver_startup_shares_local_constructor_failure():
def gather_local_error(local_error):
assert "RuntimeError: socket failed" in local_error
return [local_error, None]
receiver = _receiver_for_startup_failure(gather_local_error)
waiting_cls = MagicMock(side_effect=RuntimeError("socket failed"))
class TokenizedRequest:
rid = "request-id"
need_wait_for_mm_inputs = True
encoder_urls = ["http://encoder"]
with patch(
"sglang.srt.disaggregation.encoder.receiver.TokenizedGenerateReqInput",
TokenizedRequest,
):
ready, aborts = receiver._process_waiting_requests(
[TokenizedRequest()], waiting_cls
)
assert ready == []
assert len(aborts) == 1
assert "rank 0: RuntimeError: socket failed" in aborts[0][1]
assert aborts[0][2] == 500
assert receiver.waiting_list == []
def test_epd_encoder_reuses_scheduler_zmq_peer():
async def send_twice():
context = zmq.asyncio.Context()
@@ -129,6 +129,7 @@ def _make_tokenizer_manager(case) -> TokenizerManager:
tm.server_args.dp_size = 1
tm.disaggregation_mode = "none"
tm.rid_to_state = {}
tm.encoder_dispatch_ready = {}
tm.enable_metrics = False
tm.enable_trace = False
tm.enable_lora = False
@@ -0,0 +1,39 @@
"""CPU tests for InternS1-Pro multimodal processor behavior."""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
import torch
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.multimodal.processors.interns1pro import InternS1_1ImageProcessor
def test_epd_stores_the_image_tensor_in_the_mm_item():
processor = object.__new__(InternS1_1ImageProcessor)
processor.build_input_ids = Mock(return_value=([1, 2, 3], [(1, 2)]))
processor.IM_START_TOKEN_ID = 10
processor.IM_END_TOKEN_ID = 11
processor.mm_tokens = SimpleNamespace(
image_token_id=12,
video_token_id=13,
audio_token_id=14,
)
image_embedding = torch.zeros(2, 4)
output = processor.get_validated_mm_data(
[1, 2, 3],
{Modality.IMAGE: image_embedding},
img_grid_thw=torch.tensor([[1, 2, 2]]),
)
assert output.mm_items[0].precomputed_embeddings is image_embedding
if __name__ == "__main__":
raise SystemExit(pytest.main([__file__, "-v"]))
+1 -1
View File
@@ -727,7 +727,7 @@ def test_kimi_k3_epd_rebuild_uses_the_same_media_contract():
processor._tokenizer = _Tokenizer()
embeddings = {Modality.IMAGE: torch.arange(20, dtype=torch.float32).reshape(5, 4)}
output = processor.get_mm_data(
output = processor.get_validated_mm_data(
[1, 99, 2, 99, 3],
embeddings,
img_grid_thw=torch.tensor([[1, 2, 6], [1, 2, 4]]),
@@ -377,6 +377,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
manager = object.__new__(TokenizerManager)
manager.rid_to_state = {}
manager.encoder_dispatch_ready = {}
transport = MagicMock()
transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
manager.cuda_vmm_feature_transport = transport
@@ -405,6 +406,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
manager = object.__new__(tokenizer_manager.TokenizerManager)
manager.rid_to_state = {}
manager.encoder_dispatch_ready = {}
transport = MagicMock()
manager._dispatch_to_scheduler = MagicMock(
side_effect=RuntimeError("send failed")
@@ -440,6 +442,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
manager = object.__new__(tokenizer_manager.TokenizerManager)
manager.rid_to_state = {}
manager.encoder_dispatch_ready = {}
transport = MagicMock()
manager._dispatch_to_scheduler = MagicMock()
time_stats = MagicMock()
@@ -0,0 +1,90 @@
"""Tests for the common EPD precomputed-embedding boundary."""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import unittest
import torch
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalProcessorOutput,
)
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
class _StubProcessor(BaseMultimodalProcessor):
async def process_mm_data_async(self, *args, **kwargs):
raise NotImplementedError
def get_mm_data(self, prompt, embeddings, **kwargs):
return self.output
def _item(modality, rows, offsets):
return MultimodalDataItem(
modality=modality,
offsets=offsets,
precomputed_embeddings=torch.zeros(rows, 4),
)
class TestPrecomputedEmbeddingValidation(unittest.TestCase):
def setUp(self):
self.processor = object.__new__(_StubProcessor)
def _validate(self, items, embeddings):
self.processor.output = MultimodalProcessorOutput(
input_ids=[1, 2, 3],
mm_items=items,
)
return self.processor.get_validated_mm_data([], embeddings)
def test_accepts_exact_multi_item_layout(self):
image_embedding = torch.zeros(5, 4)
audio_embedding = torch.zeros(2, 4)
output = self._validate(
[
_item(Modality.IMAGE, 2, [(1, 2)]),
_item(Modality.IMAGE, 3, [(4, 6)]),
_item(Modality.AUDIO, 2, [(8, 9)]),
],
{
Modality.IMAGE: image_embedding,
Modality.AUDIO: audio_embedding,
},
)
self.assertEqual(len(output.mm_items), 3)
def test_rejects_item_shorter_than_prompt_offsets(self):
with self.assertRaisesRegex(RuntimeError, "expected 3 rows.*got 2"):
self._validate(
[_item(Modality.IMAGE, 2, [(1, 3)])],
{Modality.IMAGE: torch.zeros(2, 4)},
)
def test_rejects_unconsumed_trailing_rows(self):
with self.assertRaisesRegex(RuntimeError, "received 3 rows, consumed 2"):
self._validate(
[_item(Modality.IMAGE, 2, [(1, 2)])],
{Modality.IMAGE: torch.zeros(3, 4)},
)
def test_rejects_missing_modality(self):
with self.assertRaisesRegex(RuntimeError, "received 2 rows, consumed 0"):
self._validate([], {Modality.VIDEO: torch.zeros(2, 4)})
def test_rejects_unexpected_modality(self):
with self.assertRaisesRegex(RuntimeError, "unexpected embedding modality"):
self._validate(
[_item(Modality.VIDEO, 2, [(1, 2)])],
{Modality.IMAGE: torch.zeros(2, 4)},
)
if __name__ == "__main__":
unittest.main()