fix(vlm): contain EPD request lifecycle failures (#36944)

Co-authored-by: mickqian <mickqian@users.noreply.github.com>
This commit is contained in:
Mick
2026-09-06 16:05:10 +08:00
committed by GitHub
co-authored by mickqian
parent d61378af77
commit 8ef646a5c6
8 changed files with 2435 additions and 139 deletions
@@ -12,6 +12,7 @@ import logging
import multiprocessing as mp import multiprocessing as mp
import traceback import traceback
from concurrent import futures from concurrent import futures
from http import HTTPStatus
from typing import List from typing import List
import grpc import grpc
@@ -21,7 +22,12 @@ from grpc_health.v1 import health_pb2, health_pb2_grpc
from grpc_reflection.v1alpha import reflection from grpc_reflection.v1alpha import reflection
from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc
from sglang.srt.disaggregation.encoder.server import MMEncoder, launch_encoder from sglang.srt.disaggregation.encoder.runtime import validate_encode_request
from sglang.srt.disaggregation.encoder.server import (
MMEncoder,
await_task_completion_on_cancel,
launch_encoder,
)
from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle
from sglang.srt.managers.schedule_batch import Modality from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
@@ -87,6 +93,11 @@ class SGLangEncoderServer(SGLangEncoderServicer):
self.send_sockets = send_sockets self.send_sockets = send_sockets
self.server_args = server_args self.server_args = server_args
async def _dispatch_encode(self, request_dict: dict):
for socket in self.send_sockets:
await async_sock_send(socket, wrap_as_pickle(request_dict))
return await self.encoder.encode_request(request_dict, Modality.IMAGE)
async def Encode( async def Encode(
self, request: sglang_encoder_pb2.EncodeRequest, context self, request: sglang_encoder_pb2.EncodeRequest, context
) -> sglang_encoder_pb2.EncodeResponse: ) -> sglang_encoder_pb2.EncodeResponse:
@@ -98,21 +109,33 @@ class SGLangEncoderServer(SGLangEncoderServicer):
"num_parts": request.num_parts, "num_parts": request.num_parts,
"part_idx": request.part_idx, "part_idx": request.part_idx,
} }
for socket in self.send_sockets: if err := validate_encode_request(request_dict):
await async_sock_send(socket, wrap_as_pickle(request_dict)) context.set_code(grpc.StatusCode.INVALID_ARGUMENT)
context.set_details(err)
return sglang_encoder_pb2.EncodeResponse()
# gRPC encode is image-only; the request follows the configured # gRPC encode is image-only; the request follows the configured
# cache and transfer backend. # cache and transfer backend. Keep TP dispatch and rank-0 collective
# launch order identical when gRPC handlers run concurrently.
async with self.encoder.encode_dispatch_lock:
encode_task = asyncio.create_task(self._dispatch_encode(request_dict))
result = await await_task_completion_on_cancel(
encode_task, f"Encoder request {request.req_id}"
)
( (
nbytes, nbytes,
embedding_len, embedding_len,
embedding_dim, embedding_dim,
error_msg, error_msg,
error_code, error_code,
) = await self.encoder.encode_request(request_dict, Modality.IMAGE) ) = result
if error_msg is not None: if error_msg is not None:
await self.encoder.release_request(request.req_id) await self.encoder.release_request(request.req_id)
context.set_code(grpc.StatusCode.INTERNAL) context.set_code(
grpc.StatusCode.INVALID_ARGUMENT
if error_code == HTTPStatus.BAD_REQUEST
else grpc.StatusCode.INTERNAL
)
context.set_details(error_msg) context.set_details(error_msg)
return sglang_encoder_pb2.EncodeResponse() return sglang_encoder_pb2.EncodeResponse()
@@ -154,6 +177,14 @@ class SGLangEncoderServer(SGLangEncoderServicer):
return sglang_encoder_pb2.EncodeResponse() return sglang_encoder_pb2.EncodeResponse()
except asyncio.CancelledError:
try:
await asyncio.shield(self.encoder.release_request(request.req_id))
except Exception:
logger.exception(
"Failed to release cancelled encoder request %s", request.req_id
)
raise
except Exception as e: except Exception as e:
logger.error(f"Encode error: {e}") logger.error(f"Encode error: {e}")
traceback.print_exc() traceback.print_exc()
@@ -178,6 +209,14 @@ class SGLangEncoderServer(SGLangEncoderServicer):
await self.encoder.release_request(request.req_id) await self.encoder.release_request(request.req_id)
return sglang_encoder_pb2.SendResponse() return sglang_encoder_pb2.SendResponse()
except asyncio.CancelledError:
try:
await asyncio.shield(self.encoder.release_request(request.req_id))
except Exception:
logger.exception(
"Failed to release cancelled encoder request %s", request.req_id
)
raise
except Exception as e: except Exception as e:
logger.error(f"Send error: {e}") logger.error(f"Send error: {e}")
traceback.print_exc() traceback.print_exc()
@@ -32,6 +32,8 @@ from sglang.srt.disaggregation.encoder.runtime import (
execute_encode_pipeline, execute_encode_pipeline,
launch_dp_runtime, launch_dp_runtime,
launch_local_runtime, launch_local_runtime,
send_staged_embedding,
validate_encode_request,
) )
from sglang.srt.disaggregation.encoder.server import ( from sglang.srt.disaggregation.encoder.server import (
EncoderProfiler, EncoderProfiler,
@@ -260,8 +262,38 @@ def _summarise_dp_broadcast(results: List[dict]) -> Response:
) )
async def _drain_health_encode(
health_encoder: MMEncoder, encode_task: asyncio.Task, req_id: str
):
"""Finish a dispatched TP probe before releasing its state and lock."""
result = None
cleanup_failed = False
try:
result = await asyncio.shield(encode_task)
except Exception:
logger.exception("Encoder health check failed for req_id=%s", req_id)
finally:
try:
await asyncio.shield(health_encoder.release_request(req_id))
except Exception:
cleanup_failed = True
logger.exception("Encoder health cleanup failed for req_id=%s", req_id)
finally:
health_encoder.encode_dispatch_lock.release()
return None if cleanup_failed else result
@app.post("/encode") @app.post("/encode")
async def handle_encode_request(request: dict): async def handle_encode_request(request: dict):
if err := validate_encode_request(request):
return ORJSONResponse(
status_code=HTTPStatus.BAD_REQUEST,
content={
"status": "error",
"message": err,
"req_id": request.get("req_id"),
},
)
req_id = request["req_id"] req_id = request["req_id"]
start_time = time.monotonic() start_time = time.monotonic()
time_stats_json = request.pop("time_stats_json", None) time_stats_json = request.pop("time_stats_json", None)
@@ -351,7 +383,6 @@ async def handle_send_request(request: dict):
"""Mooncake-only: drive the RDMA push of a staged embedding. The zmq """Mooncake-only: drive the RDMA push of a staged embedding. The zmq
backends deliver embeddings inline during /encode and never call /send.""" backends deliver embeddings inline during /encode and never call /send."""
req_id = request["req_id"] req_id = request["req_id"]
receive_count = request.get("receive_count")
if dp_dispatcher is not None: if dp_dispatcher is not None:
try: try:
result = await dp_dispatcher.dispatch_send(request) result = await dp_dispatcher.dispatch_send(request)
@@ -373,12 +404,22 @@ async def handle_send_request(request: dict):
status_code=status_code, status_code=status_code,
) )
return ORJSONResponse(content=result.get("content")) return ORJSONResponse(content=result.get("content"))
sent = await encoder.send( try:
req_id=req_id, sent = await send_staged_embedding(
prefill_host=request["prefill_host"], encoder,
embedding_port=request["embedding_port"], request,
session_id=request["session_id"], # A pre-refcount decoder may have sibling ranks still to send.
buffer_address=request["buffer_address"], release_without_count=False,
)
except Exception as error:
logger.error("Mooncake send failed for req_id=%s: %s", req_id, error)
return ORJSONResponse(
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
content={
"status": "error",
"message": str(error),
"req_id": req_id,
},
) )
if not sent: if not sent:
# No transfer happened: fail fast rather than 200 + a phantom count. # No transfer happened: fail fast rather than 200 + a phantom count.
@@ -390,11 +431,6 @@ async def handle_send_request(request: dict):
"req_id": req_id, "req_id": req_id,
}, },
) )
# Sibling ranks share this embedding, so free it only once all have sent.
# No count means a pre-refcount decoder: leave it to the sweep, as when
# some rank never sends at all.
if receive_count:
await server_module.meta_registry.note_send_done(req_id, receive_count)
return ORJSONResponse(content=None) return ORJSONResponse(content=None)
@@ -544,10 +580,10 @@ async def health_generate():
# No processor available, fall back to liveness check only # No processor available, fall back to liveness check only
return Response(status_code=200) return Response(status_code=200)
try:
# uuid keeps rids unique across workers; a bare time.time() can collide. # uuid keeps rids unique across workers; a bare time.time() can collide.
req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}" req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}"
owns_dispatch_lock = False
try:
dummy_request = { dummy_request = {
"mm_items": mm_items, "mm_items": mm_items,
"modality": modality.name, "modality": modality.name,
@@ -560,25 +596,36 @@ async def health_generate():
# request. Serialize its broadcast and rank-0 forward with every other # request. Serialize its broadcast and rank-0 forward with every other
# collective dispatch, then recheck whether traffic made the probe # collective dispatch, then recheck whether traffic made the probe
# unnecessary while it waited for the lock. # unnecessary while it waited for the lock.
async with encoder.encode_dispatch_lock: await encoder.encode_dispatch_lock.acquire()
owns_dispatch_lock = True
if encoder.has_pending_embeddings(): if encoder.has_pending_embeddings():
return Response(status_code=200) return Response(status_code=200)
for socket in send_sockets: for socket in send_sockets:
sock_send(socket, wrap_as_pickle(dummy_request)) sock_send(socket, wrap_as_pickle(dummy_request))
_, _, _, error_msg, _ = await asyncio.wait_for( encode_task = asyncio.create_task(
encoder.encode( encoder.encode(
mm_items=mm_items, mm_items=mm_items,
modality=modality, modality=modality,
req_id=req_id, req_id=req_id,
num_parts=1, num_parts=1,
part_idx=0, part_idx=0,
), )
)
drain_task = asyncio.create_task(
_drain_health_encode(encoder, encode_task, req_id)
)
# The drain task now owns the lock and request state. A probe timeout or
# client disconnect must not let a later request overtake its TP work.
owns_dispatch_lock = False
result = await asyncio.wait_for(
asyncio.shield(drain_task),
timeout=HEALTH_CHECK_TIMEOUT, timeout=HEALTH_CHECK_TIMEOUT,
) )
# Clean up stored embedding if result is None:
await encoder.release_request(req_id) return Response(status_code=503)
_, _, _, error_msg, _ = result
if error_msg: if error_msg:
logger.error(f"Encoder health check failed: {error_msg}") logger.error(f"Encoder health check failed: {error_msg}")
@@ -592,6 +639,9 @@ async def health_generate():
except Exception as e: except Exception as e:
logger.error(f"Encoder health check failed: {e}") logger.error(f"Encoder health check failed: {e}")
return Response(status_code=503) return Response(status_code=503)
finally:
if owns_dispatch_lock:
encoder.encode_dispatch_lock.release()
@app.api_route("/start_profile", methods=["GET", "POST"]) @app.api_route("/start_profile", methods=["GET", "POST"])
@@ -12,6 +12,7 @@ import contextlib
import logging import logging
import multiprocessing as mp import multiprocessing as mp
import os import os
import sys
import time import time
import traceback import traceback
import uuid import uuid
@@ -31,6 +32,7 @@ from sglang.srt.disaggregation.encoder.server import (
EncoderProfiler, EncoderProfiler,
MMEncoder, MMEncoder,
MMError, MMError,
await_task_completion_on_cancel,
launch_encoder, launch_encoder,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -76,6 +78,46 @@ class PendingRequest:
# vary per request and can't merge into one HF processor call. # vary per request and can't merge into one HF processor call.
_BATCHABLE_MODALITIES = {Modality.IMAGE, Modality.AUDIO} _BATCHABLE_MODALITIES = {Modality.IMAGE, Modality.AUDIO}
_KIMI_K3_DEFAULT_ENCODER_MAX_BATCH_SIZE = 2 _KIMI_K3_DEFAULT_ENCODER_MAX_BATCH_SIZE = 2
_DP_RELEASE_AFTER_ENCODE = "release_after_encode"
def validate_encode_request(request: dict) -> Optional[str]:
"""Return a client-facing error before an encode request is dispatched."""
if not isinstance(request, dict):
return f"request is not a dict: {type(request).__name__}"
req_id = request.get("req_id")
if not isinstance(req_id, str) or not req_id:
return "missing or invalid req_id"
modality = request.get("modality")
if not isinstance(modality, str):
return "missing or invalid modality"
try:
Modality.from_str(modality)
except ValueError:
return f"unsupported modality: {modality}"
mm_items = request.get("mm_items")
if mm_items is None or (isinstance(mm_items, (list, tuple)) and len(mm_items) == 0):
return "missing or empty mm_items"
num_parts = request.get("num_parts")
part_idx = request.get("part_idx")
if not isinstance(num_parts, int) or isinstance(num_parts, bool) or num_parts <= 0:
return "num_parts must be a positive integer"
if (
not isinstance(part_idx, int)
or isinstance(part_idx, bool)
or part_idx < 0
or part_idx >= num_parts
):
return f"part_idx must be in [0, {num_parts})"
hashes = request.get("hashes")
if hashes is not None and not isinstance(hashes, (list, tuple, str, int, bytes)):
return f"hashes must be list/scalar, got {type(hashes).__name__}"
return None
def _resolve_encoder_batch_policy( def _resolve_encoder_batch_policy(
@@ -203,27 +245,12 @@ class EncoderScheduler:
if not p.future.done(): if not p.future.done():
p.future.set_exception(e) p.future.set_exception(e)
@staticmethod
def _validate_request_shape(req: dict) -> Optional[str]:
# Cheap pre-broadcast checks: shape errors that don't require running
# the HF processor. Once a request reaches TP workers they enter
# batch_encode and expect to join its collectives — a malformed batch
# that makes rank-0 bail mid-flight would deadlock the workers.
if not isinstance(req, dict):
return f"request is not a dict: {type(req).__name__}"
if not req.get("req_id"):
return "missing req_id"
if not req.get("mm_items"):
return "missing or empty mm_items"
if "num_parts" not in req or "part_idx" not in req:
return "missing num_parts / part_idx"
h = req.get("hashes")
if h is not None and not isinstance(h, (list, tuple, str, int, bytes)):
return f"hashes must be list/scalar, got {type(h).__name__}"
return None
async def _dispatch_group( async def _dispatch_group(
self, group: List[PendingRequest], modality: Modality self,
group: List[PendingRequest],
modality: Modality,
*,
observe_queue_wait: bool = True,
) -> None: ) -> None:
# A request may time out while queued. Never start work that no caller # A request may time out while queued. Never start work that no caller
# can observe, or its eventual staged embedding would have no owner. # can observe, or its eventual staged embedding would have no owner.
@@ -241,7 +268,7 @@ class EncoderScheduler:
# abandoned. # abandoned.
valid: List[PendingRequest] = [] valid: List[PendingRequest] = []
for p in group: for p in group:
err = self._validate_request_shape(p.request) err = validate_encode_request(p.request)
if err is None: if err is None:
valid.append(p) valid.append(p)
continue continue
@@ -255,7 +282,7 @@ class EncoderScheduler:
requests = [p.request for p in group] requests = [p.request for p in group]
start = time.time() start = time.time()
modality_str = modality.name.lower() modality_str = modality.name.lower()
if server_module.encoder_metrics_collector is not None: if observe_queue_wait and server_module.encoder_metrics_collector is not None:
for p in group: for p in group:
server_module.encoder_metrics_collector.observe_queue_wait( server_module.encoder_metrics_collector.observe_queue_wait(
max(0.0, start - p.submit_time), modality=modality_str max(0.0, start - p.submit_time), modality=modality_str
@@ -309,6 +336,22 @@ class EncoderScheduler:
p.future.set_exception(err) p.future.set_exception(err)
return return
if len(group) > 1 and all(
result[3] is not None
and result[4] is not None
and int(result[4]) == HTTPStatus.BAD_REQUEST
for result in results
):
logger.warning(
f"Retrying failed {modality.name} batch as {len(group)} "
"individual requests"
)
for pending in group:
await self._dispatch_group(
[pending], modality, observe_queue_wait=False
)
return
for p, result in zip(group, results): for p, result in zip(group, results):
if not p.future.done(): if not p.future.done():
p.future.set_result(result) p.future.set_result(result)
@@ -324,6 +367,8 @@ class EncoderScheduler:
continue continue
req = p.request req = p.request
try: try:
if err := validate_encode_request(req):
raise server_module.BadRequestError(err)
start = time.time() start = time.time()
if server_module.encoder_metrics_collector is not None: if server_module.encoder_metrics_collector is not None:
server_module.encoder_metrics_collector.observe_queue_wait( server_module.encoder_metrics_collector.observe_queue_wait(
@@ -380,6 +425,7 @@ class DPDispatcher:
self, self,
dp_size: int, dp_size: int,
dispatch_sockets: List, dispatch_sockets: List,
release_sockets: List,
result_socket, result_socket,
worker_processes: List[mp.Process], worker_processes: List[mp.Process],
enable_metrics: bool = False, enable_metrics: bool = False,
@@ -387,6 +433,7 @@ class DPDispatcher:
): ):
self.dp_size = dp_size self.dp_size = dp_size
self.dispatch_sockets = dispatch_sockets self.dispatch_sockets = dispatch_sockets
self.release_sockets = release_sockets
self.result_socket = result_socket self.result_socket = result_socket
self.worker_processes = worker_processes self.worker_processes = worker_processes
# Key = req_id for encode/broadcast, or a per-control-request key for # Key = req_id for encode/broadcast, or a per-control-request key for
@@ -405,7 +452,7 @@ class DPDispatcher:
# Set when _result_listener gives up; makes alive_ranks report empty. # Set when _result_listener gives up; makes alive_ranks report empty.
self._listener_failed = False self._listener_failed = False
# The event loop only keeps weak references to tasks, so the long-lived # The event loop only keeps weak references to tasks, so the long-lived
# loops started in `start()` need a strong reference to survive GC. # loops and fire-and-forget notifications need a strong reference.
self.background_tasks: Set[asyncio.Task] = set() self.background_tasks: Set[asyncio.Task] = set()
# Prometheus gauge: pending requests per DP rank. Lives in the main # Prometheus gauge: pending requests per DP rank. Lives in the main
@@ -461,6 +508,28 @@ class DPDispatcher:
self.req_id_to_rank.pop(req_id, None) self.req_id_to_rank.pop(req_id, None)
self._update_pending_gauge() self._update_pending_gauge()
def _release_abandoned_encode(self, rank: int, req_id: str) -> None:
"""Tell the owning worker to release an encode that lost its caller."""
async def notify_worker() -> None:
try:
await async_sock_send(
self.release_sockets[rank],
wrap_as_pickle(
{"_dp_type": _DP_RELEASE_AFTER_ENCODE, "req_id": req_id}
),
)
except Exception:
logger.exception(
"Failed to retire abandoned encoder DP request %s on rank %s",
req_id,
rank,
)
task = asyncio.create_task(notify_worker())
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
@staticmethod @staticmethod
def _send_req_key(req_id: str, request: dict) -> str: def _send_req_key(req_id: str, request: dict) -> str:
"""One in-flight /send future per decoder TP rank, keyed by the rank's """One in-flight /send future per decoder TP rank, keyed by the rank's
@@ -546,6 +615,7 @@ class DPDispatcher:
future = asyncio.get_running_loop().create_future() future = asyncio.get_running_loop().create_future()
self.pending_futures[rank][req_id] = future self.pending_futures[rank][req_id] = future
self._update_pending_gauge() self._update_pending_gauge()
dispatched = False
logger.info( logger.info(
f"MM-Encoder DP dispatch: req_id={req_id}, " f"MM-Encoder DP dispatch: req_id={req_id}, "
f"modality={request.get('modality', 'image')}, " f"modality={request.get('modality', 'image')}, "
@@ -563,6 +633,7 @@ class DPDispatcher:
await async_sock_send( await async_sock_send(
self.dispatch_sockets[rank], wrap_as_pickle(request) self.dispatch_sockets[rank], wrap_as_pickle(request)
) )
dispatched = True
except BaseException: except BaseException:
self._drop_pending_and_mapping(rank, req_id) self._drop_pending_and_mapping(rank, req_id)
self._mapping_condition.notify_all() self._mapping_condition.notify_all()
@@ -574,6 +645,8 @@ class DPDispatcher:
future, timeout=server_module.ENCODER_REQ_TIMEOUT future, timeout=server_module.ENCODER_REQ_TIMEOUT
) )
except asyncio.TimeoutError: except asyncio.TimeoutError:
if dispatched:
self._release_abandoned_encode(rank, req_id)
self._drop_pending_and_mapping(rank, req_id) self._drop_pending_and_mapping(rank, req_id)
return self._timeout_envelope( return self._timeout_envelope(
req_id, req_id,
@@ -581,6 +654,8 @@ class DPDispatcher:
f"Encoder DP rank={rank} timed out after {server_module.ENCODER_REQ_TIMEOUT}s", f"Encoder DP rank={rank} timed out after {server_module.ENCODER_REQ_TIMEOUT}s",
) )
except BaseException: except BaseException:
if dispatched:
self._release_abandoned_encode(rank, req_id)
self._drop_pending_and_mapping(rank, req_id) self._drop_pending_and_mapping(rank, req_id)
raise raise
@@ -898,6 +973,12 @@ class DPDispatcher:
return return
await asyncio.sleep(min(0.1 * consecutive_errors, 1.0)) await asyncio.sleep(min(0.1 * consecutive_errors, 1.0))
continue continue
if not isinstance(msg, dict):
logger.error(
"_result_listener received a non-dict envelope (%s); dropping",
type(msg).__name__,
)
continue
req_id = msg.get("req_id", "") req_id = msg.get("req_id", "")
dp_type = msg.get("_dp_type", "encode") dp_type = msg.get("_dp_type", "encode")
if dp_type == "send": if dp_type == "send":
@@ -1062,6 +1143,49 @@ async def _push_embedding_to_prefill(
await enc.release_request(req_id) await enc.release_request(req_id)
async def send_staged_embedding(
enc: MMEncoder,
request: dict,
*,
release_without_count: bool,
) -> bool:
"""Send one Mooncake embedding and retire its state on any failure."""
req_id = request["req_id"]
try:
sent = await enc.send(
req_id=req_id,
prefill_host=request["prefill_host"],
embedding_port=request["embedding_port"],
session_id=request["session_id"],
buffer_address=request["buffer_address"],
)
if not sent:
return False
receive_count = request.get("receive_count")
if receive_count:
destination_endpoint = NetworkAddress(
request["prefill_host"], request["embedding_port"]
).to_host_port_str()
await server_module.meta_registry.note_send_done(
req_id, receive_count, destination_endpoint
)
elif release_without_count:
await enc.release_request(req_id)
return True
except BaseException as error:
try:
await enc.release_request(req_id)
except Exception as cleanup_error:
if sys.version_info >= (3, 11):
error.add_note(f"Failed to release encoder request: {cleanup_error}")
else:
logger.exception(
"Failed to release encoder request %s after send failure", req_id
)
raise
def _record_pipeline_result(modality: Modality, status: str) -> None: def _record_pipeline_result(modality: Modality, status: str) -> None:
if server_module.encoder_metrics_collector is not None: if server_module.encoder_metrics_collector is not None:
server_module.encoder_metrics_collector.inc_requests_total( server_module.encoder_metrics_collector.inc_requests_total(
@@ -1069,6 +1193,48 @@ def _record_pipeline_result(modality: Modality, status: str) -> None:
) )
async def _publish_pipeline_error(req_id: str, error_msg: str) -> bool:
"""Report a request error without letting reporting block cleanup."""
try:
await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg)
return True
except Exception:
logger.exception("Failed to publish encoder error for req_id=%s", req_id)
return False
async def _release_failed_request(
enc: MMEncoder,
req_id: str,
*,
preserve_metadata: bool = False,
) -> None:
"""Release request resources without hiding the original request error."""
try:
await enc.release_request(req_id, preserve_metadata=preserve_metadata)
except Exception:
logger.exception("Failed to release encoder resources for req_id=%s", req_id)
async def _run_dispatched_encode(
enc: MMEncoder, request: dict, modality: Modality
) -> Tuple:
"""Finish TP encode collectives before propagating caller cancellation."""
encode_task = asyncio.create_task(
enc.encode(
mm_items=request["mm_items"],
modality=modality,
req_id=request["req_id"],
num_parts=request["num_parts"],
part_idx=request["part_idx"],
hashes=request.get("hashes"),
)
)
return await await_task_completion_on_cancel(
encode_task, f"Encoder request {request['req_id']}"
)
async def execute_encode_pipeline( async def execute_encode_pipeline(
enc: MMEncoder, enc: MMEncoder,
sched: Optional[EncoderScheduler], sched: Optional[EncoderScheduler],
@@ -1082,6 +1248,8 @@ async def execute_encode_pipeline(
and keeps the result until follow-up /send calls complete. ZMQ has no early and keeps the result until follow-up /send calls complete. ZMQ has no early
consumer: it waits for encode, sends the embedding, releases it, then returns. consumer: it waits for encode, sends the embedding, releases it, then returns.
""" """
if err := validate_encode_request(request):
raise server_module.BadRequestError(err)
req_id = request["req_id"] req_id = request["req_id"]
time_stats_json = request.pop("time_stats_json", None) time_stats_json = request.pop("time_stats_json", None)
time_stats = EncoderReqTimeStats() time_stats = EncoderReqTimeStats()
@@ -1111,14 +1279,7 @@ async def execute_encode_pipeline(
async with enc.encode_dispatch_lock: async with enc.encode_dispatch_lock:
for socket in send_sockets: for socket in send_sockets:
sock_send(socket, wrap_as_pickle(request)) sock_send(socket, wrap_as_pickle(request))
result = await enc.encode( result = await _run_dispatched_encode(enc, request, modality)
mm_items=request["mm_items"],
modality=modality,
req_id=request["req_id"],
num_parts=request["num_parts"],
part_idx=request["part_idx"],
hashes=request.get("hashes"),
)
else: else:
result = await enc.encode( result = await enc.encode(
mm_items=request["mm_items"], mm_items=request["mm_items"],
@@ -1128,27 +1289,48 @@ async def execute_encode_pipeline(
part_idx=request["part_idx"], part_idx=request["part_idx"],
hashes=request.get("hashes"), hashes=request.get("hashes"),
) )
except asyncio.CancelledError:
error_msg = "encoder request cancelled"
time_stats.trace_ctx.abort(abort_info={"reason": error_msg})
try:
await asyncio.shield(enc.release_request(req_id))
except Exception:
logger.exception("Failed to release cancelled encoder request %s", req_id)
_record_pipeline_result(modality, "error")
raise
except asyncio.TimeoutError: except asyncio.TimeoutError:
error_msg = "encoder batch timed out" error_msg = "encoder batch timed out"
time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) time_stats.trace_ctx.abort(abort_info={"reason": error_msg})
await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) error_published = await _publish_pipeline_error(req_id, error_msg)
await enc.release_request(req_id, preserve_metadata=backend == "mooncake") await _release_failed_request(
enc,
req_id,
preserve_metadata=backend == "mooncake" and error_published,
)
_record_pipeline_result(modality, "error") _record_pipeline_result(modality, "error")
raise raise
except Exception as e: except Exception as e:
error_msg = str(e) error_msg = str(e)
time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) time_stats.trace_ctx.abort(abort_info={"reason": error_msg})
await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) error_published = await _publish_pipeline_error(req_id, error_msg)
await enc.release_request(req_id, preserve_metadata=backend == "mooncake") await _release_failed_request(
enc,
req_id,
preserve_metadata=backend == "mooncake" and error_published,
)
_record_pipeline_result(modality, "error") _record_pipeline_result(modality, "error")
raise raise
nbytes, embedding_len, embedding_dim, error_msg, error_code = result nbytes, embedding_len, embedding_dim, error_msg, error_code = result
if error_msg: if error_msg:
time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) time_stats.trace_ctx.abort(abort_info={"reason": error_msg})
await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) error_published = await _publish_pipeline_error(req_id, error_msg)
if backend == "mooncake": if backend == "mooncake":
await enc.release_request(req_id, preserve_metadata=True) await _release_failed_request(
enc,
req_id,
preserve_metadata=error_published,
)
else: else:
try: try:
await _push_embedding_to_prefill( await _push_embedding_to_prefill(
@@ -1161,6 +1343,7 @@ async def execute_encode_pipeline(
f"Error-send failed for req_id={req_id}: {send_err}", f"Error-send failed for req_id={req_id}: {send_err}",
exc_info=True, exc_info=True,
) )
await _release_failed_request(enc, req_id)
_record_pipeline_result(modality, "error") _record_pipeline_result(modality, "error")
raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR)
@@ -1287,12 +1470,10 @@ async def _dp_worker_handle_request(
) from e ) from e
elif dp_type == "send": elif dp_type == "send":
req_id = request["req_id"] req_id = request["req_id"]
sent = await enc.send( sent = await send_staged_embedding(
req_id=req_id, enc,
prefill_host=request["prefill_host"], request,
embedding_port=request["embedding_port"], release_without_count=True,
session_id=request["session_id"],
buffer_address=request["buffer_address"],
) )
if not sent: if not sent:
# Error envelope, not 200 + phantom count: the decoder must # Error envelope, not 200 + phantom count: the decoder must
@@ -1300,13 +1481,6 @@ async def _dp_worker_handle_request(
raise MMError( raise MMError(
f"no staged embedding for /send req_id={req_id} (already released)" f"no staged embedding for /send req_id={req_id} (already released)"
) )
# Releasing on the first /send breaks decoder TP > 1. No count means
# a pre-refcount decoder: stay eager rather than pin until the sweep.
receive_count = request.get("receive_count")
if receive_count:
await server_module.meta_registry.note_send_done(req_id, receive_count)
else:
await enc.release_request(req_id)
content = None content = None
else: else:
content = await execute_encode_pipeline(enc, sched, request) content = await execute_encode_pipeline(enc, sched, request)
@@ -1337,7 +1511,12 @@ async def _dp_worker_handle_request(
f"req_id={request.get('req_id', '?')}: {e}", f"req_id={request.get('req_id', '?')}: {e}",
exc_info=True, exc_info=True,
) )
err_code = int(getattr(e, "code", None) or HTTPStatus.INTERNAL_SERVER_ERROR) # Only MMError carries an HTTP status in this protocol. Third-party
# exceptions may expose a callable ``code`` attribute (for example
# gRPC errors), which must not make error reporting fail a second time.
err_code = int(
e.code if isinstance(e, MMError) else HTTPStatus.INTERNAL_SERVER_ERROR
)
envelope = { envelope = {
"req_id": request.get("req_id", ""), "req_id": request.get("req_id", ""),
"_dp_type": dp_type, "_dp_type": dp_type,
@@ -1368,11 +1547,27 @@ async def _dp_worker_handle_request(
) )
async def _retire_abandoned_encode(
enc: MMEncoder,
encode_task: Optional[asyncio.Task],
req_id: str,
) -> None:
"""Retire an abandoned request without interrupting its encode work."""
try:
if encode_task is None or not encode_task.done():
await enc.abandon_request(req_id)
else:
await enc.release_request(req_id)
except Exception:
logger.exception("Failed to release abandoned encoder DP request %s", req_id)
async def run_dp_worker( async def run_dp_worker(
server_args: ServerArgs, server_args: ServerArgs,
dp_rank: int, dp_rank: int,
gpu_id: int, gpu_id: int,
dispatch_path: str, dispatch_path: str,
release_path: str,
result_path: str, result_path: str,
): ):
logger.info( logger.info(
@@ -1414,9 +1609,34 @@ async def run_dp_worker(
ctx = zmq.asyncio.Context(2) ctx = zmq.asyncio.Context(2)
recv_sock = get_zmq_socket(ctx, zmq.PULL, dispatch_path, False) recv_sock = get_zmq_socket(ctx, zmq.PULL, dispatch_path, False)
release_sock = get_zmq_socket(ctx, zmq.PULL, release_path, False)
send_sock = get_zmq_socket(ctx, zmq.PUSH, result_path, False) send_sock = get_zmq_socket(ctx, zmq.PUSH, result_path, False)
send_lock = asyncio.Lock() send_lock = asyncio.Lock()
inflight: Set[asyncio.Task] = set() inflight: Set[asyncio.Task] = set()
encode_tasks: Dict[str, asyncio.Task] = {}
release_tasks: Set[asyncio.Task] = set()
async def listen_for_releases() -> None:
# Cleanup must not wait behind the bounded encode queue: under a
# cancellation burst every normal worker slot may already be occupied.
while True:
try:
request = await async_sock_recv(release_sock)
except asyncio.CancelledError:
raise
except Exception:
logger.error(f"DP worker {dp_rank} release recv error", exc_info=True)
continue
if not isinstance(request, dict) or not request.get("req_id"):
logger.error(f"DP worker {dp_rank} received malformed release request")
continue
req_id = request["req_id"]
task = asyncio.create_task(
_retire_abandoned_encode(enc, encode_tasks.get(req_id), req_id)
)
release_tasks.add(task)
task.add_done_callback(release_tasks.discard)
# Acquire-before-recv → back-pressure propagates to the dispatcher # Acquire-before-recv → back-pressure propagates to the dispatcher
# PUSH buffer. Must be at least max_batch_size or batching degrades. # PUSH buffer. Must be at least max_batch_size or batching degrades.
max_inflight = envs.SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT.get() max_inflight = envs.SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT.get()
@@ -1428,6 +1648,7 @@ async def run_dp_worker(
) )
inflight_sem = asyncio.Semaphore(max_inflight) inflight_sem = asyncio.Semaphore(max_inflight)
sched.start() sched.start()
release_listener_task = asyncio.create_task(listen_for_releases())
logger.info(f"DP worker {dp_rank} ready") logger.info(f"DP worker {dp_rank} ready")
try: try:
@@ -1462,12 +1683,30 @@ async def run_dp_worker(
spawned = True spawned = True
inflight.add(task) inflight.add(task)
task.add_done_callback(inflight.discard) task.add_done_callback(inflight.discard)
if dp_type == "encode":
req_id = request["req_id"]
encode_tasks[req_id] = task
def forget_encode_task(
completed_task: asyncio.Task, request_id: str = req_id
) -> None:
if encode_tasks.get(request_id) is completed_task:
encode_tasks.pop(request_id, None)
enc.clear_abandoned_request(request_id)
task.add_done_callback(forget_encode_task)
finally: finally:
if not spawned: if not spawned:
inflight_sem.release() inflight_sem.release()
finally: finally:
release_listener_task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await release_listener_task
for task in inflight: for task in inflight:
task.cancel() task.cancel()
for task in release_tasks:
task.cancel()
await asyncio.gather(*inflight, *release_tasks, return_exceptions=True)
ctx.destroy(linger=0) ctx.destroy(linger=0)
@@ -1476,13 +1715,21 @@ def launch_dp_worker(
dp_rank: int, dp_rank: int,
gpu_id: int, gpu_id: int,
dispatch_path: str, dispatch_path: str,
release_path: str,
result_path: str, result_path: str,
): ):
publish(server_args, role="encoder") publish(server_args, role="encoder")
try: try:
configure_logger(server_args, prefix=f" encode_dp_worker[{dp_rank}]") configure_logger(server_args, prefix=f" encode_dp_worker[{dp_rank}]")
asyncio.run( asyncio.run(
run_dp_worker(server_args, dp_rank, gpu_id, dispatch_path, result_path) run_dp_worker(
server_args,
dp_rank,
gpu_id,
dispatch_path,
release_path,
result_path,
)
) )
except KeyboardInterrupt: except KeyboardInterrupt:
logger.info(f"DP worker {dp_rank} exiting") logger.info(f"DP worker {dp_rank} exiting")
@@ -1599,6 +1846,12 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher:
) )
for r in range(dp_size) for r in range(dp_size)
] ]
release_sockets: List[zmq.asyncio.Socket] = [
get_zmq_socket(
async_zmq_ctx, zmq.PUSH, f"ipc:///tmp/{ipc_prefix}_dp_release_{r}", True
)
for r in range(dp_size)
]
worker_processes: List[mp.Process] = [] worker_processes: List[mp.Process] = []
@@ -1628,6 +1881,7 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher:
dp_rank, dp_rank,
gpu_id, gpu_id,
f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{dp_rank}", f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{dp_rank}",
f"ipc:///tmp/{ipc_prefix}_dp_release_{dp_rank}",
result_path, result_path,
), ),
daemon=False, daemon=False,
@@ -1641,6 +1895,7 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher:
return DPDispatcher( return DPDispatcher(
dp_size, dp_size,
dispatch_sockets, dispatch_sockets,
release_sockets,
result_socket, result_socket,
worker_processes, worker_processes,
enable_metrics=get_observability().enable_metrics, enable_metrics=get_observability().enable_metrics,
@@ -1,6 +1,7 @@
import asyncio import asyncio
import concurrent.futures import concurrent.futures
import ctypes import ctypes
import hashlib
import logging import logging
import os import os
import pickle import pickle
@@ -87,6 +88,7 @@ rid_to_receive_endpoint: Dict[str, Set[str]] = dict()
rid_to_receive_count: Dict[str, int] = dict() rid_to_receive_count: Dict[str, int] = dict()
cond_dict_lock = asyncio.Lock() cond_dict_lock = asyncio.Lock()
rid_to_cond: Dict[str, asyncio.Condition] = {} rid_to_cond: Dict[str, asyncio.Condition] = {}
encode_state_condition = asyncio.Condition()
async def _get_receive_condition(req_id: str) -> asyncio.Condition: async def _get_receive_condition(req_id: str) -> asyncio.Condition:
@@ -96,6 +98,15 @@ async def _get_receive_condition(req_id: str) -> asyncio.Condition:
return rid_to_cond[req_id] return rid_to_cond[req_id]
async def _notify_receive_waiters(req_id: str) -> None:
"""Wake an existing destination waiter without creating new state."""
async with cond_dict_lock:
cond = rid_to_cond.get(req_id)
if cond is not None:
async with cond:
cond.notify_all()
ENCODER_MAX_BATCH_SIZE = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.get() ENCODER_MAX_BATCH_SIZE = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.get()
ENCODER_MAX_BATCH_SIZE_EXPLICIT = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.is_set() ENCODER_MAX_BATCH_SIZE_EXPLICIT = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.is_set()
# Watchdog: max time to wait for a batched /encode result. Bounds HTTP latency # Watchdog: max time to wait for a batched /encode result. Bounds HTTP latency
@@ -103,6 +114,34 @@ ENCODER_MAX_BATCH_SIZE_EXPLICIT = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.is_set()
ENCODER_REQ_TIMEOUT = envs.SGLANG_ENCODER_REQ_TIMEOUT.get() ENCODER_REQ_TIMEOUT = envs.SGLANG_ENCODER_REQ_TIMEOUT.get()
async def await_task_completion_on_cancel(task: asyncio.Task, operation: str):
"""Keep task-owned resources live until cancellation reaches a safe point."""
try:
return await asyncio.shield(task)
except asyncio.CancelledError:
while not task.done():
try:
await asyncio.shield(task)
except asyncio.CancelledError:
continue
except Exception:
break
if not task.cancelled() and task.exception() is not None:
logger.error(
"%s failed while draining cancellation",
operation,
exc_info=task.exception(),
)
raise
async def _await_transfer_completion(awaitable, operation: str):
"""Do not let cancellation outlive a zero-copy transfer using its buffer."""
return await await_task_completion_on_cancel(
asyncio.ensure_future(awaitable), operation
)
class EncoderMetaRegistry: class EncoderMetaRegistry:
"""Per-part metadata shared by every encoder request lifecycle. """Per-part metadata shared by every encoder request lifecycle.
@@ -117,9 +156,10 @@ class EncoderMetaRegistry:
# Backstop for state whose /send calls never all land. # Backstop for state whose /send calls never all land.
self.sweep_timeout = sweep_timeout self.sweep_timeout = sweep_timeout
self._rid_to_meta: Dict[str, dict] = {} self._rid_to_meta: Dict[str, dict] = {}
self._rid_to_send_done: Dict[str, int] = {} self._rid_to_send_done: Dict[str, Set[str]] = {}
self._pending_at: Dict[str, float] = {} self._pending_at: Dict[str, float] = {}
self._sweeper_task: Optional[asyncio.Task] = None self._sweeper_task: Optional[asyncio.Task] = None
self._stale_release_tasks: Dict[str, asyncio.Task] = {}
# Set only where the embedding also lives; None in the DP main process. # Set only where the embedding also lives; None in the DP main process.
self.on_release: Optional[Callable[[str], Awaitable[None]]] = None self.on_release: Optional[Callable[[str], Awaitable[None]]] = None
@@ -146,9 +186,36 @@ class EncoderMetaRegistry:
rid rid
for rid, ts in self._pending_at.items() for rid, ts in self._pending_at.items()
if now - ts > self.sweep_timeout if now - ts > self.sweep_timeout
and rid not in self._stale_release_tasks
] ]
for rid in stale: for rid in stale:
await self._release(rid) self._schedule_stale_release(rid)
def _schedule_stale_release(self, req_id: str) -> asyncio.Task:
"""Release one stale request without blocking cleanup of other requests."""
if task := self._stale_release_tasks.get(req_id):
return task
task = asyncio.create_task(self._release_stale(req_id))
self._stale_release_tasks[req_id] = task
task.add_done_callback(
lambda done, rid=req_id: self._finish_stale_release(rid, done)
)
return task
def _finish_stale_release(self, req_id: str, task: asyncio.Task) -> None:
if self._stale_release_tasks.get(req_id) is task:
self._stale_release_tasks.pop(req_id)
async def _release_stale(self, req_id: str) -> None:
try:
await self._release(req_id)
except Exception:
logger.exception("Failed to release stale encoder request %s", req_id)
# Keep the request eligible for a later sweep without retrying in a
# tight loop. Its metadata and buffer ownership remain intact.
async with rid_lock:
if req_id in self._pending_at:
self._pending_at[req_id] = time.monotonic()
async def publish( async def publish(
self, self,
@@ -187,12 +254,15 @@ class EncoderMetaRegistry:
) )
return self._rid_to_meta.get(req_id) return self._rid_to_meta.get(req_id)
async def note_send_done(self, req_id: str, receive_count: int) -> None: async def note_send_done(
"""Count one completed ``/send``; release everything at receive_count.""" self, req_id: str, receive_count: int, destination_endpoint: str
) -> None:
"""Count one destination once; release after every receiver has sent."""
async with rid_lock: async with rid_lock:
count = self._rid_to_send_done.get(req_id, 0) + 1 completed = self._rid_to_send_done.setdefault(req_id, set())
self._rid_to_send_done[req_id] = count completed.add(destination_endpoint)
if count >= receive_count: all_done = len(completed) >= receive_count
if all_done:
await self._release(req_id) await self._release(req_id)
async def _release(self, req_id: str) -> None: async def _release(self, req_id: str) -> None:
@@ -249,6 +319,36 @@ class EncodeContext(msgspec.Struct):
is_health_check: bool is_health_check: bool
def _preprocess_layout_digest(ctx: EncodeContext) -> tuple[int, int]:
"""Hash metadata that must agree before TP ranks enter model forward."""
def normalize(value):
if isinstance(value, torch.Tensor):
value = value.detach().cpu().numpy()
if isinstance(value, np.ndarray):
return (
str(value.dtype),
tuple(value.shape),
tuple(value.reshape(-1).tolist()),
)
if isinstance(value, (list, tuple)):
return tuple(normalize(item) for item in value)
if isinstance(value, np.generic):
return value.item()
return value
signature = (
tuple(ctx.items_per_req),
tuple(ctx.preprocess_result.token_counts),
normalize(ctx.preprocess_result.grid_thw),
)
digest = hashlib.blake2b(pickle.dumps(signature), digest_size=16).digest()
return (
int.from_bytes(digest[:8], byteorder="little", signed=True),
int.from_bytes(digest[8:], byteorder="little", signed=True),
)
@dataclass @dataclass
class ReqState: class ReqState:
"""The result and in-flight work for one encoder request.""" """The result and in-flight work for one encoder request."""
@@ -582,6 +682,9 @@ class MMEncoder:
) )
self.req_states: Dict[str, ReqState] = {} self.req_states: Dict[str, ReqState] = {}
# A DP caller can disappear before its encode creates ReqState.
# Preserve that release intent until _acquire_encode_ref runs.
self.abandoned_req_ids: Set[str] = set()
# Need to ensure the NCCL launch order on rank0 matches the dispatch order rank>0 # Need to ensure the NCCL launch order on rank0 matches the dispatch order rank>0
self.encode_dispatch_lock = asyncio.Lock() self.encode_dispatch_lock = asyncio.Lock()
@@ -641,8 +744,22 @@ class MMEncoder:
state = ReqState(req_id) state = ReqState(req_id)
self.req_states[req_id] = state self.req_states[req_id] = state
state.active_encodes += 1 state.active_encodes += 1
if req_id in self.abandoned_req_ids:
state.release_requested = True
self.abandoned_req_ids.discard(req_id)
return state return state
async def abandon_request(self, req_id: str) -> None:
"""Release now, or remember the release until encode state exists."""
self.abandoned_req_ids.add(req_id)
if req_id in self.req_states:
self.abandoned_req_ids.discard(req_id)
await self.release_request(req_id)
def clear_abandoned_request(self, req_id: str) -> None:
"""Drop an unused release marker after the worker task exits."""
self.abandoned_req_ids.discard(req_id)
async def _release_encode_ref(self, state: Optional[ReqState]) -> None: async def _release_encode_ref(self, state: Optional[ReqState]) -> None:
if state is None: if state is None:
return return
@@ -718,8 +835,16 @@ class MMEncoder:
async with state.lifecycle_condition: async with state.lifecycle_condition:
state.release_requested = True state.release_requested = True
state.preserve_metadata_on_release |= preserve_metadata state.preserve_metadata_on_release |= preserve_metadata
if state.active_encodes > 0: encode_is_active = state.active_encodes > 0
# ``send_with_url`` may be waiting for a destination that will never
# arrive after its HTTP caller disappears. Wake it so the worker slot
# is retired together with the staged embedding.
await _notify_receive_waiters(req_id)
if encode_is_active:
return return
async with state.lifecycle_condition:
await state.lifecycle_condition.wait_for(lambda: state.active_sends == 0) await state.lifecycle_condition.wait_for(lambda: state.active_sends == 0)
if self.req_states.get(req_id) is not state: if self.req_states.get(req_id) is not state:
return return
@@ -735,6 +860,29 @@ class MMEncoder:
expected_destination_count: int, expected_destination_count: int,
destination_urls: Iterable[str], destination_urls: Iterable[str],
) -> None: ) -> None:
state = self.req_states.get(req_id)
if state is None:
# registration can beat /encode or its queued batch; only encode creates state
try:
async with encode_state_condition:
await asyncio.wait_for(
encode_state_condition.wait_for(
lambda: req_id in self.req_states
),
timeout=ENCODER_REQ_TIMEOUT,
)
except asyncio.TimeoutError as exc:
raise MMError(
f"Timed out waiting for encoder request to start: {req_id}",
code=HTTPStatus.GATEWAY_TIMEOUT,
) from exc
state = self.req_states.get(req_id)
if state is None:
raise BadRequestError(f"Encoder request is not active: {req_id}")
async with state.lifecycle_condition:
if self.req_states.get(req_id) is not state or state.release_requested:
raise BadRequestError(f"Encoder request is not active: {req_id}")
async with rid_lock: async with rid_lock:
if req_id not in rid_to_receive_endpoint: if req_id not in rid_to_receive_endpoint:
rid_to_receive_endpoint[req_id] = set() rid_to_receive_endpoint[req_id] = set()
@@ -746,7 +894,6 @@ class MMEncoder:
f"registered {registered_count}, got {expected_destination_count}" f"registered {registered_count}, got {expected_destination_count}"
) )
rid_to_receive_endpoint[req_id].update(destination_urls) rid_to_receive_endpoint[req_id].update(destination_urls)
cond = await _get_receive_condition(req_id) cond = await _get_receive_condition(req_id)
async with cond: async with cond:
cond.notify_all() cond.notify_all()
@@ -974,10 +1121,14 @@ class MMEncoder:
preprocess_result, preprocess_result,
items_per_req, items_per_req,
) = await self.preprocessor.process_batch_mm_items(requests, modality) ) = await self.preprocessor.process_batch_mm_items(requests, modality)
except MMError:
raise
except NotImplementedError as e: except NotImplementedError as e:
raise InternalError(f"Not implemented error: {str(e)}") raise InternalError(f"Not implemented error: {str(e)}")
except Exception as e: except (TypeError, ValueError) as e:
raise BadRequestError(f"Failed to process mm items: {str(e)}") raise BadRequestError(f"Failed to process mm items: {str(e)}")
except Exception as e:
raise InternalError(f"Failed to process mm items: {str(e)}")
if len(items_per_req) != len(requests) or any(n <= 0 for n in items_per_req): if len(items_per_req) != len(requests) or any(n <= 0 for n in items_per_req):
raise InternalError( raise InternalError(
@@ -1053,6 +1204,134 @@ class MMEncoder:
is_health_check=is_health_check, is_health_check=is_health_check,
) )
async def _prepare_encode_context_on_all_ranks(
self,
requests: List[dict],
modality: Modality,
*,
use_global_cache: bool,
is_health_check: bool = False,
) -> EncodeContext:
"""Prepare one context consistently before TP ranks enter model forward."""
ctx = None
local_error = None
error_phase = 0
try:
ctx = await self._prepare_encode_context(
requests,
modality,
use_global_cache=use_global_cache,
is_health_check=is_health_check,
)
except Exception as e:
local_error = e
error_phase = 1
if local_error is None:
try:
assert ctx is not None
await self._publish_preprocess_metadata(ctx, requests)
except Exception as e:
local_error = e
error_phase = 2
if local_error is None:
assert ctx is not None
layout_digest = _preprocess_layout_digest(ctx)
else:
layout_digest = (0, 0)
statuses = self._sync_tp_prepare_status(
local_error,
error_phase=error_phase,
layout_digest=layout_digest,
)
expected_layout = tuple(statuses[0][2:].tolist())
mismatch_rank = next(
(
rank
for rank, rank_status in enumerate(statuses[1:], start=1)
if tuple(rank_status[2:].tolist()) != expected_layout
),
None,
)
if mismatch_rank is not None:
raise InternalError(
"Encoder preprocessing produced inconsistent layouts across TP "
f"ranks 0 and {mismatch_rank}"
)
assert ctx is not None
return ctx
def _sync_tp_prepare_status(
self,
local_error: Optional[Exception],
*,
error_phase: int,
layout_digest: tuple[int, int],
) -> List[torch.Tensor]:
"""Raise the same preparation error on every TP rank."""
tp_group = get_tp_group()
error_code = (
int(
local_error.code
if isinstance(local_error, MMError)
else HTTPStatus.INTERNAL_SERVER_ERROR
)
if local_error is not None
else 0
)
local_status = torch.tensor(
[error_code, error_phase, *layout_digest], dtype=torch.int64
)
statuses = [torch.empty_like(local_status) for _ in range(tp_group.world_size)]
if tp_group.world_size > 1:
torch.distributed.all_gather(
statuses,
local_status,
group=tp_group.cpu_group,
)
else:
statuses[0].copy_(local_status)
failures = [
(
rank,
int(rank_status[0].item()),
int(rank_status[1].item()),
)
for rank, rank_status in enumerate(statuses)
if rank_status[0].item() != 0
]
if not failures:
return statuses
errors = (
tp_group.all_gather_object(
str(local_error) if local_error is not None else None
)
if tp_group.world_size > 1
else [str(local_error)]
)
rank, failure_code, failure_phase = next(
(
(rank, rank_error_code, rank_error_phase)
for rank, rank_error_code, rank_error_phase in failures
if rank_error_code != HTTPStatus.BAD_REQUEST
),
failures[0],
)
phase = (
"Encoder metadata publication"
if failure_phase == 2
else "Encoder preprocessing"
)
message = f"{phase} failed on TP rank {rank}: {errors[rank]}"
if failure_code == HTTPStatus.BAD_REQUEST:
raise BadRequestError(message)
raise InternalError(message)
def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor):
if get_parallel().tp_size > 1: if get_parallel().tp_size > 1:
torch.distributed.broadcast( torch.distributed.broadcast(
@@ -1737,7 +2016,10 @@ class MMEncoder:
# Queue sends in order under the lock, then wait for buffer # Queue sends in order under the lock, then wait for buffer
# ownership independently so libzmq can pipeline the connection. # ownership independently so libzmq can pipeline the connection.
try: try:
await asyncio.to_thread(tracker.wait, self.send_timeout) await _await_transfer_completion(
asyncio.to_thread(tracker.wait, self.send_timeout),
f"ZMQ transfer for req_id={mm_data.req_id}",
)
except Exception: except Exception:
if self.scheduler_send_sockets.get(endpoint) is sock: if self.scheduler_send_sockets.get(endpoint) is sock:
self.scheduler_send_sockets.pop(endpoint, None) self.scheduler_send_sockets.pop(endpoint, None)
@@ -1772,7 +2054,10 @@ class MMEncoder:
finally: finally:
sock.close(linger=5000) sock.close(linger=5000)
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket) await _await_transfer_completion(
asyncio.get_running_loop().run_in_executor(self.executor, send_with_socket),
f"ZMQ transfer for req_id={mm_data.req_id}",
)
if ( if (
encoder_metrics_collector is not None encoder_metrics_collector is not None
and get_disagg().encoder_transfer_backend != "mooncake" and get_disagg().encoder_transfer_backend != "mooncake"
@@ -1790,23 +2075,16 @@ class MMEncoder:
size: int, size: int,
) -> int: ) -> int:
"""Keep the send active until its blocking transfer stops using the MR.""" """Keep the send active until its blocking transfer stops using the MR."""
transfer_task = asyncio.create_task( return await _await_transfer_completion(
asyncio.to_thread( asyncio.to_thread(
self.engine.transfer_sync, self.engine.transfer_sync,
session_id, session_id,
source_address, source_address,
destination_address, destination_address,
size, size,
),
f"Mooncake transfer to session={session_id}",
) )
)
try:
return await asyncio.shield(transfer_task)
except asyncio.CancelledError:
try:
await transfer_task
except Exception:
pass
raise
def _register_shared_mr(self, mm_data: EmbeddingData, embedding: torch.Tensor): def _register_shared_mr(self, mm_data: EmbeddingData, embedding: torch.Tensor):
"""Register one MR shared by every rank's /send; _send re-registers on failure.""" """Register one MR shared by every rank's /send; _send re-registers on failure."""
@@ -1943,13 +2221,15 @@ class MMEncoder:
keep_on_gpu = self.use_mooncake and not is_health_check keep_on_gpu = self.use_mooncake and not is_health_check
use_global_cache = self.mm_global_cache is not None and not is_health_check use_global_cache = self.mm_global_cache is not None and not is_health_check
try: try:
ctx = await self._prepare_encode_context( if self.rank == 0:
async with encode_state_condition:
encode_state_condition.notify_all()
ctx = await self._prepare_encode_context_on_all_ranks(
requests, requests,
modality, modality,
use_global_cache=use_global_cache, use_global_cache=use_global_cache,
is_health_check=is_health_check, is_health_check=is_health_check,
) )
await self._publish_preprocess_metadata(ctx, requests)
mm_embedding = await self._compute_embedding(ctx, keep_on_gpu=keep_on_gpu) mm_embedding = await self._compute_embedding(ctx, keep_on_gpu=keep_on_gpu)
if self.profiler is not None: if self.profiler is not None:
@@ -2034,6 +2314,9 @@ class MMEncoder:
try: try:
while True: while True:
if state.release_requested:
break
async with rid_lock: async with rid_lock:
current_targets = rid_to_receive_endpoint.get(req_id, set()).copy() current_targets = rid_to_receive_endpoint.get(req_id, set()).copy()
expected_count = rid_to_receive_count.get(req_id) expected_count = rid_to_receive_count.get(req_id)
@@ -26,8 +26,7 @@ class TestOpenAICompletionRustParity(CustomTestCase):
api_key = "sk-123456" api_key = "sk-123456"
def _get_logprobs(self, *, rust_frontend): def _get_logprobs(self, *, rust_frontend):
# Prefill CUDA graph pads the batch, so numerics follow whichever # compare identical prefill shapes, without graph padding or warmup cache hits
# requests share the forward pass; the assertions below need equality.
process = popen_launch_server( process = popen_launch_server(
self.model, self.model,
DEFAULT_URL_FOR_TEST, DEFAULT_URL_FOR_TEST,
@@ -38,6 +37,7 @@ class TestOpenAICompletionRustParity(CustomTestCase):
"--random-seed", "--random-seed",
"42", "42",
"--disable-prefill-cuda-graph", "--disable-prefill-cuda-graph",
"--disable-radix-cache",
], ],
) )
try: try:
File diff suppressed because it is too large Load Diff
@@ -17,6 +17,8 @@ class _FakeEncoder:
self.embedding_to_send = {} self.embedding_to_send = {}
self.encode_dispatch_lock = asyncio.Lock() self.encode_dispatch_lock = asyncio.Lock()
self.encode_calls = [] self.encode_calls = []
self.released = []
self.release_event = asyncio.Event()
def has_pending_embeddings(self): def has_pending_embeddings(self):
return bool(self.embedding_to_send) return bool(self.embedding_to_send)
@@ -28,8 +30,9 @@ class _FakeEncoder:
self.encode_calls.append(kwargs) self.encode_calls.append(kwargs)
return 1, 1, 1, None, None return 1, 1, 1, None, None
async def release_request(self, _req_id): async def release_request(self, req_id):
return None self.released.append(req_id)
self.release_event.set()
def _install_tp_encoder(monkeypatch, encoder): def _install_tp_encoder(monkeypatch, encoder):
@@ -84,5 +87,87 @@ def test_health_encode_rechecks_busy_state_after_waiting(monkeypatch):
asyncio.run(run_test()) asyncio.run(run_test())
def test_health_timeout_keeps_dispatch_order_until_encode_drains(monkeypatch):
async def run_test():
encoder = _FakeEncoder()
_install_tp_encoder(monkeypatch, encoder)
encode_started = asyncio.Event()
finish_encode = asyncio.Event()
async def encode(**kwargs):
encoder.encode_calls.append(kwargs)
encode_started.set()
await finish_encode.wait()
return 1, 1, 1, None, None
encoder.encode = encode
monkeypatch.setattr(http_server, "HEALTH_CHECK_TIMEOUT", 0.01)
response = await http_server.health_generate()
assert response.status_code == 503
assert encode_started.is_set()
assert encoder.encode_dispatch_lock.locked()
assert encoder.released == []
finish_encode.set()
await asyncio.wait_for(encoder.release_event.wait(), timeout=1)
await asyncio.wait_for(encoder.encode_dispatch_lock.acquire(), timeout=1)
encoder.encode_dispatch_lock.release()
assert len(encoder.released) == 1
asyncio.run(run_test())
def test_cancelled_health_request_does_not_cancel_dispatched_encode(monkeypatch):
async def run_test():
encoder = _FakeEncoder()
_install_tp_encoder(monkeypatch, encoder)
encode_started = asyncio.Event()
finish_encode = asyncio.Event()
async def encode(**kwargs):
encoder.encode_calls.append(kwargs)
encode_started.set()
await finish_encode.wait()
return 1, 1, 1, None, None
encoder.encode = encode
task = asyncio.create_task(http_server.health_generate())
await asyncio.wait_for(encode_started.wait(), timeout=1)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert encoder.encode_dispatch_lock.locked()
assert encoder.released == []
finish_encode.set()
await asyncio.wait_for(encoder.release_event.wait(), timeout=1)
await asyncio.wait_for(encoder.encode_dispatch_lock.acquire(), timeout=1)
encoder.encode_dispatch_lock.release()
assert len(encoder.released) == 1
asyncio.run(run_test())
def test_health_cleanup_failure_releases_dispatch_lock(monkeypatch):
async def run_test():
encoder = _FakeEncoder()
_install_tp_encoder(monkeypatch, encoder)
async def release_request(req_id):
encoder.released.append(req_id)
raise RuntimeError("cleanup failed")
encoder.release_request = release_request
response = await http_server.health_generate()
assert response.status_code == 503
assert not encoder.encode_dispatch_lock.locked()
assert len(encoder.released) == 1
asyncio.run(run_test())
if __name__ == "__main__": if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"])) sys.exit(pytest.main([__file__, "-v"]))
@@ -1,5 +1,7 @@
import asyncio import asyncio
import sys import sys
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest import pytest
@@ -7,7 +9,9 @@ from sglang.srt.disaggregation.encoder.runtime import (
EncoderScheduler, EncoderScheduler,
PendingRequest, PendingRequest,
_resolve_encoder_batch_policy, _resolve_encoder_batch_policy,
validate_encode_request,
) )
from sglang.srt.managers.schedule_batch import Modality
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu") register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -108,6 +112,113 @@ def test_scheduler_coalesces_concurrent_submissions():
asyncio.run(run_test()) asyncio.run(run_test())
def test_scheduler_isolates_bad_request_from_failed_fused_batch():
class FakeEncoder:
def __init__(self):
self.encode_dispatch_lock = asyncio.Lock()
self.batches = []
async def batch_encode(self, requests, _modality):
req_ids = [request["req_id"] for request in requests]
self.batches.append(req_ids)
if len(requests) > 1 or req_ids == ["bad"]:
return [(0, 0, 0, "bad image", 400) for _ in requests]
return [(1, 2, 3, None, None)]
async def run_test():
encoder = FakeEncoder()
scheduler = EncoderScheduler(
encoder=encoder,
send_sockets=[],
max_batch_size=8,
coalesce_same_turn=True,
)
collector = SimpleNamespace(observe_queue_wait=Mock())
with patch(
"sglang.srt.disaggregation.encoder.runtime.server_module.encoder_metrics_collector",
collector,
):
scheduler.start()
try:
requests = [
{
"req_id": req_id,
"modality": "image",
"mm_items": [object()],
"num_parts": 1,
"part_idx": 0,
}
for req_id in ("bad", "good")
]
results = await asyncio.gather(
*(scheduler.submit(request) for request in requests)
)
finally:
await scheduler.stop()
assert encoder.batches == [["bad", "good"], ["bad"], ["good"]]
assert results == [(0, 0, 0, "bad image", 400), (1, 2, 3, None, None)]
assert collector.observe_queue_wait.call_count == len(requests)
asyncio.run(run_test())
@pytest.mark.parametrize(
("update", "expected"),
[
({"req_id": ""}, "missing or invalid req_id"),
({"modality": "text"}, "unsupported modality"),
({"mm_items": []}, "missing or empty mm_items"),
({"num_parts": 0}, "num_parts must be a positive integer"),
({"part_idx": 1}, "part_idx must be in [0, 1)"),
],
)
def test_validate_encode_request_rejects_invalid_fields(update, expected):
request = {
"req_id": "request",
"modality": "image",
"mm_items": [object()],
"num_parts": 1,
"part_idx": 0,
}
request.update(update)
assert expected in validate_encode_request(request)
def test_video_request_is_validated_before_tp_broadcast():
class FakeSocket:
pass
class FakeEncoder:
async def encode(self, **_kwargs):
raise AssertionError("invalid request must not reach the encoder")
async def run_test():
scheduler = EncoderScheduler(
encoder=FakeEncoder(),
send_sockets=[FakeSocket()],
max_batch_size=1,
)
pending = PendingRequest(
{
"req_id": "bad-video",
"modality": "video",
"mm_items": [object()],
"num_parts": 1,
"part_idx": 1,
},
asyncio.get_running_loop(),
)
await scheduler._dispatch_per_request([pending], Modality.VIDEO)
with pytest.raises(Exception, match="part_idx must be in"):
pending.future.result()
asyncio.run(run_test())
@pytest.mark.parametrize( @pytest.mark.parametrize(
("model_type", "configured", "explicit", "expected"), ("model_type", "configured", "explicit", "expected"),
[ [