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()
|
||||
|
||||
Reference in New Issue
Block a user