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