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