fix(vlm): contain EPD request lifecycle failures (#36944)
Co-authored-by: mickqian <mickqian@users.noreply.github.com>
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -17,6 +17,8 @@ class _FakeEncoder:
|
||||
self.embedding_to_send = {}
|
||||
self.encode_dispatch_lock = asyncio.Lock()
|
||||
self.encode_calls = []
|
||||
self.released = []
|
||||
self.release_event = asyncio.Event()
|
||||
|
||||
def has_pending_embeddings(self):
|
||||
return bool(self.embedding_to_send)
|
||||
@@ -28,8 +30,9 @@ class _FakeEncoder:
|
||||
self.encode_calls.append(kwargs)
|
||||
return 1, 1, 1, None, None
|
||||
|
||||
async def release_request(self, _req_id):
|
||||
return None
|
||||
async def release_request(self, req_id):
|
||||
self.released.append(req_id)
|
||||
self.release_event.set()
|
||||
|
||||
|
||||
def _install_tp_encoder(monkeypatch, encoder):
|
||||
@@ -84,5 +87,87 @@ def test_health_encode_rechecks_busy_state_after_waiting(monkeypatch):
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
def test_health_timeout_keeps_dispatch_order_until_encode_drains(monkeypatch):
|
||||
async def run_test():
|
||||
encoder = _FakeEncoder()
|
||||
_install_tp_encoder(monkeypatch, encoder)
|
||||
encode_started = asyncio.Event()
|
||||
finish_encode = asyncio.Event()
|
||||
|
||||
async def encode(**kwargs):
|
||||
encoder.encode_calls.append(kwargs)
|
||||
encode_started.set()
|
||||
await finish_encode.wait()
|
||||
return 1, 1, 1, None, None
|
||||
|
||||
encoder.encode = encode
|
||||
monkeypatch.setattr(http_server, "HEALTH_CHECK_TIMEOUT", 0.01)
|
||||
|
||||
response = await http_server.health_generate()
|
||||
assert response.status_code == 503
|
||||
assert encode_started.is_set()
|
||||
assert encoder.encode_dispatch_lock.locked()
|
||||
assert encoder.released == []
|
||||
|
||||
finish_encode.set()
|
||||
await asyncio.wait_for(encoder.release_event.wait(), timeout=1)
|
||||
await asyncio.wait_for(encoder.encode_dispatch_lock.acquire(), timeout=1)
|
||||
encoder.encode_dispatch_lock.release()
|
||||
assert len(encoder.released) == 1
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
def test_cancelled_health_request_does_not_cancel_dispatched_encode(monkeypatch):
|
||||
async def run_test():
|
||||
encoder = _FakeEncoder()
|
||||
_install_tp_encoder(monkeypatch, encoder)
|
||||
encode_started = asyncio.Event()
|
||||
finish_encode = asyncio.Event()
|
||||
|
||||
async def encode(**kwargs):
|
||||
encoder.encode_calls.append(kwargs)
|
||||
encode_started.set()
|
||||
await finish_encode.wait()
|
||||
return 1, 1, 1, None, None
|
||||
|
||||
encoder.encode = encode
|
||||
task = asyncio.create_task(http_server.health_generate())
|
||||
await asyncio.wait_for(encode_started.wait(), timeout=1)
|
||||
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
assert encoder.encode_dispatch_lock.locked()
|
||||
assert encoder.released == []
|
||||
|
||||
finish_encode.set()
|
||||
await asyncio.wait_for(encoder.release_event.wait(), timeout=1)
|
||||
await asyncio.wait_for(encoder.encode_dispatch_lock.acquire(), timeout=1)
|
||||
encoder.encode_dispatch_lock.release()
|
||||
assert len(encoder.released) == 1
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
def test_health_cleanup_failure_releases_dispatch_lock(monkeypatch):
|
||||
async def run_test():
|
||||
encoder = _FakeEncoder()
|
||||
_install_tp_encoder(monkeypatch, encoder)
|
||||
|
||||
async def release_request(req_id):
|
||||
encoder.released.append(req_id)
|
||||
raise RuntimeError("cleanup failed")
|
||||
|
||||
encoder.release_request = release_request
|
||||
response = await http_server.health_generate()
|
||||
|
||||
assert response.status_code == 503
|
||||
assert not encoder.encode_dispatch_lock.locked()
|
||||
assert len(encoder.released) == 1
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import asyncio
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -7,7 +9,9 @@ from sglang.srt.disaggregation.encoder.runtime import (
|
||||
EncoderScheduler,
|
||||
PendingRequest,
|
||||
_resolve_encoder_batch_policy,
|
||||
validate_encode_request,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import Modality
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
@@ -108,6 +112,113 @@ def test_scheduler_coalesces_concurrent_submissions():
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
def test_scheduler_isolates_bad_request_from_failed_fused_batch():
|
||||
class FakeEncoder:
|
||||
def __init__(self):
|
||||
self.encode_dispatch_lock = asyncio.Lock()
|
||||
self.batches = []
|
||||
|
||||
async def batch_encode(self, requests, _modality):
|
||||
req_ids = [request["req_id"] for request in requests]
|
||||
self.batches.append(req_ids)
|
||||
if len(requests) > 1 or req_ids == ["bad"]:
|
||||
return [(0, 0, 0, "bad image", 400) for _ in requests]
|
||||
return [(1, 2, 3, None, None)]
|
||||
|
||||
async def run_test():
|
||||
encoder = FakeEncoder()
|
||||
scheduler = EncoderScheduler(
|
||||
encoder=encoder,
|
||||
send_sockets=[],
|
||||
max_batch_size=8,
|
||||
coalesce_same_turn=True,
|
||||
)
|
||||
collector = SimpleNamespace(observe_queue_wait=Mock())
|
||||
with patch(
|
||||
"sglang.srt.disaggregation.encoder.runtime.server_module.encoder_metrics_collector",
|
||||
collector,
|
||||
):
|
||||
scheduler.start()
|
||||
try:
|
||||
requests = [
|
||||
{
|
||||
"req_id": req_id,
|
||||
"modality": "image",
|
||||
"mm_items": [object()],
|
||||
"num_parts": 1,
|
||||
"part_idx": 0,
|
||||
}
|
||||
for req_id in ("bad", "good")
|
||||
]
|
||||
results = await asyncio.gather(
|
||||
*(scheduler.submit(request) for request in requests)
|
||||
)
|
||||
finally:
|
||||
await scheduler.stop()
|
||||
|
||||
assert encoder.batches == [["bad", "good"], ["bad"], ["good"]]
|
||||
assert results == [(0, 0, 0, "bad image", 400), (1, 2, 3, None, None)]
|
||||
assert collector.observe_queue_wait.call_count == len(requests)
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("update", "expected"),
|
||||
[
|
||||
({"req_id": ""}, "missing or invalid req_id"),
|
||||
({"modality": "text"}, "unsupported modality"),
|
||||
({"mm_items": []}, "missing or empty mm_items"),
|
||||
({"num_parts": 0}, "num_parts must be a positive integer"),
|
||||
({"part_idx": 1}, "part_idx must be in [0, 1)"),
|
||||
],
|
||||
)
|
||||
def test_validate_encode_request_rejects_invalid_fields(update, expected):
|
||||
request = {
|
||||
"req_id": "request",
|
||||
"modality": "image",
|
||||
"mm_items": [object()],
|
||||
"num_parts": 1,
|
||||
"part_idx": 0,
|
||||
}
|
||||
request.update(update)
|
||||
|
||||
assert expected in validate_encode_request(request)
|
||||
|
||||
|
||||
def test_video_request_is_validated_before_tp_broadcast():
|
||||
class FakeSocket:
|
||||
pass
|
||||
|
||||
class FakeEncoder:
|
||||
async def encode(self, **_kwargs):
|
||||
raise AssertionError("invalid request must not reach the encoder")
|
||||
|
||||
async def run_test():
|
||||
scheduler = EncoderScheduler(
|
||||
encoder=FakeEncoder(),
|
||||
send_sockets=[FakeSocket()],
|
||||
max_batch_size=1,
|
||||
)
|
||||
pending = PendingRequest(
|
||||
{
|
||||
"req_id": "bad-video",
|
||||
"modality": "video",
|
||||
"mm_items": [object()],
|
||||
"num_parts": 1,
|
||||
"part_idx": 1,
|
||||
},
|
||||
asyncio.get_running_loop(),
|
||||
)
|
||||
|
||||
await scheduler._dispatch_per_request([pending], Modality.VIDEO)
|
||||
|
||||
with pytest.raises(Exception, match="part_idx must be in"):
|
||||
pending.future.result()
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model_type", "configured", "explicit", "expected"),
|
||||
[
|
||||
|
||||
Reference in New Issue
Block a user