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