[EPD] feat: encoder DP mode with per-rank subprocess workers (#26576)
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
import contextlib
|
import contextlib
|
||||||
|
import copy
|
||||||
import ctypes
|
import ctypes
|
||||||
import functools
|
import functools
|
||||||
import logging
|
import logging
|
||||||
@@ -57,17 +58,18 @@ from sglang.srt.utils import (
|
|||||||
load_video,
|
load_video,
|
||||||
random_uuid,
|
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 (
|
from sglang.srt.utils.network import (
|
||||||
NetworkAddress,
|
NetworkAddress,
|
||||||
config_socket,
|
config_socket,
|
||||||
|
get_free_port,
|
||||||
get_local_ip_auto,
|
get_local_ip_auto,
|
||||||
get_zmq_socket,
|
get_zmq_socket,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
HEALTH_CHECK_TIMEOUT = 10
|
HEALTH_CHECK_TIMEOUT = 30
|
||||||
|
|
||||||
# Minimal 32x32 black PNG for health check dummy encode
|
# Minimal 32x32 black PNG for health check dummy encode
|
||||||
MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
|
MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
|
||||||
@@ -2371,10 +2373,712 @@ encoder: Optional[MMEncoder] = None
|
|||||||
send_sockets: List[zmq.Socket] = []
|
send_sockets: List[zmq.Socket] = []
|
||||||
encoder_scheduler: Optional[EncoderScheduler] = None
|
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
|
@contextlib.asynccontextmanager
|
||||||
async def _lifespan(app: FastAPI):
|
async def _lifespan(app: FastAPI):
|
||||||
global encoder_scheduler
|
global encoder_scheduler
|
||||||
|
if dp_dispatcher is not None:
|
||||||
|
dp_dispatcher.start()
|
||||||
|
yield
|
||||||
|
return
|
||||||
if encoder is not None:
|
if encoder is not None:
|
||||||
encoder_scheduler = EncoderScheduler(
|
encoder_scheduler = EncoderScheduler(
|
||||||
encoder, send_sockets, max_batch_size=ENCODER_MAX_BATCH_SIZE
|
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):
|
def launch_server(server_args: ServerArgs):
|
||||||
configure_logger(server_args, prefix=" encode_server")
|
configure_logger(server_args, prefix=" encode_server")
|
||||||
|
if server_args.dp_size > 1:
|
||||||
|
_launch_server_dp(server_args)
|
||||||
|
return
|
||||||
|
|
||||||
global encoder
|
global encoder
|
||||||
ctx = mp.get_context("spawn")
|
ctx = mp.get_context("spawn")
|
||||||
zmq_ctx = zmq.Context(10)
|
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)
|
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 def get_condition(rid):
|
||||||
async with cond_dict_lock:
|
async with cond_dict_lock:
|
||||||
if rid not in rid_to_cond:
|
if rid not in rid_to_cond:
|
||||||
@@ -2462,6 +3269,40 @@ async def get_condition(rid):
|
|||||||
async def handle_encode_request(request: dict):
|
async def handle_encode_request(request: dict):
|
||||||
req_id = request["req_id"]
|
req_id = request["req_id"]
|
||||||
start_time = time.monotonic()
|
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:
|
try:
|
||||||
# when multiple decoder TP ranks POST /encode
|
# when multiple decoder TP ranks POST /encode
|
||||||
# with the same req_id, only the first triggers the VIT forward;
|
# 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")
|
@app.post("/send")
|
||||||
async def handle_send_request(request: dict):
|
async def handle_send_request(request: dict):
|
||||||
# mooncake backend
|
# 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(
|
await encoder.send(
|
||||||
req_id=request["req_id"],
|
req_id=request["req_id"],
|
||||||
prefill_host=request["prefill_host"],
|
prefill_host=request["prefill_host"],
|
||||||
@@ -2668,6 +3532,24 @@ async def health_generate():
|
|||||||
Performs a dummy encode to verify the encoder is functional.
|
Performs a dummy encode to verify the encoder is functional.
|
||||||
Returns 200 if the encoder is healthy, 503 otherwise.
|
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:
|
if encoder is None:
|
||||||
return Response(status_code=503)
|
return Response(status_code=503)
|
||||||
|
|
||||||
@@ -2734,6 +3616,30 @@ async def health_generate():
|
|||||||
|
|
||||||
@app.api_route("/start_profile", methods=["GET", "POST"])
|
@app.api_route("/start_profile", methods=["GET", "POST"])
|
||||||
async def start_profile_async(obj: Optional[ProfileReqInput] = None):
|
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:
|
if encoder is None:
|
||||||
return Response(content="encoder not ready\n", status_code=503)
|
return Response(content="encoder not ready\n", status_code=503)
|
||||||
req = None
|
req = None
|
||||||
@@ -2772,6 +3678,12 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None):
|
|||||||
|
|
||||||
@app.api_route("/stop_profile", methods=["GET", "POST"])
|
@app.api_route("/stop_profile", methods=["GET", "POST"])
|
||||||
async def stop_profile_async():
|
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:
|
if encoder is None:
|
||||||
return Response(content="encoder not ready\n", status_code=503)
|
return Response(content="encoder not ready\n", status_code=503)
|
||||||
if encoder.profiler is None:
|
if encoder.profiler is None:
|
||||||
|
|||||||
@@ -779,6 +779,7 @@ class Envs:
|
|||||||
# Persistent receiver-side GPU embedding pool size for mooncake EPD transport.
|
# Persistent receiver-side GPU embedding pool size for mooncake EPD transport.
|
||||||
# 0 disables (per-request register/deregister). 4096 = 4GB default per TP
|
# 0 disables (per-request register/deregister). 4096 = 4GB default per TP
|
||||||
SGLANG_EMBEDDING_POOL_SIZE_MB = EnvInt(4096)
|
SGLANG_EMBEDDING_POOL_SIZE_MB = EnvInt(4096)
|
||||||
|
SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT = EnvInt(64)
|
||||||
|
|
||||||
# Elastic EP Backup Port
|
# Elastic EP Backup Port
|
||||||
SGLANG_BACKUP_PORT_BASE = EnvInt(10000)
|
SGLANG_BACKUP_PORT_BASE = EnvInt(10000)
|
||||||
|
|||||||
Reference in New Issue
Block a user