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