fix(vlm): harden EPD receiver validation and liveness (#36945)
This commit is contained in:
@@ -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"]))
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user