From e541bc3881b09258a64726f417fe69285391f22b Mon Sep 17 00:00:00 2001 From: Zhonghua Deng Date: Thu, 4 Jun 2026 12:37:41 +0800 Subject: [PATCH] [EPD] feat: encoder DP mode with per-rank subprocess workers (#26576) --- .../srt/disaggregation/encode_server.py | 916 +++++++++++++++++- python/sglang/srt/environ.py | 1 + 2 files changed, 915 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 265f5c6fb..02b888165 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -1,6 +1,7 @@ import asyncio import concurrent.futures import contextlib +import copy import ctypes import functools import logging @@ -57,17 +58,18 @@ from sglang.srt.utils import ( load_video, random_uuid, ) -from sglang.srt.utils.common import configure_logger +from sglang.srt.utils.common import configure_logger, maybe_reindex_device_id from sglang.srt.utils.network import ( NetworkAddress, config_socket, + get_free_port, get_local_ip_auto, get_zmq_socket, ) logger = logging.getLogger(__name__) -HEALTH_CHECK_TIMEOUT = 10 +HEALTH_CHECK_TIMEOUT = 30 # Minimal 32x32 black PNG for health check dummy encode MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg==" @@ -2371,10 +2373,712 @@ encoder: Optional[MMEncoder] = None send_sockets: List[zmq.Socket] = [] encoder_scheduler: Optional[EncoderScheduler] = None +# DP mode (--dp-size > 1): each rank runs as a subprocess with its own +# MMEncoder on its own GPU; the main process only routes via ZMQ so the +# asyncio event loop is never blocked by GPU work. +dp_dispatcher: Optional["DPDispatcher"] = None + + +async def _push_embedding_to_prefill(enc: MMEncoder, request: dict) -> None: + # No-op for mooncake (its /send is separate). embedding_port=None is + # rejected upfront, so ports is always a concrete list here. + req_id = request["req_id"] + backend = enc.server_args.encoder_transfer_backend + + if backend == "zmq_to_tokenizer": + await enc.send( + req_id=req_id, + prefill_host=request["prefill_host"], + embedding_port=request["embedding_port"], + ) + enc.embedding_to_send.pop(req_id, None) + return + + if backend == "zmq_to_scheduler": + ports = request["embedding_port"] + assert isinstance(ports, list) + await asyncio.gather( + *( + enc.send( + req_id=req_id, + prefill_host=request["prefill_host"], + embedding_port=p, + ) + for p in ports + ) + ) + enc.embedding_to_send.pop(req_id, None) + + +async def _dp_worker_encode_and_send( + enc: MMEncoder, + sched: Optional[EncoderScheduler], + request: dict, +) -> Optional[dict]: + # Mooncake returns metadata for main to forward; zmq inlines the send. + # Soft errors raise MMError so the dispatcher route maps them to HTTP. + req_id = request["req_id"] + request["enter_time"] = time.time() + modality = Modality.from_str(request["modality"]) + backend = enc.server_args.encoder_transfer_backend + + # URL state lives in main process module globals; workers don't see it. + if backend == "zmq_to_scheduler" and request.get("embedding_port") is None: + raise MMError( + "Encoder DP mode does not support zmq_to_scheduler with " + "embedding_port=None (URL state isn't synchronised to workers). " + "Provide an explicit embedding_port list, switch to mooncake / " + "zmq_to_tokenizer, or run without --dp-size.", + code=HTTPStatus.BAD_REQUEST, + ) + + encode_coro = ( + sched.submit(request) + if sched is not None and modality in _BATCHABLE_MODALITIES + else enc.encode_request(request, modality) + ) + nbytes, embedding_len, embedding_dim, error_msg, error_code = await encode_coro + + if error_msg: + # zmq backends still forward an error EmbeddingData to P so it + # doesn't block; send failures here are swallowed. + try: + await _push_embedding_to_prefill(enc, request) + except Exception as e: + logger.error( + f"DP error-send failed for req_id={req_id}: {e}", exc_info=True + ) + # Free the error EmbeddingData stored during encode, or it leaks in + # embedding_to_send and pins /health into "busy" (a non-empty + # embedding_to_send reads as busy, skipping the probe). Neither path + # guarantees cleanup on its own: mooncake's _push_embedding_to_prefill + # is a no-op, and a swallowed zmq send failure above skips its own pop. + # zmq lacks the inflight attrs so _cleanup_inflight_encode_state would + # early-return on it — pop directly. Mirrors the non-DP error path. + if backend == "mooncake": + await enc._cleanup_inflight_encode_state(req_id) + else: + enc.embedding_to_send.pop(req_id, None) + raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) + + if backend == "mooncake": + request.pop("mm_items", None) + request.update( + embedding_size=nbytes, + embedding_len=embedding_len, + embedding_dim=embedding_dim, + ) + # Free the held embedding if the follow-up /send never arrives (same + # send_timeout cleanup the non-DP path uses). + enc._schedule_inflight_encode_cleanup(req_id) + return request + + await _push_embedding_to_prefill(enc, request) + return None + + +async def _dp_worker_health_encode(enc: MMEncoder) -> None: + """Functional health probe run on a DP worker. + + Process-liveness (proc.sentinel) can't see a worker that's alive but + wedged — hung GPU, NCCL deadlock, stalled ZMQ, or a blocked event loop. + When idle, run a tiny dummy encode to exercise the VIT forward and surface + those stalls. No prefill destination: the embedding is discarded, mirroring + the non-DP /health path. Raises on encode failure so the worker envelope + carries ``_error`` back to the dispatcher. + """ + # Busy worker: in-flight traffic already proves liveness, so skip the probe + # and report healthy — same `embedding_to_send` signal the non-DP /health + # path uses. A wedged-but-busy worker never reaches here (it can't service + # the recv), so the dispatcher's broadcast still times out → 503. + if enc.embedding_to_send: + return None + + if enc.image_processor is not None: + mm_items = [f"data:image/png;base64,{MINIMUM_PNG_PICTURE_BASE64}"] + modality = Modality.IMAGE + elif enc.audio_processor is not None: + mm_items = [f"data:audio/wav;base64,{MINIMUM_WAV_SILENCE_BASE64}"] + modality = Modality.AUDIO + else: + # No processor → can't functionally probe; liveness alone is healthy. + return None + + req_id = f"{HEALTH_CHECK_RID_PREFIX}_{time.time()}" + try: + _, _, _, error_msg, error_code = await enc.encode( + mm_items=mm_items, + modality=modality, + req_id=req_id, + num_parts=1, + part_idx=0, + ) + finally: + # Never leave the dummy embedding sitting in the send map. + enc.embedding_to_send.pop(req_id, None) + + if error_msg: + raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) + + +class DPDispatcher: + """Routes encode requests across DP ranks by least-pending count.""" + + def __init__( + self, + dp_size: int, + dispatch_sockets: List, + result_socket, + worker_processes: List[mp.Process], + ): + self.dp_size = dp_size + self.dispatch_sockets = dispatch_sockets + self.result_socket = result_socket + self.worker_processes = worker_processes + # Key = req_id for encode/broadcast, req_id + "_send" for mooncake /send. + self.pending_futures: List[Dict[str, asyncio.Future]] = [ + {} for _ in range(dp_size) + ] + self.req_id_to_rank: Dict[str, int] = {} + self._rr_counter = 0 + self._broadcast_counter = 0 + self._dead_ranks: Set[int] = set() + # req_id -> monotonic ts a mooncake mapping has waited for its /send. + self._pending_send_at: Dict[str, float] = {} + # Set when _result_listener gives up; makes alive_ranks report empty. + self._listener_failed = False + + @property + def pending_counts(self) -> List[int]: + return [len(d) for d in self.pending_futures] + + @property + def alive_ranks(self) -> List[int]: + # Empty if the result listener died; else ranks not marked dead. + if self._listener_failed: + return [] + return [r for r in range(self.dp_size) if r not in self._dead_ranks] + + @property + def all_ranks_alive(self) -> bool: + # Strict health (only /health uses this); routing still degrades. + return len(self.alive_ranks) == self.dp_size + + def start(self) -> None: + logger.info(f"DP dispatcher started: {self.dp_size} ranks (all remote)") + asyncio.create_task(self._result_listener()) + asyncio.create_task(self._worker_watchdog()) + asyncio.create_task(self._cleanup_stale_mappings()) + + def _drop_pending_and_mapping(self, rank: int, req_id: str) -> None: + # dispatch / broadcast failure: no follow-up /send expected. + self.pending_futures[rank].pop(req_id, None) + self.req_id_to_rank.pop(req_id, None) + + def _fail_pending_for_rank(self, rank: int, reason: str, error_type: str) -> None: + # Resolve a rank's outstanding futures with 503 so awaiters don't hang. + pending = self.pending_futures[rank] + for key, future in list(pending.items()): + if not future.done(): + future.set_result( + { + "req_id": key.removesuffix("_send"), + "_dp_type": "send" if key.endswith("_send") else "encode", + "content": None, + "_error": reason, + "_error_type": error_type, + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + ) + pending.pop(key, None) + + def _fail_all_pending(self, reason: str, error_type: str) -> None: + for rank in range(self.dp_size): + self._fail_pending_for_rank(rank, reason, error_type) + self.req_id_to_rank.clear() + self._pending_send_at.clear() + + @staticmethod + def _timeout_envelope(req_id: str, dp_type: str, reason: str) -> dict: + return { + "req_id": req_id, + "_dp_type": dp_type, + "content": None, + "_error": reason, + "_error_type": "TimeoutError", + "_error_code": int(HTTPStatus.GATEWAY_TIMEOUT), + } + + async def dispatch(self, request: dict) -> dict: + counts = self.pending_counts + # Skip ranks whose worker process has died. + alive_ranks = self.alive_ranks + if not alive_ranks: + raise MMError( + "All encoder DP workers are dead.", + code=HTTPStatus.SERVICE_UNAVAILABLE, + ) + min_p = min(counts[r] for r in alive_ranks) + candidates = [r for r in alive_ranks if counts[r] == min_p] + rank = candidates[self._rr_counter % len(candidates)] + self._rr_counter += 1 + req_id = request["req_id"] + self.req_id_to_rank[req_id] = rank + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][req_id] = future + logger.info( + f"MM-Encoder DP dispatch: req_id={req_id}, " + f"modality={request.get('modality', 'image')}, " + f"dp_rank={rank}, pending={self.pending_counts}" + ) + + try: + await self.dispatch_sockets[rank].send_pyobj(request) + # An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip + # the watchdog, so bound the wait explicitly. + return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT) + except asyncio.TimeoutError: + self._drop_pending_and_mapping(rank, req_id) + return self._timeout_envelope( + req_id, + "encode", + f"Encoder DP rank={rank} timed out after {ENCODER_REQ_TIMEOUT}s", + ) + except BaseException: + self._drop_pending_and_mapping(rank, req_id) + raise + + async def dispatch_send(self, request: dict) -> dict: + req_id = request["req_id"] + # /send arrived → stop tracking it for stale-mapping GC. + self._pending_send_at.pop(req_id, None) + if self._listener_failed: + return { + "req_id": req_id, + "_error": "encoder DP result listener stopped; cannot route /send", + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + rank = self.req_id_to_rank.get(req_id) + if rank is None: + logger.warning( + f"MM-Encoder dispatch_send: unknown req_id={req_id}, " + f"cannot route to worker" + ) + return {"req_id": req_id, "_error": f"Unknown req_id: {req_id}"} + if rank in self._dead_ranks: + # Worker died between encode and /send; embedding is gone. + self.req_id_to_rank.pop(req_id, None) + return { + "req_id": req_id, + "_error": f"DP worker rank={rank} died before /send for req_id={req_id}", + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + key = req_id + "_send" + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][key] = future + request["_dp_type"] = "send" + logger.info( + f"MM-Encoder DP dispatch_send: req_id={req_id}, " + f"dp_rank={rank}, pending={self.pending_counts}" + ) + try: + await self.dispatch_sockets[rank].send_pyobj(request) + return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT) + except asyncio.TimeoutError: + self.pending_futures[rank].pop(key, None) + self.req_id_to_rank.pop(req_id, None) + return self._timeout_envelope( + req_id, + "send", + f"Encoder DP rank={rank} /send timed out after {ENCODER_REQ_TIMEOUT}s", + ) + except BaseException: + self.pending_futures[rank].pop(key, None) + self.req_id_to_rank.pop(req_id, None) + raise + + async def broadcast( + self, request: dict, timeout: Optional[float] = None + ) -> List[dict]: + # Skip dead ranks: a PUSH to a gone worker would just buffer and then + # surface as a spurious per-rank timeout. All dead → 503 (same as + # dispatch), which the profile endpoints turn into an HTTP error. + eff_timeout = timeout if timeout is not None else ENCODER_REQ_TIMEOUT + alive_ranks = self.alive_ranks + if not alive_ranks: + raise MMError( + "All encoder DP workers are dead.", + code=HTTPStatus.SERVICE_UNAVAILABLE, + ) + batch_id = self._broadcast_counter + self._broadcast_counter += 1 + rank_keys: List[Tuple[int, str]] = [] + futures: List[asyncio.Future] = [] + dp_type = request.get("_dp_type", "unknown") + try: + for rank in alive_ranks: + req_id = f"_broadcast_{batch_id}_{rank}" + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][req_id] = future + self.req_id_to_rank[req_id] = rank + rank_keys.append((rank, req_id)) + request_copy = {**request, "req_id": req_id} + await self.dispatch_sockets[rank].send_pyobj(request_copy) + futures.append(future) + # Concurrent wait → total bounded by eff_timeout, not + # dp_size × eff_timeout. + outcomes = await asyncio.gather( + *(asyncio.wait_for(fut, timeout=eff_timeout) for fut in futures), + return_exceptions=True, + ) + results: List[dict] = [] + for (rank, req_id), outcome in zip(rank_keys, outcomes): + if isinstance(outcome, asyncio.TimeoutError): + self._drop_pending_and_mapping(rank, req_id) + results.append( + self._timeout_envelope( + req_id, + dp_type, + f"Encoder DP rank={rank} broadcast timed out " + f"after {eff_timeout}s", + ) + ) + elif isinstance(outcome, BaseException): + self._drop_pending_and_mapping(rank, req_id) + raise outcome + else: + results.append(outcome) + return results + except BaseException: + for rank, req_id in rank_keys: + self._drop_pending_and_mapping(rank, req_id) + raise + + async def _worker_watchdog(self) -> None: + # proc.sentinel becomes readable on process exit; fail this rank's + # pending futures so awaiters don't hang on a dead worker. + loop = asyncio.get_running_loop() + watch: Dict[int, asyncio.Future] = {} + for rank, proc in enumerate(self.worker_processes): + fut: asyncio.Future = loop.create_future() + + # add_reader is level-triggered, so remove_reader inside the + # callback to avoid spinning every loop iteration. + def _on_exit(r=rank, f=fut, p=proc, lp=loop): + try: + lp.remove_reader(p.sentinel) + except (ValueError, OSError): + pass + if not f.done(): + f.set_result(r) + + try: + loop.add_reader(proc.sentinel, _on_exit) + except (ValueError, OSError): + continue + watch[rank] = fut + + while watch: + done, _ = await asyncio.wait( + watch.values(), return_when=asyncio.FIRST_COMPLETED + ) + for fut in done: + rank = fut.result() + proc = self.worker_processes[rank] + logger.error( + f"DP worker rank={rank} (pid={proc.pid}) exited " + f"with code={proc.exitcode}; failing pending requests" + ) + self._dead_ranks.add(rank) + reason = f"DP worker rank={rank} died (exitcode={proc.exitcode})" + self._fail_pending_for_rank(rank, reason, "WorkerDied") + self.req_id_to_rank = { + r: rk for r, rk in self.req_id_to_rank.items() if rk != rank + } + watch.pop(rank, None) + + async def _result_listener(self) -> None: + # Bounded back-off + give-up so a torn-down context exits in ~3s + # rather than spinning forever on recv errors. + consecutive_errors = 0 + while True: + try: + msg = await self.result_socket.recv_pyobj() + consecutive_errors = 0 + except asyncio.CancelledError: + raise + except Exception: + consecutive_errors += 1 + logger.error("_result_listener recv error", exc_info=True) + if consecutive_errors >= 30: + logger.error( + "_result_listener giving up after 30 consecutive errors" + ) + self._listener_failed = True + self._fail_all_pending( + "encoder DP result listener stopped after repeated " + "recv errors", + "ResultListenerStopped", + ) + return + await asyncio.sleep(min(0.1 * consecutive_errors, 1.0)) + continue + req_id = msg.get("req_id", "") + dp_type = msg.get("_dp_type", "encode") + key = (req_id + "_send") if dp_type == "send" else req_id + rank = self.req_id_to_rank.get(req_id) + if rank is None or key not in self.pending_futures[rank]: + logger.warning( + f"_result_listener: no pending future for " + f"req_id={req_id}, dp_type={dp_type}, dropping" + ) + continue + future = self.pending_futures[rank].pop(key) + # Only mooncake encode (content=request dict) needs the mapping + # kept for the follow-up /send. + keep_mapping = dp_type == "encode" and msg.get("content") is not None + if keep_mapping: + self._pending_send_at[req_id] = time.monotonic() + else: + self.req_id_to_rank.pop(req_id, None) + try: + future.set_result(msg) + + except asyncio.InvalidStateError: + logger.warning( + f"_result_listener: future already done for " + f"req_id={req_id}, dp_type={dp_type}" + ) + + async def _cleanup_stale_mappings(self) -> None: + # Evict req_id->rank mappings whose /send never came. The worker frees + # its own embedding via the send_timeout cleanup scheduled at encode, + # so both sides key off the same timeout. + ttl = envs.SGLANG_ENCODER_SEND_TIMEOUT.get() + interval = max(ttl / 4, 30) + while True: + await asyncio.sleep(interval) + now = time.monotonic() + stale = [rid for rid, ts in self._pending_send_at.items() if now - ts > ttl] + for rid in stale: + self._pending_send_at.pop(rid, None) + self.req_id_to_rank.pop(rid, None) + if stale: + logger.warning( + f"Evicted {len(stale)} stale encoder DP /send mapping(s) " + f"with no /send within {ttl}s" + ) + + +async def _dp_worker_handle_profile( + enc: MMEncoder, dp_rank: int, dp_type: str, request: dict +) -> dict: + prefix = f"dp_rank={dp_rank}: " + if dp_type == "start_profile": + obj = request.get("profile_req") + # `is None` (not `if not obj`) so empty dict still raises. + req = ( + ProfileReq(**obj) + if obj is not None + else ProfileReq(ProfileReqType.START_PROFILE) + ) + if enc.profiler is None: + enc.profiler = EncoderProfiler(dp_rank) + ok, msg = enc.profiler.start(req) + detail = ( + f"started profiling, output_dir={enc.profiler.output_dir}" if ok else msg + ) + else: # stop_profile + if enc.profiler is None: + return {"ok": False, "msg": prefix + "profiling not initialized"} + ok, msg = enc.profiler.stop() + detail = "stopped profiling" if ok else msg + return {"ok": ok, "msg": prefix + detail} + + +async def _dp_worker_handle_request( + enc: MMEncoder, + sched: EncoderScheduler, + send_sock, + send_lock: asyncio.Lock, + dp_rank: int, + request: dict, + dp_type: str, +) -> None: + t0 = time.time() + try: + if dp_type in ("start_profile", "stop_profile"): + content = await _dp_worker_handle_profile(enc, dp_rank, dp_type, request) + elif dp_type == "health_encode": + content = await _dp_worker_health_encode(enc) + elif dp_type == "send": + req_id = request["req_id"] + 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"], + ) + # cancels the scheduled cleanup + frees embedding/forward state + await enc._cleanup_inflight_encode_state(req_id) + content = None + else: + content = await _dp_worker_encode_and_send(enc, sched, request) + + logger.info( + f"MM-Encoder [dp_rank={dp_rank}] {dp_type} done: " + f"req_id={request.get('req_id', '?')}, " + f"modality={request.get('modality', 'image')}, " + f"cost={(time.time() - t0) * 1000:.1f}ms" + ) + envelope = { + "req_id": request.get("req_id", ""), + "_dp_type": dp_type, + "content": content, + } + except Exception as e: + logger.error( + f"DP worker {dp_rank} error on {dp_type} " + f"req_id={request.get('req_id', '?')}: {e}", + exc_info=True, + ) + err_code = int(getattr(e, "code", None) or HTTPStatus.INTERNAL_SERVER_ERROR) + envelope = { + "req_id": request.get("req_id", ""), + "_dp_type": dp_type, + "content": None, + "_error": str(e), + "_error_type": type(e).__name__, + "_error_code": err_code, + } + + # pyzmq async send_pyobj isn't safe for concurrent senders. + try: + async with send_lock: + await send_sock.send_pyobj(envelope) + except Exception: + logger.error( + f"DP worker {dp_rank} failed to send envelope for " + f"req_id={request.get('req_id', '?')}", + exc_info=True, + ) + + +async def run_dp_worker( + server_args: ServerArgs, + dp_rank: int, + gpu_id: int, + dispatch_path: str, + result_path: str, +): + logger.info( + f"DP worker {dp_rank} starting on gpu_id={gpu_id} " + f"(CUDA_VISIBLE_DEVICES={os.environ.get('CUDA_VISIBLE_DEVICES', 'unset')})" + ) + + # gpu_id is the device chosen by maybe_reindex_device_id in the parent: + # 0 when CVD is pinned to one GPU, else the absolute id. rank=0, so + # MMEncoder runs set_device(base_gpu_id). + args = copy.deepcopy(server_args) + args.base_gpu_id = gpu_id + args.tp_size = 1 + enc = MMEncoder(args, dist_init_method=f"tcp://127.0.0.1:{get_free_port()}", rank=0) + sched = EncoderScheduler( + encoder=enc, send_sockets=[], max_batch_size=ENCODER_MAX_BATCH_SIZE + ) + + ctx = zmq.asyncio.Context(2) + recv_sock = get_zmq_socket(ctx, zmq.PULL, dispatch_path, False) + send_sock = get_zmq_socket(ctx, zmq.PUSH, result_path, False) + send_lock = asyncio.Lock() + inflight: Set[asyncio.Task] = set() + # Acquire-before-recv → back-pressure propagates to the dispatcher + # PUSH buffer. Must be ≥ ENCODER_MAX_BATCH_SIZE or batching degrades. + max_inflight = envs.SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT.get() + if max_inflight < ENCODER_MAX_BATCH_SIZE: + logger.warning( + f"SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT={max_inflight} is below " + f"ENCODER_MAX_BATCH_SIZE={ENCODER_MAX_BATCH_SIZE}; the encoder " + f"will never assemble a full batch." + ) + inflight_sem = asyncio.Semaphore(max_inflight) + sched.start() + logger.info(f"DP worker {dp_rank} ready") + + # Task-per-request so EncoderScheduler.pending_queue accumulates and + # actual cross-request batching can happen. + try: + while True: + await inflight_sem.acquire() + # Released by _run on success or the outer finally if not spawned. + spawned = False + try: + try: + request = await recv_sock.recv_pyobj() + except asyncio.CancelledError: + raise + except Exception: + logger.error(f"DP worker {dp_rank} recv error", exc_info=True) + continue + if not isinstance(request, dict): + logger.error( + f"DP worker {dp_rank} received non-dict request " + f"({type(request).__name__}); dropping" + ) + continue + dp_type = request.pop("_dp_type", "encode") + + async def _run(req=request, t=dp_type): + try: + await _dp_worker_handle_request( + enc, sched, send_sock, send_lock, dp_rank, req, t + ) + finally: + inflight_sem.release() + + task = asyncio.create_task(_run()) + # Ownership transferred to _run; mark before any op that could + # raise (theoretical: set.add / add_done_callback) and cause a + # double-release. + spawned = True + inflight.add(task) + task.add_done_callback(inflight.discard) + finally: + if not spawned: + inflight_sem.release() + finally: + # Close zmq on exception/cancellation (normal stop is parent SIGKILL). + for task in inflight: + task.cancel() + ctx.destroy(linger=0) + + +def launch_dp_worker( + server_args: ServerArgs, + dp_rank: int, + gpu_id: int, + dispatch_path: str, + result_path: str, +): + try: + configure_logger(server_args, prefix=f" encode_dp_worker[{dp_rank}]") + asyncio.run( + run_dp_worker(server_args, dp_rank, gpu_id, dispatch_path, result_path) + ) + except KeyboardInterrupt: + logger.info(f"DP worker {dp_rank} exiting") + except Exception: + traceback.print_exc() + @contextlib.asynccontextmanager async def _lifespan(app: FastAPI): global encoder_scheduler + if dp_dispatcher is not None: + dp_dispatcher.start() + yield + return if encoder is not None: encoder_scheduler = EncoderScheduler( encoder, send_sockets, max_batch_size=ENCODER_MAX_BATCH_SIZE @@ -2425,6 +3129,10 @@ def launch_encoder(server_args, schedule_path, dist_init_method, rank): def launch_server(server_args: ServerArgs): configure_logger(server_args, prefix=" encode_server") + if server_args.dp_size > 1: + _launch_server_dp(server_args) + return + global encoder ctx = mp.get_context("spawn") zmq_ctx = zmq.Context(10) @@ -2451,6 +3159,105 @@ def launch_server(server_args: ServerArgs): uvicorn.run(app, host=server_args.host, port=server_args.port) +def _launch_server_dp(server_args: ServerArgs): + global dp_dispatcher + + if server_args.dp_size <= 1 or server_args.tp_size != 1: + raise ValueError( + "Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " + f"dp_size={server_args.dp_size}, tp_size={server_args.tp_size}." + ) + dp_size = server_args.dp_size + logger.info(f"Launching encoder in DP mode: dp_size={dp_size}") + + ctx = mp.get_context("spawn") + ipc_prefix = random_uuid() + async_zmq_ctx = zmq.asyncio.Context(dp_size + 1) + + result_path = f"ipc:///tmp/{ipc_prefix}_dp_result" + result_socket = get_zmq_socket(async_zmq_ctx, zmq.PULL, result_path, True) + + dispatch_sockets: List[zmq.asyncio.Socket] = [ + get_zmq_socket( + async_zmq_ctx, zmq.PUSH, f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{r}", True + ) + for r in range(dp_size) + ] + + # Register atexit BEFORE spawn loop so partial spawns get reaped on + # exception (atexit holds the list ref and reads it at exit time). + import atexit + + worker_processes: List[mp.Process] = [] + + def _kill_workers(): + for p in worker_processes: + if p.is_alive(): + p.kill() + for p in worker_processes: + p.join(timeout=5) + + atexit.register(_kill_workers) + + for dp_rank in range(dp_size): + gpu_id = server_args.base_gpu_id + dp_rank + # Pin the device parent-side around spawn (same convention as the + # scheduler launcher and DP controller) so the child inherits + # CUDA_VISIBLE_DEVICES from its first instruction, before any import + # can enumerate CUDA. No-op unless SGLANG_ONE_VISIBLE_DEVICE_PER_PROCESS + # is set, in which case gpu_id is reindexed to 0 and CVD is pinned. + with maybe_reindex_device_id(gpu_id) as gpu_id: + proc = ctx.Process( + target=launch_dp_worker, + args=( + server_args, + dp_rank, + gpu_id, + f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{dp_rank}", + result_path, + ), + daemon=False, + ) + proc.start() + worker_processes.append(proc) + + dp_dispatcher = DPDispatcher( + dp_size, + dispatch_sockets, + result_socket, + worker_processes, + ) + + uvicorn.run(app, host=server_args.host, port=server_args.port) + + +def _summarise_dp_broadcast(results: List[dict]) -> Response: + # Treat missing/None content as failure so a stuck rank doesn't hide + # behind the others' "ok". Status = the most severe per-rank error code + # (5xx beats 4xx) rather than a blanket 400, so a worker's 500/503/504 + # isn't misreported as a client error. + msgs: List[str] = [] + error_codes: List[int] = [] + for r in results: + content = r.get("content") + if isinstance(content, dict): + msgs.append(content.get("msg", "")) + if not content.get("ok"): + # Worker ran but reported a logical failure; no transport code, + # so treat as a bad request (matches the non-DP profile path). + error_codes.append(int(r.get("_error_code") or HTTPStatus.BAD_REQUEST)) + else: + msgs.append(r.get("_error", "unknown error")) + error_codes.append( + int(r.get("_error_code") or HTTPStatus.INTERNAL_SERVER_ERROR) + ) + status_code = 200 if not error_codes else max(error_codes) + return Response( + content="\n".join(msgs) + "\n", + status_code=status_code, + ) + + async def get_condition(rid): async with cond_dict_lock: if rid not in rid_to_cond: @@ -2462,6 +3269,40 @@ async def get_condition(rid): async def handle_encode_request(request: dict): req_id = request["req_id"] start_time = time.monotonic() + if dp_dispatcher is not None: + try: + result = await dp_dispatcher.dispatch(request) + except MMError as e: + # Surface MMError.code (503 when all workers dead) instead of + # FastAPI's default 500. + logger.error(f"DP dispatch refused req_id={req_id}: {e}") + return ORJSONResponse( + status_code=int(e.code), + content={"status": "error", "message": str(e), "req_id": req_id}, + ) + if result.get("_error"): + error_type = result.get("_error_type", "") + # `or` (not `dict.get(key, default)`) so explicit None falls back too. + status_code = result.get("_error_code") or ( + HTTPStatus.BAD_REQUEST + if error_type == "ValueError" + else HTTPStatus.INTERNAL_SERVER_ERROR + ) + logger.error(f"DP worker error for req_id={req_id}: {result['_error']}") + return ORJSONResponse( + status_code=status_code, + content={ + "status": "error", + "message": result["_error"], + "req_id": req_id, + }, + ) + elapsed = time.monotonic() - start_time + logger.info( + f"[{req_id}] /encode completed in {elapsed:.3f}s, " + f"modality={request.get('modality', 'image')}" + ) + return ORJSONResponse(content=result.get("content")) try: # when multiple decoder TP ranks POST /encode # with the same req_id, only the first triggers the VIT forward; @@ -2631,6 +3472,29 @@ async def handle_encode_request(request: dict): @app.post("/send") async def handle_send_request(request: dict): # mooncake backend + if dp_dispatcher is not None: + try: + result = await dp_dispatcher.dispatch_send(request) + except MMError as e: + req_id = request.get("req_id", "?") + logger.error(f"DP dispatch_send refused req_id={req_id}: {e}") + return Response( + content=f"Encoder DP worker send error: {e}", + status_code=int(e.code), + ) + if result.get("_error"): + req_id = request.get("req_id", "?") + status_code = result.get("_error_code") or int( + HTTPStatus.INTERNAL_SERVER_ERROR + ) + logger.error( + f"DP worker send error for req_id={req_id}: {result['_error']}" + ) + return Response( + content=f"Encoder DP worker send error: {result['_error']}", + status_code=status_code, + ) + return ORJSONResponse(content=result.get("content")) await encoder.send( req_id=request["req_id"], prefill_host=request["prefill_host"], @@ -2668,6 +3532,24 @@ async def health_generate(): Performs a dummy encode to verify the encoder is functional. Returns 200 if the encoder is healthy, 503 otherwise. """ + if dp_dispatcher is not None: + # Strict: any dead (exited) rank fails health → orchestrator restarts. + if not dp_dispatcher.all_ranks_alive: + return Response(status_code=503) + # Process-liveness (proc.sentinel) can't see a worker that's alive but + # wedged (hung GPU / NCCL deadlock / stalled ZMQ). Probe every rank with + # a tiny dummy encode; each worker runs it only when idle and otherwise + # reports healthy at once, keeping the probe off the GPU under load. + try: + results = await dp_dispatcher.broadcast( + {"_dp_type": "health_encode"}, + timeout=HEALTH_CHECK_TIMEOUT, + ) + except MMError: + return Response(status_code=503) + if any(r.get("_error") for r in results): + return Response(status_code=503) + return Response(status_code=200) if encoder is None: return Response(status_code=503) @@ -2734,6 +3616,30 @@ async def health_generate(): @app.api_route("/start_profile", methods=["GET", "POST"]) async def start_profile_async(obj: Optional[ProfileReqInput] = None): + if dp_dispatcher is not None: + profile_req = None + if obj is not None: + profile_req = { + "type": ProfileReqType.START_PROFILE, + "output_dir": obj.output_dir, + "start_step": obj.start_step, + "num_steps": obj.num_steps, + "activities": obj.activities, + "with_stack": obj.with_stack, + "record_shapes": obj.record_shapes, + "profile_by_stage": obj.profile_by_stage, + "profile_id": str(time.time()), + "merge_profiles": obj.merge_profiles, + "profile_prefix": obj.profile_prefix, + "profile_stages": obj.profile_stages, + } + try: + results = await dp_dispatcher.broadcast( + {"_dp_type": "start_profile", "profile_req": profile_req} + ) + except MMError as e: + return Response(content=f"{e}\n", status_code=int(e.code)) + return _summarise_dp_broadcast(results) if encoder is None: return Response(content="encoder not ready\n", status_code=503) req = None @@ -2772,6 +3678,12 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None): @app.api_route("/stop_profile", methods=["GET", "POST"]) async def stop_profile_async(): + if dp_dispatcher is not None: + try: + results = await dp_dispatcher.broadcast({"_dp_type": "stop_profile"}) + except MMError as e: + return Response(content=f"{e}\n", status_code=int(e.code)) + return _summarise_dp_broadcast(results) if encoder is None: return Response(content="encoder not ready\n", status_code=503) if encoder.profiler is None: diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 82ec4ab7d..de4264c95 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -779,6 +779,7 @@ class Envs: # Persistent receiver-side GPU embedding pool size for mooncake EPD transport. # 0 disables (per-request register/deregister). 4096 = 4GB default per TP SGLANG_EMBEDDING_POOL_SIZE_MB = EnvInt(4096) + SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT = EnvInt(64) # Elastic EP Backup Port SGLANG_BACKUP_PORT_BASE = EnvInt(10000)