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
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
rid = extract_original_req_id(recv_obj.req_id)
try:
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
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,31 +2335,81 @@ class MMReceiverBase(ABC):
# tokenizer never set encoder_urls (legacy / static path).
encode_urls = recv_req.encoder_urls or list(self.encode_urls)
waiting_req = waiting_cls(
rid=recv_req.rid,
recv_req=recv_req,
mm_processor=self.mm_processor,
encoder_urls=encode_urls,
model_type=self.model_type,
host_name=self.hostname,
receive_count=self.tp_size,
zmq_context=(
None
if self.scheduler_recv_socket is not None
else self.scheduler_context
),
embedding_port=self.scheduler_embedding_port,
**extra_kwargs,
)
if self.scheduler_recv_socket is not None:
waiting_req = None
local_error = None
try:
waiting_req = waiting_cls(
rid=recv_req.rid,
recv_req=recv_req,
mm_processor=self.mm_processor,
encoder_urls=encode_urls,
model_type=self.model_type,
host_name=self.hostname,
receive_count=self.tp_size,
zmq_context=(
None
if self.scheduler_recv_socket is not None
else self.scheduler_context
),
embedding_port=self.scheduler_embedding_port,
**extra_kwargs,
)
self.waiting_by_rid[waiting_req.rid] = waiting_req
waiting_req.send_encode_request()
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(
responses, "HTTP request", f"req_id={req_id}", encode_requests
)
is None
error = await _extract_encoder_error(
responses, "HTTP request", f"req_id={req_id}", encode_requests
)
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]
+33 -11
View File
@@ -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.dispatched:
try:
self.abort_request(rid)
except Exception:
logger.exception("Failed to abort request %s during cleanup", rid)
else:
del self.rid_to_state[rid]
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
)
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],
)
]