[EPD] feat: encoder DP mode with per-rank subprocess workers (#26576)

This commit is contained in:
Zhonghua Deng
2026-06-04 12:37:41 +08:00
committed by GitHub
parent 71c759ebb7
commit e541bc3881
2 changed files with 915 additions and 2 deletions
@@ -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:
+1
View File
@@ -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)