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
@@ -30,7 +30,11 @@ from sglang.srt.distributed.parallel_state import (
get_mooncake_transfer_engine,
)
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReqInput
from sglang.srt.managers.io_struct import (
EncoderDispatchErrorReq,
GenerateReqInput,
TokenizedGenerateReqInput,
)
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.multimodal.cache import media_preprocess_kwargs
@@ -62,6 +66,22 @@ if TYPE_CHECKING:
from sglang.srt.managers.scheduler import Scheduler
class _ReceiveRegistrationRunner:
"""Run encoder receive-URL registration off the scheduler thread."""
def __init__(self, name: str):
self.loop = asyncio.new_event_loop()
self.thread = threading.Thread(target=self._run, daemon=True, name=name)
self.thread.start()
def _run(self) -> None:
asyncio.set_event_loop(self.loop)
self.loop.run_forever()
def submit(self, coroutine):
return asyncio.run_coroutine_threadsafe(coroutine, self.loop)
def _mark_keep_device_embedding(mm_inputs) -> None:
"""Tell general_mm_embed_routine not to copy embeddings back to CPU."""
if mm_inputs is None:
@@ -356,15 +376,15 @@ def _normalize_embedding_ports(embedding_port):
return [embedding_port]
def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count):
async def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count):
import grpc
from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc
timeout_secs = envs.SGLANG_ENCODER_GRPC_TIMEOUT_SECS.get()
channel = grpc.insecure_channel(target)
channel = grpc.aio.insecure_channel(target)
stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel)
try:
stub.SchedulerReceiveUrl(
await stub.SchedulerReceiveUrl(
sglang_encoder_pb2.SchedulerReceiveUrlRequest(
req_id=req_id,
receive_url=receive_url,
@@ -373,7 +393,7 @@ def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count):
timeout=timeout_secs,
)
finally:
channel.close()
await channel.close()
def _grpc_encode_request(target, encode_request):
@@ -402,6 +422,24 @@ def _grpc_encode_request(target, encode_request):
channel.close()
async def _gather_blocking_grpc_calls(calls):
"""Wait for synchronous gRPC threads to stop before propagating cancellation."""
future = asyncio.gather(*calls, return_exceptions=True)
try:
results = await asyncio.shield(future)
except asyncio.CancelledError:
results = await asyncio.shield(future)
for result in results:
if isinstance(result, Exception):
logger.error("gRPC call failed while draining cancellation: %s", result)
raise
for result in results:
if isinstance(result, Exception):
raise result
return results
class EmbeddingData:
def __init__(
self,
@@ -610,6 +648,7 @@ class MultiModalEmbeddingData(EmbeddingData):
model_type: Optional[str] = None,
):
"""Create MultiModalEmbeddingData from an EmbeddingData instance."""
_validate_embedding_part(embedding_data)
# Only forward known optional attrs (e.g. video metadata) so they land on the instance
extra = {}
for attr in video_meta_attrs_for(model_type):
@@ -677,13 +716,7 @@ class MultiModalEmbeddingData(EmbeddingData):
return kwargs
def add(self, embedding_data: EmbeddingData):
if self.req_id != embedding_data.req_id:
logger.warning(
f"Dropping embedding data with mismatched req_id: "
f"expected {self.req_id}, got {embedding_data.req_id}"
)
return False
assert not self.ready_list[embedding_data.part_idx]
_validate_embedding_part(embedding_data, current=self)
pid = embedding_data.part_idx
self.ready_list[pid] = True
self.modality_list[pid] = embedding_data.modality
@@ -696,6 +729,44 @@ class MultiModalEmbeddingData(EmbeddingData):
self._set_image_meta_for_part(pid, embedding_data)
def _validate_embedding_part(
embedding_data: EmbeddingData,
current: Optional[MultiModalEmbeddingData] = None,
) -> None:
"""Reject malformed part metadata before indexing aggregation buffers."""
if not isinstance(embedding_data, EmbeddingData):
raise ValueError(f"expected EmbeddingData, got {type(embedding_data).__name__}")
if (
not isinstance(embedding_data.num_parts, int)
or isinstance(embedding_data.num_parts, bool)
or embedding_data.num_parts <= 0
):
raise ValueError("num_parts must be a positive integer")
if (
not isinstance(embedding_data.part_idx, int)
or isinstance(embedding_data.part_idx, bool)
or embedding_data.part_idx < 0
or embedding_data.part_idx >= embedding_data.num_parts
):
raise ValueError(
f"part_idx must be in [0, {embedding_data.num_parts}), "
f"got {embedding_data.part_idx}"
)
if current is None:
return
if current.req_id != embedding_data.req_id:
raise ValueError(
f"embedding req_id mismatch: expected {current.req_id}, "
f"got {embedding_data.req_id}"
)
if current.num_parts != embedding_data.num_parts:
raise ValueError(
f"num_parts changed from {current.num_parts} to {embedding_data.num_parts}"
)
if current.ready_list[embedding_data.part_idx]:
raise ValueError(f"duplicate embedding part {embedding_data.part_idx}")
def _aggregate_embedding_part(current, recv_obj, model_type):
"""Fold one received part into the aggregate (the first part creates it)."""
if current is None:
@@ -732,6 +803,41 @@ def extract_original_req_id(part_req_id: str) -> str:
return part_req_id
def _resolve_embedding_part_request_id(
embedding_data: object, expected_req_id: Optional[str] = None
) -> Optional[str]:
"""Validate and normalize an embedding part ID for safe routing."""
expected = (
f" for expected rid={expected_req_id}" if expected_req_id is not None else ""
)
if not isinstance(embedding_data, EmbeddingData):
logger.warning("Dropping non-embedding data%s", expected)
return None
if not isinstance(embedding_data.req_id, str):
logger.warning("Dropping embedding data with a non-string req_id%s", expected)
return None
original_req_id = extract_original_req_id(embedding_data.req_id)
if expected_req_id is not None and original_req_id != expected_req_id:
logger.warning(
"Dropping stale embedding data: expected rid=%s, got rid=%s "
"(likely from ZMQ port reuse)",
expected_req_id,
embedding_data.req_id,
)
return None
embedding_data.req_id = original_req_id
return original_req_id
def _embedding_part_matches_request(
embedding_data: object, expected_req_id: str
) -> bool:
"""Normalize a matching part ID; reject stale data from a reused socket."""
return (
_resolve_embedding_part_request_id(embedding_data, expected_req_id) is not None
)
def _encoder_media_item(mm_item: dict):
"""Keep per-media options aligned while preserving the legacy URL shape."""
item = {
@@ -784,6 +890,7 @@ class WaitingMMRequestBase(ABC):
embedding_pool: Optional["EmbeddingPool"] = None,
zmq_context=None,
embedding_port=None,
registration_runner: Optional[_ReceiveRegistrationRunner] = None,
):
self.rid = rid
self.recv_req = recv_req
@@ -820,6 +927,10 @@ class WaitingMMRequestBase(ABC):
# Success-path finalizer handle so abort can release the slot early.
self._mm_finalizer: Optional[weakref.finalize] = None
self._pool_full_warned = False
self.registration_runner = registration_runner
self.registration_future = None
self.registration_error = None
self.registration_lock = threading.Lock()
@abstractmethod
def send_encode_request(self) -> None:
@@ -829,6 +940,15 @@ class WaitingMMRequestBase(ABC):
if self.status != WaitingMMRequestStatus.PENDING:
return
with self.registration_lock:
registration_error, self.registration_error = (
self.registration_error,
None,
)
if registration_error is not None:
self._fail_and_release(*registration_error)
return
# A complete request can remain pending while the GPU pool is full.
# Retry assembly on every scheduler tick, including shared-socket mode.
if self.recv_embedding_data is not None and self.recv_embedding_data.ready:
@@ -860,6 +980,8 @@ class WaitingMMRequestBase(ABC):
try:
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if not self._is_valid_embedding_part(recv_obj):
return
if getattr(recv_obj, "error_msg", None) is not None:
logger.warning(
f"Received error signal from encoder for {self.rid}: "
@@ -867,8 +989,6 @@ class WaitingMMRequestBase(ABC):
)
self._fail_and_release(recv_obj.error_msg, recv_obj.error_code)
return
if not self._is_valid_embedding_part(recv_obj):
return
# ZMQ materializes frame 1; RDMA already wrote the registered buffer.
self._extract_embedding_from_buffer(recv_obj, parts)
self.recv_embedding_data = _aggregate_embedding_part(
@@ -896,29 +1016,12 @@ class WaitingMMRequestBase(ABC):
self.error_msg = error_msg
self.error_code = error_code
self.status = WaitingMMRequestStatus.FAIL
self._cleanup_gpu_buffer()
self.release_resources()
self.close_recv_socket()
async def _check_encoder_responses(self, responses, endpoint: str) -> bool:
"""Validate gathered encoder responses; on the first error, FAIL the
request and release its resources. Returns True if all succeeded."""
msg = await _extract_encoder_error(responses, endpoint, f"rid={self.rid}")
if msg is None:
return True
self._fail_and_release(msg)
return False
def _is_valid_embedding_part(self, recv_obj) -> bool:
"""Check for and drop stale or out-of-sync payloads; normalize the part req_id to the original rid."""
original_req_id = extract_original_req_id(recv_obj.req_id)
if original_req_id != self.recv_req.rid:
logger.warning(
f"Dropping stale embedding data: expected rid={self.recv_req.rid}, "
f"got rid={recv_obj.req_id} (likely from ZMQ port reuse)"
)
return False
recv_obj.req_id = original_req_id
return True
return _embedding_part_matches_request(recv_obj, self.recv_req.rid)
@abstractmethod
def _extract_embedding_from_buffer(self, recv_obj, parts) -> None:
@@ -964,8 +1067,8 @@ class WaitingMMRequestBase(ABC):
return True
def _finish_assemble(self, recv_embedding) -> None:
"""get_mm_data → bind pool slot → publish onto recv_req → SUCCESS."""
mm_inputs = self.mm_processor.get_mm_data(
"""Build validated mm data, bind its pool slot, then publish it."""
mm_inputs = self.mm_processor.get_validated_mm_data(
_select_mm_processor_prompt(self.recv_req, self.mm_processor),
recv_embedding,
**self.recv_embedding_data.get_mm_extra_meta(),
@@ -1005,6 +1108,12 @@ class WaitingMMRequestBase(ABC):
def release_resources(self):
"""Free pool/GPU resources on abort/fail/timeout. Idempotent."""
registration_future, self.registration_future = (
self.registration_future,
None,
)
if registration_future is not None and not registration_future.done():
registration_future.cancel()
self._cleanup_gpu_buffer()
finalizer, self._mm_finalizer = self._mm_finalizer, None
if finalizer is not None:
@@ -1014,8 +1123,37 @@ class WaitingMMRequestBase(ABC):
# For zmq_to_scheduler: embedding parts arrive as ZMQ payload frames and
# are optionally staged into the GPU EmbeddingPool.
class WaitingZmqRequest(WaitingMMRequestBase):
def send_encode_request(self):
def _start_registration(self, coroutine) -> None:
if self.registration_runner is None:
coroutine.close()
self._fail_and_release(
"Encoder receive registration runner is unavailable",
int(HTTPStatus.INTERNAL_SERVER_ERROR),
)
return
self.registration_future = self.registration_runner.submit(coroutine)
self.registration_future.add_done_callback(self._on_registration_done)
def _on_registration_done(self, future) -> None:
if future.cancelled():
return
error = future.exception()
if error is None:
return
logger.error(
"Failed to register encoder receive URL for rid=%s: %s",
self.rid,
error,
exc_info=error,
)
with self.registration_lock:
self.registration_error = (
f"Failed to register receive URL with encoder: {error}",
int(HTTPStatus.BAD_GATEWAY),
)
def send_encode_request(self):
async def _send_single_request(session, url, payload):
try:
async with session.post(url, json=payload) as response:
@@ -1053,7 +1191,7 @@ class WaitingZmqRequest(WaitingMMRequestBase):
encoder_url = self.encoder_urls[idx]
target_url = f"{encoder_url}/scheduler_receive_url"
payload = {
"req_id": part_req_id, # use part_req_id to match encode request
"req_id": part_req_id,
"receive_count": receive_count,
"receive_url": NetworkAddress(
host_name, embedding_port
@@ -1091,15 +1229,9 @@ class WaitingZmqRequest(WaitingMMRequestBase):
logger.debug(f"Request {i} succeeded.")
failed = [r for r in results if isinstance(r, BaseException)]
if failed:
# A rank without a registered receive URL can never be
# pushed to; fail via the normal completion path now
# instead of pending until the embedding wait times out.
self._fail_and_release(
f"Failed to register receive URL with encoder: {failed[0]!r}",
int(HTTPStatus.BAD_GATEWAY),
)
raise failed[0]
asyncio.run(
self._start_registration(
send_embedding_port(
self.recv_req.rid,
self.receive_count,
@@ -1180,8 +1312,7 @@ class WaitingZmqRequestGrpc(WaitingZmqRequest):
target_url = f"{encoder_url}/SchedulerReceiveUrl"
logger.info(f"Preparing to send to {target_url}")
tasks.append(
asyncio.to_thread(
_grpc_scheduler_receive_url,
_grpc_scheduler_receive_url(
_grpc_target(encoder_url),
req_id,
receive_url,
@@ -1200,8 +1331,11 @@ class WaitingZmqRequestGrpc(WaitingZmqRequest):
logger.error(f"Request {i} failed: {result}")
else:
logger.debug(f"Request {i} succeeded.")
failed = [r for r in results if isinstance(r, BaseException)]
if failed:
raise failed[0]
asyncio.run(
self._start_registration(
send_embedding_port(
self.recv_req.rid,
self.receive_count,
@@ -1249,6 +1383,8 @@ class WaitingRDMARequest(WaitingMMRequestBase):
self._buffer_lock = threading.Lock()
self._terminal = False
self._receive_running = False
self._receive_error = None
self._receive_error_lock = threading.Lock()
def send_encode_request(self):
# Base-class hook. The tokenizer owns /encode, so this rank only pulls
@@ -1261,13 +1397,34 @@ class WaitingRDMARequest(WaitingMMRequestBase):
asyncio.run(self._pull_meta_and_receive_embedding())
except Exception as e:
logger.error(f"RDMA receive failed for rid={self.rid}: {e}")
self._fail_and_release(str(e))
self._record_receive_error(str(e))
finally:
with self._buffer_lock:
self._receive_running = False
if self._terminal:
self._release_buffer_locked()
def _record_receive_error(self, error_msg, error_code=None) -> None:
"""Pass a worker-thread failure to the scheduler thread."""
with self._receive_error_lock:
self._receive_error = (error_msg, error_code)
def _try_recv_mm_data(self):
with self._receive_error_lock:
receive_error, self._receive_error = self._receive_error, None
if receive_error is not None:
self._fail_and_release(*receive_error)
return
super()._try_recv_mm_data()
async def _check_encoder_responses(self, responses, endpoint: str) -> bool:
"""Record network failures for the scheduler thread to consume."""
error = await _extract_encoder_error(responses, endpoint, f"rid={self.rid}")
if error is None:
return True
self._record_receive_error(*error)
return False
async def _pull_meta_and_receive_embedding(self):
"""Pull per-part sizes, allocate the landing buffer, then drive /send.
@@ -1334,7 +1491,7 @@ class WaitingRDMARequest(WaitingMMRequestBase):
)
if alloc_result is None:
# Oversize or alloc timeout — fatal for this request.
self._fail_and_release(
self._record_receive_error(
f"EmbeddingPool could not allocate "
f"{total_bytes // (1024 * 1024)}MB (oversize or "
f"timeout). Raise SGLANG_EMBEDDING_POOL_SIZE_MB."
@@ -1449,7 +1606,7 @@ class WaitingRDMARequest(WaitingMMRequestBase):
async def _extract_encoder_error(responses, endpoint, context, encode_requests=None):
"""Return the first error among gathered encoder responses, or None.
"""Return the first ``(message, status)`` error, or None.
Pure check — logs each error but has no other side effects; the caller
decides how to react. ``encode_requests`` optionally enriches each log
@@ -1467,13 +1624,16 @@ async def _extract_encoder_error(responses, endpoint, context, encode_requests=N
logger.error(
f"Encoder {endpoint} timeout ({timeout_val}s) for {ctx} (request {i})"
)
return f"Encoder {endpoint} timeout ({timeout_val}s)"
return (
f"Encoder {endpoint} timeout ({timeout_val}s)",
int(HTTPStatus.GATEWAY_TIMEOUT),
)
if isinstance(resp, Exception):
logger.error(
f"Encoder {endpoint} failed for {ctx} (request {i}): {resp}",
exc_info=resp,
)
return str(resp)
return str(resp), int(HTTPStatus.BAD_GATEWAY)
if resp.status != 200:
try:
err = await resp.json()
@@ -1481,7 +1641,7 @@ async def _extract_encoder_error(responses, endpoint, context, encode_requests=N
except Exception:
msg = await resp.text()
logger.error(f"Encoder {endpoint} returned error {resp.status}: {msg}")
return msg
return msg, int(resp.status)
return None
@@ -1739,12 +1899,16 @@ class MMReceiverBase(ABC):
self.hostname = get_local_ip_auto()
self.waiting_list: List[WaitingMMRequestBase] = []
self.waiting_by_rid: Dict[str, WaitingMMRequestBase] = {}
self.registration_runner = None
self.scheduler_embedding_port = None
self.scheduler_recv_socket = None
if (
self.encoder_transfer_backend == "zmq_to_scheduler"
and scheduler is not None
):
self.registration_runner = _ReceiveRegistrationRunner(
f"encoder-receive-registration-{tp_rank}"
)
(
self.scheduler_embedding_port,
self.scheduler_recv_socket,
@@ -1887,6 +2051,10 @@ class MMReceiverBase(ABC):
self, request_obj, mm_processor, prompt, need_wait_for_mm_inputs=True
):
req_id = None
recv_socket = None
encode_task = None
recv_task = None
send_time = time.monotonic()
try:
# ``self.encode_urls`` is shared by reference with the bootstrap
# server (when running) so it always reflects the current set.
@@ -1931,13 +2099,13 @@ class MMReceiverBase(ABC):
done
and recv_task not in done
and (
encode_task.exception() is not None or encode_task.result() is False
encode_task.exception() is not None
or encode_task.result() is not None
)
):
logger.warning(
f"[{req_id}] Encoder dispatch failed; skipping embedding wait"
)
recv_task.cancel()
return None
result = await asyncio.wait_for(
recv_task,
@@ -1950,6 +2118,15 @@ class MMReceiverBase(ABC):
elapsed = time.monotonic() - send_time
logger.warning(f"[{req_id}] Embedding recv timeout after {elapsed:.3f}s")
return None
finally:
tasks = [task for task in (encode_task, recv_task) if task is not None]
for task in tasks:
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
if recv_socket is not None:
recv_socket.close(linger=0)
async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt):
"""zmq_to_tokenizer receive: embedding parts arrive as 2-frame ZMQ
@@ -1966,6 +2143,8 @@ class MMReceiverBase(ABC):
if not parts:
continue
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if not _embedding_part_matches_request(recv_obj, req_id):
continue
if getattr(recv_obj, "error_msg", None) is not None:
logger.warning(
f"Encoder error for req_id={req_id}: {recv_obj.error_msg} "
@@ -1973,8 +2152,6 @@ class MMReceiverBase(ABC):
)
return None
logger.debug("recv_obj=%s", recv_obj)
# Normalize the part req_id to the original for aggregation.
recv_obj.req_id = extract_original_req_id(recv_obj.req_id)
if len(parts) < 2:
logger.error(
"zmq_to_tokenizer expected 2-part message, got %d parts",
@@ -1993,18 +2170,31 @@ class MMReceiverBase(ABC):
)
recv_embedding = recv_embedding_data.get_embedding(is_concat=True)
return mm_processor.get_mm_data(
return mm_processor.get_validated_mm_data(
prompt,
recv_embedding,
**recv_embedding_data.get_mm_extra_meta(),
)
except Exception:
logger.exception(
"Failed to receive encoder embeddings for req_id=%s", req_id
)
return None
finally:
recv_socket.close()
def send_encode_request(self, obj, time_stats_json=None):
self._send_encode_request(obj, time_stats_json=time_stats_json)
def send_encode_request(
self, obj, time_stats_json=None, on_dispatch_error=None
) -> Optional[threading.Event]:
return self._send_encode_request(
obj,
time_stats_json=time_stats_json,
on_dispatch_error=on_dispatch_error,
)
def _send_encode_request(self, obj, time_stats_json=None):
def _send_encode_request(
self, obj, time_stats_json=None, on_dispatch_error=None
) -> Optional[threading.Event]:
mm_data = self._extract_url_data(obj)
if obj.rid is None:
obj.rid = uuid.uuid4().hex
@@ -2028,6 +2218,7 @@ class MMReceiverBase(ABC):
# Freeze the encoder URL snapshot onto obj so the scheduler
# subprocess uses the same list when indexing encoder_idx.
obj.encoder_urls = encode_urls
scheduler_dispatch_ready = threading.Event()
encode_thread = threading.Thread(
target=self._run_encode_in_thread,
@@ -2038,10 +2229,13 @@ class MMReceiverBase(ABC):
num_items_assigned,
encode_urls,
time_stats_json,
scheduler_dispatch_ready,
on_dispatch_error,
),
daemon=True,
)
encode_thread.start()
return scheduler_dispatch_ready
else:
# No encoder URLs available (bootstrap may not have any registered yet);
# reset the flag so the scheduler does not wait for embeddings that will
@@ -2053,6 +2247,7 @@ class MMReceiverBase(ABC):
"processing without encoder disaggregation."
)
obj.need_wait_for_mm_inputs = False
return None
def _sync_fail_info_across_tp(self, waiting_req: WaitingMMRequestBase) -> None:
"""Share encoder error fields across TP ranks before abort.
@@ -2092,8 +2287,18 @@ class MMReceiverBase(ABC):
except zmq.Again:
return
try:
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
rid = extract_original_req_id(recv_obj.req_id)
except Exception as error:
logger.warning(
"Dropping malformed embedding data from the shared "
"scheduler socket: %s",
error,
)
continue
rid = _resolve_embedding_part_request_id(recv_obj)
if rid is None:
continue
waiting_req = self.waiting_by_rid.get(rid)
if waiting_req is None:
logger.warning(
@@ -2104,7 +2309,21 @@ class MMReceiverBase(ABC):
def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs):
new_recv_reqs = []
abort_reqs = []
for recv_req in recv_reqs:
if isinstance(recv_req, EncoderDispatchErrorReq):
waiting_req = self.waiting_by_rid.get(recv_req.rid)
if waiting_req is None:
logger.debug(
"Ignoring encoder dispatch error for inactive request %s",
recv_req.rid,
)
else:
waiting_req._fail_and_release(
recv_req.error_msg, recv_req.error_code
)
continue
if (
isinstance(recv_req, TokenizedGenerateReqInput)
and recv_req.need_wait_for_mm_inputs is True
@@ -2116,6 +2335,9 @@ class MMReceiverBase(ABC):
# tokenizer never set encoder_urls (legacy / static path).
encode_urls = recv_req.encoder_urls or list(self.encode_urls)
waiting_req = None
local_error = None
try:
waiting_req = waiting_cls(
rid=recv_req.rid,
recv_req=recv_req,
@@ -2132,15 +2354,62 @@ class MMReceiverBase(ABC):
embedding_port=self.scheduler_embedding_port,
**extra_kwargs,
)
if self.scheduler_recv_socket is not None:
self.waiting_by_rid[waiting_req.rid] = waiting_req
waiting_req.send_encode_request()
except Exception as error:
local_error = f"{type(error).__name__}: {error}"
logger.exception(
"Failed to start multimodal receive for rid=%s", recv_req.rid
)
# The status all-reduce below requires every TP rank to append
# exactly the same requests. Agree on startup before appending.
rank_errors = (
[local_error]
if self.tp_size <= 1
else self.tp_group.all_gather_object(local_error)
)
failed_ranks = [
rank for rank, error in enumerate(rank_errors) if error is not None
]
if failed_ranks:
details = "; ".join(
f"rank {rank}: {rank_errors[rank]}" for rank in failed_ranks
)
error_msg = f"Failed to start multimodal receive ({details})"
logger.error(error_msg)
if waiting_req is not None:
try:
waiting_req.release_resources()
except Exception:
logger.exception(
"Failed to release multimodal receive resources "
"for rid=%s",
waiting_req.rid,
)
try:
waiting_req.close_recv_socket()
except Exception:
logger.exception(
"Failed to close multimodal receive socket for rid=%s",
waiting_req.rid,
)
self.waiting_by_rid.pop(waiting_req.rid, None)
abort_reqs.append(
(
self.create_req(recv_req),
error_msg,
HTTPStatus.INTERNAL_SERVER_ERROR,
)
)
continue
self.waiting_list.append(waiting_req)
else:
new_recv_reqs.append(recv_req)
if len(self.waiting_list) == 0:
return new_recv_reqs, []
return new_recv_reqs, abort_reqs
self._drain_scheduler_embeddings()
current_time = time.time()
@@ -2164,7 +2433,6 @@ class MMReceiverBase(ABC):
)
new_waiting = []
abort_reqs = []
for i, waiting_req in enumerate(self.waiting_list):
status_value = local_status[i].item()
if status_value == WaitingMMRequestStatus.SUCCESS:
@@ -2199,6 +2467,7 @@ class MMReceiverBase(ABC):
else: # status_value == WaitingMMRequestStatus.PENDING
new_waiting.append(waiting_req)
continue
waiting_req.close_recv_socket()
self.waiting_by_rid.pop(waiting_req.rid, None)
self.waiting_list = new_waiting
@@ -2212,12 +2481,14 @@ class MMReceiverBase(ABC):
num_items_assigned,
encode_urls=None,
time_stats_json=None,
scheduler_dispatch_ready=None,
on_dispatch_error=None,
):
# ``embedding_port`` is always None on this path: zmq_to_scheduler /
# mooncake ranks register their receive ports with the encoder later
# via /scheduler_receive_url, so the dispatch itself carries no port.
try:
asyncio.run(
dispatch_error = asyncio.run(
self.encode(
req_id=req_id,
mm_data=mm_data,
@@ -2230,6 +2501,15 @@ class MMReceiverBase(ABC):
)
except Exception as e:
logger.error(f"Encode failed for request {req_id}: {e}", exc_info=True)
dispatch_error = EncoderDispatchErrorReq(
rid=req_id,
error_msg=str(e),
error_code=int(HTTPStatus.BAD_GATEWAY),
)
if dispatch_error is not None and on_dispatch_error is not None:
scheduler_dispatch_ready.wait()
on_dispatch_error(dispatch_error)
def create_req(self, recv_req: TokenizedGenerateReqInput):
req = Req(
@@ -2422,7 +2702,10 @@ class MMReceiverHTTP(MMReceiverBase):
embedding_pool=self.embedding_pool,
)
return self._process_waiting_requests(
recv_reqs, WaitingZmqRequest, embedding_pool=self.embedding_pool
recv_reqs,
WaitingZmqRequest,
embedding_pool=self.embedding_pool,
registration_runner=self.registration_runner,
)
async def encode(
@@ -2511,11 +2794,15 @@ class MMReceiverHTTP(MMReceiverBase):
# zmq_to_tokenizer is pushed to our PULL socket during /encode,
# zmq_to_scheduler to the ports its ranks registered, and mooncake
# by RDMA once those ranks have pulled sizes and driven /send.
return (
await _extract_encoder_error(
error = await _extract_encoder_error(
responses, "HTTP request", f"req_id={req_id}", encode_requests
)
is None
if error is None:
return None
return EncoderDispatchErrorReq(
rid=req_id,
error_msg=error[0],
error_code=error[1],
)
@@ -2551,7 +2838,11 @@ class MMReceiverGrpc(MMReceiverBase):
# For zmq_to_scheduler
def process_waiting_requests(self, recv_reqs):
return self._process_waiting_requests(recv_reqs, WaitingZmqRequestGrpc)
return self._process_waiting_requests(
recv_reqs,
WaitingZmqRequestGrpc,
registration_runner=self.registration_runner,
)
async def encode(
self,
@@ -2621,7 +2912,7 @@ class MMReceiverGrpc(MMReceiverBase):
)
for encode_request in encode_requests
]
await asyncio.gather(*grpc_tasks)
await _gather_blocking_grpc_calls(grpc_tasks)
def _validate_transport_mode(transport_mode: str, encoder_urls):
+7
View File
@@ -2053,6 +2053,13 @@ class AbortReq(BaseReq, kw_only=True):
self.rid = ""
class EncoderDispatchErrorReq(BaseReq, kw_only=True):
"""Tokenizer-to-scheduler failure for one EPD encoder dispatch."""
error_msg: str
error_code: int
class ActiveRanksOutput(BaseReq, kw_only=True):
status: List[bool]
@@ -72,6 +72,7 @@ from sglang.srt.managers.io_struct import (
ContinueGenerationReqInput,
ElasticScaleUpdateReq,
EmbeddingReqInput,
EncoderDispatchErrorReq,
FreezeGCReq,
GenerateReqInput,
HealthCheckOutput,
@@ -588,6 +589,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def init_running_status(self):
# Request states
self.rid_to_state: Dict[str, ReqState] = {}
self.encoder_dispatch_ready: Dict[str, threading.Event] = {}
self.event_loop = None
self.asyncio_tasks = set()
@@ -1579,6 +1581,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self._dispatch_to_scheduler(tokenized_obj)
self._mark_state_dispatched(tokenized_obj.rid)
dispatched = True
dispatch_ready = self.encoder_dispatch_ready.pop(tokenized_obj.rid, None)
if dispatch_ready is not None:
dispatch_ready.set()
tokenized_obj.time_stats = time_stats
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
finally:
@@ -3495,15 +3500,28 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
"""
for rid in rids:
state = self.rid_to_state.get(rid)
if state is None:
continue
if state is not None:
if state.dispatched:
try:
self.abort_request(rid)
except Exception:
logger.exception("Failed to abort request %s during cleanup", rid)
logger.exception(
"Failed to abort request %s during cleanup", rid
)
else:
del self.rid_to_state[rid]
dispatch_ready = self.encoder_dispatch_ready.pop(rid, None)
if dispatch_ready is not None:
dispatch_ready.set()
def _forward_encoder_dispatch_error(self, error: EncoderDispatchErrorReq) -> None:
if error.rid in self.rid_to_state:
self._dispatch_to_scheduler(error)
def _schedule_encoder_dispatch_error(self, error: EncoderDispatchErrorReq) -> None:
self.event_loop.call_soon_threadsafe(
self._forward_encoder_dispatch_error, error
)
def _should_dispatch_to_encoder(
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
@@ -3554,9 +3572,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if state is not None:
time_stats_json = state.time_stats.encode_json()
self.mm_receiver.send_encode_request(
obj, time_stats_json=time_stats_json
dispatch_ready = self.mm_receiver.send_encode_request(
obj,
time_stats_json=time_stats_json,
on_dispatch_error=self._schedule_encoder_dispatch_error,
)
if dispatch_ready is not None:
self.encoder_dispatch_ready[obj.rid] = dispatch_ready
else:
obj.need_wait_for_mm_inputs = False
@@ -647,6 +647,69 @@ class BaseMultimodalProcessor(ABC):
video_token_id=getattr(self, "VIDEO_TOKEN_ID", None),
)
def get_validated_mm_data(
self,
prompt,
embeddings: Dict[Modality, torch.Tensor],
**kwargs,
) -> MultimodalProcessorOutput:
"""Build EPD multimodal inputs and validate the embedding layout.
Model processors may override ``get_mm_data`` to rebuild their prompt
layout. This shared wrapper ensures every override consumes exactly the
encoder rows it received before the result reaches the scheduler.
"""
output = self.get_mm_data(prompt, embeddings, **kwargs)
self._validate_precomputed_embedding_layout(output, embeddings)
return output
@staticmethod
def _validate_precomputed_embedding_layout(
output: MultimodalProcessorOutput,
embeddings: Dict[Modality, torch.Tensor],
) -> None:
consumed_per_modality = {modality: 0 for modality in embeddings}
for item in output.mm_items:
embedding = item.precomputed_embeddings
if not isinstance(embedding, torch.Tensor):
raise RuntimeError(
"EPD multimodal items must contain tensor embeddings; "
f"got {type(embedding).__name__} for "
f"{item.modality.name.lower()}"
)
num_rows = embedding.shape[0]
if item.offsets is not None:
expected_rows = sum(end - start + 1 for start, end in item.offsets)
if num_rows != expected_rows:
raise RuntimeError(
"Precomputed multimodal embedding length mismatch for "
f"{item.modality.name.lower()}: expected {expected_rows} "
f"rows from prompt offsets, got {num_rows}"
)
if item.modality not in consumed_per_modality:
raise RuntimeError(
"EPD processor returned an unexpected embedding modality: "
f"{item.modality.name.lower()}"
)
consumed_per_modality[item.modality] += num_rows
for modality, embedding in embeddings.items():
if not isinstance(embedding, torch.Tensor):
raise RuntimeError(
"EPD encoder output must contain tensor embeddings; "
f"got {type(embedding).__name__} for {modality.name.lower()}"
)
consumed_rows = consumed_per_modality[modality]
if consumed_rows != embedding.shape[0]:
raise RuntimeError(
"Precomputed multimodal embedding consumption mismatch for "
f"{modality.name.lower()}: received {embedding.shape[0]} rows, "
f"consumed {consumed_rows}"
)
def _resolve_processor(self, processor=None):
if processor is None:
return self._processor, self._tokenizer
@@ -26,7 +26,7 @@ class InternS1_1ImageProcessor(QwenVLImageProcessor):
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=offsets,
precomputed_embeddings=embeddings,
precomputed_embeddings=embeddings[Modality.IMAGE],
)
]
@@ -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()