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,
|
get_mooncake_transfer_engine,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
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.multimodal_processor import get_mm_processor, import_processors
|
||||||
from sglang.srt.managers.schedule_batch import Modality, Req
|
from sglang.srt.managers.schedule_batch import Modality, Req
|
||||||
from sglang.srt.multimodal.cache import media_preprocess_kwargs
|
from sglang.srt.multimodal.cache import media_preprocess_kwargs
|
||||||
@@ -62,6 +66,22 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.managers.scheduler import Scheduler
|
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:
|
def _mark_keep_device_embedding(mm_inputs) -> None:
|
||||||
"""Tell general_mm_embed_routine not to copy embeddings back to CPU."""
|
"""Tell general_mm_embed_routine not to copy embeddings back to CPU."""
|
||||||
if mm_inputs is None:
|
if mm_inputs is None:
|
||||||
@@ -356,15 +376,15 @@ def _normalize_embedding_ports(embedding_port):
|
|||||||
return [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
|
import grpc
|
||||||
from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc
|
from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc
|
||||||
|
|
||||||
timeout_secs = envs.SGLANG_ENCODER_GRPC_TIMEOUT_SECS.get()
|
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)
|
stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel)
|
||||||
try:
|
try:
|
||||||
stub.SchedulerReceiveUrl(
|
await stub.SchedulerReceiveUrl(
|
||||||
sglang_encoder_pb2.SchedulerReceiveUrlRequest(
|
sglang_encoder_pb2.SchedulerReceiveUrlRequest(
|
||||||
req_id=req_id,
|
req_id=req_id,
|
||||||
receive_url=receive_url,
|
receive_url=receive_url,
|
||||||
@@ -373,7 +393,7 @@ def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count):
|
|||||||
timeout=timeout_secs,
|
timeout=timeout_secs,
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
channel.close()
|
await channel.close()
|
||||||
|
|
||||||
|
|
||||||
def _grpc_encode_request(target, encode_request):
|
def _grpc_encode_request(target, encode_request):
|
||||||
@@ -402,6 +422,24 @@ def _grpc_encode_request(target, encode_request):
|
|||||||
channel.close()
|
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:
|
class EmbeddingData:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -610,6 +648,7 @@ class MultiModalEmbeddingData(EmbeddingData):
|
|||||||
model_type: Optional[str] = None,
|
model_type: Optional[str] = None,
|
||||||
):
|
):
|
||||||
"""Create MultiModalEmbeddingData from an EmbeddingData instance."""
|
"""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
|
# Only forward known optional attrs (e.g. video metadata) so they land on the instance
|
||||||
extra = {}
|
extra = {}
|
||||||
for attr in video_meta_attrs_for(model_type):
|
for attr in video_meta_attrs_for(model_type):
|
||||||
@@ -677,13 +716,7 @@ class MultiModalEmbeddingData(EmbeddingData):
|
|||||||
return kwargs
|
return kwargs
|
||||||
|
|
||||||
def add(self, embedding_data: EmbeddingData):
|
def add(self, embedding_data: EmbeddingData):
|
||||||
if self.req_id != embedding_data.req_id:
|
_validate_embedding_part(embedding_data, current=self)
|
||||||
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]
|
|
||||||
pid = embedding_data.part_idx
|
pid = embedding_data.part_idx
|
||||||
self.ready_list[pid] = True
|
self.ready_list[pid] = True
|
||||||
self.modality_list[pid] = embedding_data.modality
|
self.modality_list[pid] = embedding_data.modality
|
||||||
@@ -696,6 +729,44 @@ class MultiModalEmbeddingData(EmbeddingData):
|
|||||||
self._set_image_meta_for_part(pid, embedding_data)
|
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):
|
def _aggregate_embedding_part(current, recv_obj, model_type):
|
||||||
"""Fold one received part into the aggregate (the first part creates it)."""
|
"""Fold one received part into the aggregate (the first part creates it)."""
|
||||||
if current is None:
|
if current is None:
|
||||||
@@ -732,6 +803,41 @@ def extract_original_req_id(part_req_id: str) -> str:
|
|||||||
return part_req_id
|
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):
|
def _encoder_media_item(mm_item: dict):
|
||||||
"""Keep per-media options aligned while preserving the legacy URL shape."""
|
"""Keep per-media options aligned while preserving the legacy URL shape."""
|
||||||
item = {
|
item = {
|
||||||
@@ -784,6 +890,7 @@ class WaitingMMRequestBase(ABC):
|
|||||||
embedding_pool: Optional["EmbeddingPool"] = None,
|
embedding_pool: Optional["EmbeddingPool"] = None,
|
||||||
zmq_context=None,
|
zmq_context=None,
|
||||||
embedding_port=None,
|
embedding_port=None,
|
||||||
|
registration_runner: Optional[_ReceiveRegistrationRunner] = None,
|
||||||
):
|
):
|
||||||
self.rid = rid
|
self.rid = rid
|
||||||
self.recv_req = recv_req
|
self.recv_req = recv_req
|
||||||
@@ -820,6 +927,10 @@ class WaitingMMRequestBase(ABC):
|
|||||||
# Success-path finalizer handle so abort can release the slot early.
|
# Success-path finalizer handle so abort can release the slot early.
|
||||||
self._mm_finalizer: Optional[weakref.finalize] = None
|
self._mm_finalizer: Optional[weakref.finalize] = None
|
||||||
self._pool_full_warned = False
|
self._pool_full_warned = False
|
||||||
|
self.registration_runner = registration_runner
|
||||||
|
self.registration_future = None
|
||||||
|
self.registration_error = None
|
||||||
|
self.registration_lock = threading.Lock()
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def send_encode_request(self) -> None:
|
def send_encode_request(self) -> None:
|
||||||
@@ -829,6 +940,15 @@ class WaitingMMRequestBase(ABC):
|
|||||||
if self.status != WaitingMMRequestStatus.PENDING:
|
if self.status != WaitingMMRequestStatus.PENDING:
|
||||||
return
|
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.
|
# A complete request can remain pending while the GPU pool is full.
|
||||||
# Retry assembly on every scheduler tick, including shared-socket mode.
|
# Retry assembly on every scheduler tick, including shared-socket mode.
|
||||||
if self.recv_embedding_data is not None and self.recv_embedding_data.ready:
|
if self.recv_embedding_data is not None and self.recv_embedding_data.ready:
|
||||||
@@ -860,6 +980,8 @@ class WaitingMMRequestBase(ABC):
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
|
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:
|
if getattr(recv_obj, "error_msg", None) is not None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Received error signal from encoder for {self.rid}: "
|
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)
|
self._fail_and_release(recv_obj.error_msg, recv_obj.error_code)
|
||||||
return
|
return
|
||||||
if not self._is_valid_embedding_part(recv_obj):
|
|
||||||
return
|
|
||||||
# ZMQ materializes frame 1; RDMA already wrote the registered buffer.
|
# ZMQ materializes frame 1; RDMA already wrote the registered buffer.
|
||||||
self._extract_embedding_from_buffer(recv_obj, parts)
|
self._extract_embedding_from_buffer(recv_obj, parts)
|
||||||
self.recv_embedding_data = _aggregate_embedding_part(
|
self.recv_embedding_data = _aggregate_embedding_part(
|
||||||
@@ -896,29 +1016,12 @@ class WaitingMMRequestBase(ABC):
|
|||||||
self.error_msg = error_msg
|
self.error_msg = error_msg
|
||||||
self.error_code = error_code
|
self.error_code = error_code
|
||||||
self.status = WaitingMMRequestStatus.FAIL
|
self.status = WaitingMMRequestStatus.FAIL
|
||||||
self._cleanup_gpu_buffer()
|
self.release_resources()
|
||||||
self.close_recv_socket()
|
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:
|
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."""
|
"""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)
|
return _embedding_part_matches_request(recv_obj, self.recv_req.rid)
|
||||||
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
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def _extract_embedding_from_buffer(self, recv_obj, parts) -> None:
|
def _extract_embedding_from_buffer(self, recv_obj, parts) -> None:
|
||||||
@@ -964,8 +1067,8 @@ class WaitingMMRequestBase(ABC):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def _finish_assemble(self, recv_embedding) -> None:
|
def _finish_assemble(self, recv_embedding) -> None:
|
||||||
"""get_mm_data → bind pool slot → publish onto recv_req → SUCCESS."""
|
"""Build validated mm data, bind its pool slot, then publish it."""
|
||||||
mm_inputs = self.mm_processor.get_mm_data(
|
mm_inputs = self.mm_processor.get_validated_mm_data(
|
||||||
_select_mm_processor_prompt(self.recv_req, self.mm_processor),
|
_select_mm_processor_prompt(self.recv_req, self.mm_processor),
|
||||||
recv_embedding,
|
recv_embedding,
|
||||||
**self.recv_embedding_data.get_mm_extra_meta(),
|
**self.recv_embedding_data.get_mm_extra_meta(),
|
||||||
@@ -1005,6 +1108,12 @@ class WaitingMMRequestBase(ABC):
|
|||||||
|
|
||||||
def release_resources(self):
|
def release_resources(self):
|
||||||
"""Free pool/GPU resources on abort/fail/timeout. Idempotent."""
|
"""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()
|
self._cleanup_gpu_buffer()
|
||||||
finalizer, self._mm_finalizer = self._mm_finalizer, None
|
finalizer, self._mm_finalizer = self._mm_finalizer, None
|
||||||
if finalizer is not 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
|
# For zmq_to_scheduler: embedding parts arrive as ZMQ payload frames and
|
||||||
# are optionally staged into the GPU EmbeddingPool.
|
# are optionally staged into the GPU EmbeddingPool.
|
||||||
class WaitingZmqRequest(WaitingMMRequestBase):
|
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):
|
async def _send_single_request(session, url, payload):
|
||||||
try:
|
try:
|
||||||
async with session.post(url, json=payload) as response:
|
async with session.post(url, json=payload) as response:
|
||||||
@@ -1053,7 +1191,7 @@ class WaitingZmqRequest(WaitingMMRequestBase):
|
|||||||
encoder_url = self.encoder_urls[idx]
|
encoder_url = self.encoder_urls[idx]
|
||||||
target_url = f"{encoder_url}/scheduler_receive_url"
|
target_url = f"{encoder_url}/scheduler_receive_url"
|
||||||
payload = {
|
payload = {
|
||||||
"req_id": part_req_id, # use part_req_id to match encode request
|
"req_id": part_req_id,
|
||||||
"receive_count": receive_count,
|
"receive_count": receive_count,
|
||||||
"receive_url": NetworkAddress(
|
"receive_url": NetworkAddress(
|
||||||
host_name, embedding_port
|
host_name, embedding_port
|
||||||
@@ -1091,15 +1229,9 @@ class WaitingZmqRequest(WaitingMMRequestBase):
|
|||||||
logger.debug(f"Request {i} succeeded.")
|
logger.debug(f"Request {i} succeeded.")
|
||||||
failed = [r for r in results if isinstance(r, BaseException)]
|
failed = [r for r in results if isinstance(r, BaseException)]
|
||||||
if failed:
|
if failed:
|
||||||
# A rank without a registered receive URL can never be
|
raise failed[0]
|
||||||
# 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),
|
|
||||||
)
|
|
||||||
|
|
||||||
asyncio.run(
|
self._start_registration(
|
||||||
send_embedding_port(
|
send_embedding_port(
|
||||||
self.recv_req.rid,
|
self.recv_req.rid,
|
||||||
self.receive_count,
|
self.receive_count,
|
||||||
@@ -1180,8 +1312,7 @@ class WaitingZmqRequestGrpc(WaitingZmqRequest):
|
|||||||
target_url = f"{encoder_url}/SchedulerReceiveUrl"
|
target_url = f"{encoder_url}/SchedulerReceiveUrl"
|
||||||
logger.info(f"Preparing to send to {target_url}")
|
logger.info(f"Preparing to send to {target_url}")
|
||||||
tasks.append(
|
tasks.append(
|
||||||
asyncio.to_thread(
|
_grpc_scheduler_receive_url(
|
||||||
_grpc_scheduler_receive_url,
|
|
||||||
_grpc_target(encoder_url),
|
_grpc_target(encoder_url),
|
||||||
req_id,
|
req_id,
|
||||||
receive_url,
|
receive_url,
|
||||||
@@ -1200,8 +1331,11 @@ class WaitingZmqRequestGrpc(WaitingZmqRequest):
|
|||||||
logger.error(f"Request {i} failed: {result}")
|
logger.error(f"Request {i} failed: {result}")
|
||||||
else:
|
else:
|
||||||
logger.debug(f"Request {i} succeeded.")
|
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(
|
send_embedding_port(
|
||||||
self.recv_req.rid,
|
self.recv_req.rid,
|
||||||
self.receive_count,
|
self.receive_count,
|
||||||
@@ -1249,6 +1383,8 @@ class WaitingRDMARequest(WaitingMMRequestBase):
|
|||||||
self._buffer_lock = threading.Lock()
|
self._buffer_lock = threading.Lock()
|
||||||
self._terminal = False
|
self._terminal = False
|
||||||
self._receive_running = False
|
self._receive_running = False
|
||||||
|
self._receive_error = None
|
||||||
|
self._receive_error_lock = threading.Lock()
|
||||||
|
|
||||||
def send_encode_request(self):
|
def send_encode_request(self):
|
||||||
# Base-class hook. The tokenizer owns /encode, so this rank only pulls
|
# 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())
|
asyncio.run(self._pull_meta_and_receive_embedding())
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"RDMA receive failed for rid={self.rid}: {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:
|
finally:
|
||||||
with self._buffer_lock:
|
with self._buffer_lock:
|
||||||
self._receive_running = False
|
self._receive_running = False
|
||||||
if self._terminal:
|
if self._terminal:
|
||||||
self._release_buffer_locked()
|
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):
|
async def _pull_meta_and_receive_embedding(self):
|
||||||
"""Pull per-part sizes, allocate the landing buffer, then drive /send.
|
"""Pull per-part sizes, allocate the landing buffer, then drive /send.
|
||||||
|
|
||||||
@@ -1334,7 +1491,7 @@ class WaitingRDMARequest(WaitingMMRequestBase):
|
|||||||
)
|
)
|
||||||
if alloc_result is None:
|
if alloc_result is None:
|
||||||
# Oversize or alloc timeout — fatal for this request.
|
# Oversize or alloc timeout — fatal for this request.
|
||||||
self._fail_and_release(
|
self._record_receive_error(
|
||||||
f"EmbeddingPool could not allocate "
|
f"EmbeddingPool could not allocate "
|
||||||
f"{total_bytes // (1024 * 1024)}MB (oversize or "
|
f"{total_bytes // (1024 * 1024)}MB (oversize or "
|
||||||
f"timeout). Raise SGLANG_EMBEDDING_POOL_SIZE_MB."
|
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):
|
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
|
Pure check — logs each error but has no other side effects; the caller
|
||||||
decides how to react. ``encode_requests`` optionally enriches each log
|
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(
|
logger.error(
|
||||||
f"Encoder {endpoint} timeout ({timeout_val}s) for {ctx} (request {i})"
|
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):
|
if isinstance(resp, Exception):
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Encoder {endpoint} failed for {ctx} (request {i}): {resp}",
|
f"Encoder {endpoint} failed for {ctx} (request {i}): {resp}",
|
||||||
exc_info=resp,
|
exc_info=resp,
|
||||||
)
|
)
|
||||||
return str(resp)
|
return str(resp), int(HTTPStatus.BAD_GATEWAY)
|
||||||
if resp.status != 200:
|
if resp.status != 200:
|
||||||
try:
|
try:
|
||||||
err = await resp.json()
|
err = await resp.json()
|
||||||
@@ -1481,7 +1641,7 @@ async def _extract_encoder_error(responses, endpoint, context, encode_requests=N
|
|||||||
except Exception:
|
except Exception:
|
||||||
msg = await resp.text()
|
msg = await resp.text()
|
||||||
logger.error(f"Encoder {endpoint} returned error {resp.status}: {msg}")
|
logger.error(f"Encoder {endpoint} returned error {resp.status}: {msg}")
|
||||||
return msg
|
return msg, int(resp.status)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
@@ -1739,12 +1899,16 @@ class MMReceiverBase(ABC):
|
|||||||
self.hostname = get_local_ip_auto()
|
self.hostname = get_local_ip_auto()
|
||||||
self.waiting_list: List[WaitingMMRequestBase] = []
|
self.waiting_list: List[WaitingMMRequestBase] = []
|
||||||
self.waiting_by_rid: Dict[str, WaitingMMRequestBase] = {}
|
self.waiting_by_rid: Dict[str, WaitingMMRequestBase] = {}
|
||||||
|
self.registration_runner = None
|
||||||
self.scheduler_embedding_port = None
|
self.scheduler_embedding_port = None
|
||||||
self.scheduler_recv_socket = None
|
self.scheduler_recv_socket = None
|
||||||
if (
|
if (
|
||||||
self.encoder_transfer_backend == "zmq_to_scheduler"
|
self.encoder_transfer_backend == "zmq_to_scheduler"
|
||||||
and scheduler is not None
|
and scheduler is not None
|
||||||
):
|
):
|
||||||
|
self.registration_runner = _ReceiveRegistrationRunner(
|
||||||
|
f"encoder-receive-registration-{tp_rank}"
|
||||||
|
)
|
||||||
(
|
(
|
||||||
self.scheduler_embedding_port,
|
self.scheduler_embedding_port,
|
||||||
self.scheduler_recv_socket,
|
self.scheduler_recv_socket,
|
||||||
@@ -1887,6 +2051,10 @@ class MMReceiverBase(ABC):
|
|||||||
self, request_obj, mm_processor, prompt, need_wait_for_mm_inputs=True
|
self, request_obj, mm_processor, prompt, need_wait_for_mm_inputs=True
|
||||||
):
|
):
|
||||||
req_id = None
|
req_id = None
|
||||||
|
recv_socket = None
|
||||||
|
encode_task = None
|
||||||
|
recv_task = None
|
||||||
|
send_time = time.monotonic()
|
||||||
try:
|
try:
|
||||||
# ``self.encode_urls`` is shared by reference with the bootstrap
|
# ``self.encode_urls`` is shared by reference with the bootstrap
|
||||||
# server (when running) so it always reflects the current set.
|
# server (when running) so it always reflects the current set.
|
||||||
@@ -1931,13 +2099,13 @@ class MMReceiverBase(ABC):
|
|||||||
done
|
done
|
||||||
and recv_task not in done
|
and recv_task not in done
|
||||||
and (
|
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(
|
logger.warning(
|
||||||
f"[{req_id}] Encoder dispatch failed; skipping embedding wait"
|
f"[{req_id}] Encoder dispatch failed; skipping embedding wait"
|
||||||
)
|
)
|
||||||
recv_task.cancel()
|
|
||||||
return None
|
return None
|
||||||
result = await asyncio.wait_for(
|
result = await asyncio.wait_for(
|
||||||
recv_task,
|
recv_task,
|
||||||
@@ -1950,6 +2118,15 @@ class MMReceiverBase(ABC):
|
|||||||
elapsed = time.monotonic() - send_time
|
elapsed = time.monotonic() - send_time
|
||||||
logger.warning(f"[{req_id}] Embedding recv timeout after {elapsed:.3f}s")
|
logger.warning(f"[{req_id}] Embedding recv timeout after {elapsed:.3f}s")
|
||||||
return None
|
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):
|
async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt):
|
||||||
"""zmq_to_tokenizer receive: embedding parts arrive as 2-frame ZMQ
|
"""zmq_to_tokenizer receive: embedding parts arrive as 2-frame ZMQ
|
||||||
@@ -1966,6 +2143,8 @@ class MMReceiverBase(ABC):
|
|||||||
if not parts:
|
if not parts:
|
||||||
continue
|
continue
|
||||||
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
|
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:
|
if getattr(recv_obj, "error_msg", None) is not None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Encoder error for req_id={req_id}: {recv_obj.error_msg} "
|
f"Encoder error for req_id={req_id}: {recv_obj.error_msg} "
|
||||||
@@ -1973,8 +2152,6 @@ class MMReceiverBase(ABC):
|
|||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
logger.debug("recv_obj=%s", recv_obj)
|
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:
|
if len(parts) < 2:
|
||||||
logger.error(
|
logger.error(
|
||||||
"zmq_to_tokenizer expected 2-part message, got %d parts",
|
"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)
|
recv_embedding = recv_embedding_data.get_embedding(is_concat=True)
|
||||||
return mm_processor.get_mm_data(
|
return mm_processor.get_validated_mm_data(
|
||||||
prompt,
|
prompt,
|
||||||
recv_embedding,
|
recv_embedding,
|
||||||
**recv_embedding_data.get_mm_extra_meta(),
|
**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:
|
finally:
|
||||||
recv_socket.close()
|
recv_socket.close()
|
||||||
|
|
||||||
def send_encode_request(self, obj, time_stats_json=None):
|
def send_encode_request(
|
||||||
self._send_encode_request(obj, time_stats_json=time_stats_json)
|
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)
|
mm_data = self._extract_url_data(obj)
|
||||||
if obj.rid is None:
|
if obj.rid is None:
|
||||||
obj.rid = uuid.uuid4().hex
|
obj.rid = uuid.uuid4().hex
|
||||||
@@ -2028,6 +2218,7 @@ class MMReceiverBase(ABC):
|
|||||||
# Freeze the encoder URL snapshot onto obj so the scheduler
|
# Freeze the encoder URL snapshot onto obj so the scheduler
|
||||||
# subprocess uses the same list when indexing encoder_idx.
|
# subprocess uses the same list when indexing encoder_idx.
|
||||||
obj.encoder_urls = encode_urls
|
obj.encoder_urls = encode_urls
|
||||||
|
scheduler_dispatch_ready = threading.Event()
|
||||||
|
|
||||||
encode_thread = threading.Thread(
|
encode_thread = threading.Thread(
|
||||||
target=self._run_encode_in_thread,
|
target=self._run_encode_in_thread,
|
||||||
@@ -2038,10 +2229,13 @@ class MMReceiverBase(ABC):
|
|||||||
num_items_assigned,
|
num_items_assigned,
|
||||||
encode_urls,
|
encode_urls,
|
||||||
time_stats_json,
|
time_stats_json,
|
||||||
|
scheduler_dispatch_ready,
|
||||||
|
on_dispatch_error,
|
||||||
),
|
),
|
||||||
daemon=True,
|
daemon=True,
|
||||||
)
|
)
|
||||||
encode_thread.start()
|
encode_thread.start()
|
||||||
|
return scheduler_dispatch_ready
|
||||||
else:
|
else:
|
||||||
# No encoder URLs available (bootstrap may not have any registered yet);
|
# No encoder URLs available (bootstrap may not have any registered yet);
|
||||||
# reset the flag so the scheduler does not wait for embeddings that will
|
# reset the flag so the scheduler does not wait for embeddings that will
|
||||||
@@ -2053,6 +2247,7 @@ class MMReceiverBase(ABC):
|
|||||||
"processing without encoder disaggregation."
|
"processing without encoder disaggregation."
|
||||||
)
|
)
|
||||||
obj.need_wait_for_mm_inputs = False
|
obj.need_wait_for_mm_inputs = False
|
||||||
|
return None
|
||||||
|
|
||||||
def _sync_fail_info_across_tp(self, waiting_req: WaitingMMRequestBase) -> None:
|
def _sync_fail_info_across_tp(self, waiting_req: WaitingMMRequestBase) -> None:
|
||||||
"""Share encoder error fields across TP ranks before abort.
|
"""Share encoder error fields across TP ranks before abort.
|
||||||
@@ -2092,8 +2287,18 @@ class MMReceiverBase(ABC):
|
|||||||
except zmq.Again:
|
except zmq.Again:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
|
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)
|
waiting_req = self.waiting_by_rid.get(rid)
|
||||||
if waiting_req is None:
|
if waiting_req is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -2104,7 +2309,21 @@ class MMReceiverBase(ABC):
|
|||||||
|
|
||||||
def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs):
|
def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs):
|
||||||
new_recv_reqs = []
|
new_recv_reqs = []
|
||||||
|
abort_reqs = []
|
||||||
for recv_req in recv_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 (
|
if (
|
||||||
isinstance(recv_req, TokenizedGenerateReqInput)
|
isinstance(recv_req, TokenizedGenerateReqInput)
|
||||||
and recv_req.need_wait_for_mm_inputs is True
|
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).
|
# tokenizer never set encoder_urls (legacy / static path).
|
||||||
encode_urls = recv_req.encoder_urls or list(self.encode_urls)
|
encode_urls = recv_req.encoder_urls or list(self.encode_urls)
|
||||||
|
|
||||||
|
waiting_req = None
|
||||||
|
local_error = None
|
||||||
|
try:
|
||||||
waiting_req = waiting_cls(
|
waiting_req = waiting_cls(
|
||||||
rid=recv_req.rid,
|
rid=recv_req.rid,
|
||||||
recv_req=recv_req,
|
recv_req=recv_req,
|
||||||
@@ -2132,15 +2354,62 @@ class MMReceiverBase(ABC):
|
|||||||
embedding_port=self.scheduler_embedding_port,
|
embedding_port=self.scheduler_embedding_port,
|
||||||
**extra_kwargs,
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
if self.scheduler_recv_socket is not None:
|
|
||||||
self.waiting_by_rid[waiting_req.rid] = waiting_req
|
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)
|
self.waiting_list.append(waiting_req)
|
||||||
else:
|
else:
|
||||||
new_recv_reqs.append(recv_req)
|
new_recv_reqs.append(recv_req)
|
||||||
|
|
||||||
if len(self.waiting_list) == 0:
|
if len(self.waiting_list) == 0:
|
||||||
return new_recv_reqs, []
|
return new_recv_reqs, abort_reqs
|
||||||
|
|
||||||
self._drain_scheduler_embeddings()
|
self._drain_scheduler_embeddings()
|
||||||
current_time = time.time()
|
current_time = time.time()
|
||||||
@@ -2164,7 +2433,6 @@ class MMReceiverBase(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
new_waiting = []
|
new_waiting = []
|
||||||
abort_reqs = []
|
|
||||||
for i, waiting_req in enumerate(self.waiting_list):
|
for i, waiting_req in enumerate(self.waiting_list):
|
||||||
status_value = local_status[i].item()
|
status_value = local_status[i].item()
|
||||||
if status_value == WaitingMMRequestStatus.SUCCESS:
|
if status_value == WaitingMMRequestStatus.SUCCESS:
|
||||||
@@ -2199,6 +2467,7 @@ class MMReceiverBase(ABC):
|
|||||||
else: # status_value == WaitingMMRequestStatus.PENDING
|
else: # status_value == WaitingMMRequestStatus.PENDING
|
||||||
new_waiting.append(waiting_req)
|
new_waiting.append(waiting_req)
|
||||||
continue
|
continue
|
||||||
|
waiting_req.close_recv_socket()
|
||||||
self.waiting_by_rid.pop(waiting_req.rid, None)
|
self.waiting_by_rid.pop(waiting_req.rid, None)
|
||||||
|
|
||||||
self.waiting_list = new_waiting
|
self.waiting_list = new_waiting
|
||||||
@@ -2212,12 +2481,14 @@ class MMReceiverBase(ABC):
|
|||||||
num_items_assigned,
|
num_items_assigned,
|
||||||
encode_urls=None,
|
encode_urls=None,
|
||||||
time_stats_json=None,
|
time_stats_json=None,
|
||||||
|
scheduler_dispatch_ready=None,
|
||||||
|
on_dispatch_error=None,
|
||||||
):
|
):
|
||||||
# ``embedding_port`` is always None on this path: zmq_to_scheduler /
|
# ``embedding_port`` is always None on this path: zmq_to_scheduler /
|
||||||
# mooncake ranks register their receive ports with the encoder later
|
# mooncake ranks register their receive ports with the encoder later
|
||||||
# via /scheduler_receive_url, so the dispatch itself carries no port.
|
# via /scheduler_receive_url, so the dispatch itself carries no port.
|
||||||
try:
|
try:
|
||||||
asyncio.run(
|
dispatch_error = asyncio.run(
|
||||||
self.encode(
|
self.encode(
|
||||||
req_id=req_id,
|
req_id=req_id,
|
||||||
mm_data=mm_data,
|
mm_data=mm_data,
|
||||||
@@ -2230,6 +2501,15 @@ class MMReceiverBase(ABC):
|
|||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Encode failed for request {req_id}: {e}", exc_info=True)
|
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):
|
def create_req(self, recv_req: TokenizedGenerateReqInput):
|
||||||
req = Req(
|
req = Req(
|
||||||
@@ -2422,7 +2702,10 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
embedding_pool=self.embedding_pool,
|
embedding_pool=self.embedding_pool,
|
||||||
)
|
)
|
||||||
return self._process_waiting_requests(
|
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(
|
async def encode(
|
||||||
@@ -2511,11 +2794,15 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
# zmq_to_tokenizer is pushed to our PULL socket during /encode,
|
# zmq_to_tokenizer is pushed to our PULL socket during /encode,
|
||||||
# zmq_to_scheduler to the ports its ranks registered, and mooncake
|
# zmq_to_scheduler to the ports its ranks registered, and mooncake
|
||||||
# by RDMA once those ranks have pulled sizes and driven /send.
|
# by RDMA once those ranks have pulled sizes and driven /send.
|
||||||
return (
|
error = await _extract_encoder_error(
|
||||||
await _extract_encoder_error(
|
|
||||||
responses, "HTTP request", f"req_id={req_id}", encode_requests
|
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
|
# For zmq_to_scheduler
|
||||||
def process_waiting_requests(self, recv_reqs):
|
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(
|
async def encode(
|
||||||
self,
|
self,
|
||||||
@@ -2621,7 +2912,7 @@ class MMReceiverGrpc(MMReceiverBase):
|
|||||||
)
|
)
|
||||||
for encode_request in encode_requests
|
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):
|
def _validate_transport_mode(transport_mode: str, encoder_urls):
|
||||||
|
|||||||
@@ -2053,6 +2053,13 @@ class AbortReq(BaseReq, kw_only=True):
|
|||||||
self.rid = ""
|
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):
|
class ActiveRanksOutput(BaseReq, kw_only=True):
|
||||||
status: List[bool]
|
status: List[bool]
|
||||||
|
|
||||||
|
|||||||
@@ -72,6 +72,7 @@ from sglang.srt.managers.io_struct import (
|
|||||||
ContinueGenerationReqInput,
|
ContinueGenerationReqInput,
|
||||||
ElasticScaleUpdateReq,
|
ElasticScaleUpdateReq,
|
||||||
EmbeddingReqInput,
|
EmbeddingReqInput,
|
||||||
|
EncoderDispatchErrorReq,
|
||||||
FreezeGCReq,
|
FreezeGCReq,
|
||||||
GenerateReqInput,
|
GenerateReqInput,
|
||||||
HealthCheckOutput,
|
HealthCheckOutput,
|
||||||
@@ -588,6 +589,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
def init_running_status(self):
|
def init_running_status(self):
|
||||||
# Request states
|
# Request states
|
||||||
self.rid_to_state: Dict[str, ReqState] = {}
|
self.rid_to_state: Dict[str, ReqState] = {}
|
||||||
|
self.encoder_dispatch_ready: Dict[str, threading.Event] = {}
|
||||||
self.event_loop = None
|
self.event_loop = None
|
||||||
self.asyncio_tasks = set()
|
self.asyncio_tasks = set()
|
||||||
|
|
||||||
@@ -1579,6 +1581,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self._dispatch_to_scheduler(tokenized_obj)
|
self._dispatch_to_scheduler(tokenized_obj)
|
||||||
self._mark_state_dispatched(tokenized_obj.rid)
|
self._mark_state_dispatched(tokenized_obj.rid)
|
||||||
dispatched = True
|
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 = time_stats
|
||||||
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
|
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
|
||||||
finally:
|
finally:
|
||||||
@@ -3495,15 +3500,28 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
"""
|
"""
|
||||||
for rid in rids:
|
for rid in rids:
|
||||||
state = self.rid_to_state.get(rid)
|
state = self.rid_to_state.get(rid)
|
||||||
if state is None:
|
if state is not None:
|
||||||
continue
|
|
||||||
if state.dispatched:
|
if state.dispatched:
|
||||||
try:
|
try:
|
||||||
self.abort_request(rid)
|
self.abort_request(rid)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Failed to abort request %s during cleanup", rid)
|
logger.exception(
|
||||||
|
"Failed to abort request %s during cleanup", rid
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
del self.rid_to_state[rid]
|
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(
|
def _should_dispatch_to_encoder(
|
||||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||||
@@ -3554,9 +3572,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
if state is not None:
|
if state is not None:
|
||||||
time_stats_json = state.time_stats.encode_json()
|
time_stats_json = state.time_stats.encode_json()
|
||||||
|
|
||||||
self.mm_receiver.send_encode_request(
|
dispatch_ready = self.mm_receiver.send_encode_request(
|
||||||
obj, time_stats_json=time_stats_json
|
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:
|
else:
|
||||||
obj.need_wait_for_mm_inputs = False
|
obj.need_wait_for_mm_inputs = False
|
||||||
|
|
||||||
|
|||||||
@@ -647,6 +647,69 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
video_token_id=getattr(self, "VIDEO_TOKEN_ID", None),
|
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):
|
def _resolve_processor(self, processor=None):
|
||||||
if processor is None:
|
if processor is None:
|
||||||
return self._processor, self._tokenizer
|
return self._processor, self._tokenizer
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ class InternS1_1ImageProcessor(QwenVLImageProcessor):
|
|||||||
MultimodalDataItem(
|
MultimodalDataItem(
|
||||||
modality=Modality.IMAGE,
|
modality=Modality.IMAGE,
|
||||||
offsets=offsets,
|
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
|
import unittest
|
||||||
from array import array
|
from array import array
|
||||||
|
from http import HTTPStatus
|
||||||
from types import SimpleNamespace
|
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.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.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
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")
|
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):
|
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):
|
def test_extra_key_and_cache_salt_are_forwarded(self):
|
||||||
scheduler = SimpleNamespace(
|
scheduler = SimpleNamespace(
|
||||||
model_config=SimpleNamespace(hf_eos_token_id={2}, vocab_size=128),
|
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.extra_key, "classification")
|
||||||
self.assertEqual(req.cache_salt, "tenant-a")
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import time
|
|||||||
from array import array
|
from array import array
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
@@ -23,8 +23,11 @@ from sglang.srt.disaggregation.encoder.preprocessor import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.encoder.receiver import (
|
from sglang.srt.disaggregation.encoder.receiver import (
|
||||||
EmbeddingData,
|
EmbeddingData,
|
||||||
|
MMReceiverGrpc,
|
||||||
MMReceiverHTTP,
|
MMReceiverHTTP,
|
||||||
MultiModalEmbeddingData,
|
MultiModalEmbeddingData,
|
||||||
|
WaitingMMRequestStatus,
|
||||||
|
WaitingZmqRequest,
|
||||||
_encoder_media_item,
|
_encoder_media_item,
|
||||||
_select_mm_processor_prompt,
|
_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():
|
def test_kimi_k3_epd_aggregates_original_image_sizes_in_part_order():
|
||||||
first = EmbeddingData(
|
first = EmbeddingData(
|
||||||
req_id="request",
|
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():
|
def test_kimi_k3_encoder_prefers_grid_thws_and_uses_temporal_pool_length():
|
||||||
grid_thws = torch.tensor([[3, 8, 12]])
|
grid_thws = torch.tensor([[3, 8, 12]])
|
||||||
stale_grid = torch.tensor([[1, 2, 2]])
|
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():
|
def test_epd_scheduler_routes_many_requests_over_one_receive_socket():
|
||||||
context = zmq.Context()
|
context = zmq.Context()
|
||||||
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
|
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
|
||||||
@@ -761,6 +1038,8 @@ def test_epd_scheduler_routes_many_requests_over_one_receive_socket():
|
|||||||
sender = context.socket(zmq.PUSH)
|
sender = context.socket(zmq.PUSH)
|
||||||
try:
|
try:
|
||||||
sender.connect(f"tcp://127.0.0.1:{port}")
|
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):
|
for i in range(32):
|
||||||
mm_data = EmbeddingData(
|
mm_data = EmbeddingData(
|
||||||
req_id=f"rid-{i}_local_part_0",
|
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()
|
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():
|
def test_epd_encoder_reuses_scheduler_zmq_peer():
|
||||||
async def send_twice():
|
async def send_twice():
|
||||||
context = zmq.asyncio.Context()
|
context = zmq.asyncio.Context()
|
||||||
|
|||||||
@@ -129,6 +129,7 @@ def _make_tokenizer_manager(case) -> TokenizerManager:
|
|||||||
tm.server_args.dp_size = 1
|
tm.server_args.dp_size = 1
|
||||||
tm.disaggregation_mode = "none"
|
tm.disaggregation_mode = "none"
|
||||||
tm.rid_to_state = {}
|
tm.rid_to_state = {}
|
||||||
|
tm.encoder_dispatch_ready = {}
|
||||||
tm.enable_metrics = False
|
tm.enable_metrics = False
|
||||||
tm.enable_trace = False
|
tm.enable_trace = False
|
||||||
tm.enable_lora = 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()
|
processor._tokenizer = _Tokenizer()
|
||||||
embeddings = {Modality.IMAGE: torch.arange(20, dtype=torch.float32).reshape(5, 4)}
|
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],
|
[1, 99, 2, 99, 3],
|
||||||
embeddings,
|
embeddings,
|
||||||
img_grid_thw=torch.tensor([[1, 2, 6], [1, 2, 4]]),
|
img_grid_thw=torch.tensor([[1, 2, 6], [1, 2, 4]]),
|
||||||
|
|||||||
@@ -377,6 +377,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
|
|
||||||
manager = object.__new__(TokenizerManager)
|
manager = object.__new__(TokenizerManager)
|
||||||
manager.rid_to_state = {}
|
manager.rid_to_state = {}
|
||||||
|
manager.encoder_dispatch_ready = {}
|
||||||
transport = MagicMock()
|
transport = MagicMock()
|
||||||
transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
|
transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
|
||||||
manager.cuda_vmm_feature_transport = transport
|
manager.cuda_vmm_feature_transport = transport
|
||||||
@@ -405,6 +406,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
|
|
||||||
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
||||||
manager.rid_to_state = {}
|
manager.rid_to_state = {}
|
||||||
|
manager.encoder_dispatch_ready = {}
|
||||||
transport = MagicMock()
|
transport = MagicMock()
|
||||||
manager._dispatch_to_scheduler = MagicMock(
|
manager._dispatch_to_scheduler = MagicMock(
|
||||||
side_effect=RuntimeError("send failed")
|
side_effect=RuntimeError("send failed")
|
||||||
@@ -440,6 +442,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
|
|
||||||
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
||||||
manager.rid_to_state = {}
|
manager.rid_to_state = {}
|
||||||
|
manager.encoder_dispatch_ready = {}
|
||||||
transport = MagicMock()
|
transport = MagicMock()
|
||||||
manager._dispatch_to_scheduler = MagicMock()
|
manager._dispatch_to_scheduler = MagicMock()
|
||||||
time_stats = 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