Support EPD error handling (#16670)

Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
siyu
2026-01-21 18:47:00 +08:00
committed by GitHub
parent e7224e9681
commit 7520b92927
4 changed files with 310 additions and 117 deletions

View File

@@ -4,7 +4,8 @@ import pickle
import random
import threading
import uuid
from typing import List, Optional
from enum import IntEnum
from typing import TYPE_CHECKING, List, Optional
import aiohttp
import torch
@@ -16,15 +17,28 @@ from sglang.srt.disaggregation.mooncake.transfer_engine import MooncakeTransferE
from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_local_ip_auto, get_zmq_socket_on_host
from sglang.srt.utils.hf_transformers_utils import get_processor
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from sglang.srt.managers.scheduler import Scheduler
class EmbeddingData:
def __init__(self, req_id, num_parts, part_idx, image_grid_dim, embedding=None):
def __init__(
self,
req_id,
num_parts,
part_idx,
image_grid_dim,
embedding=None,
error_msg=None,
error_code=None,
):
self.req_id = req_id
self.num_parts = num_parts
self.part_idx = part_idx
@@ -42,6 +56,8 @@ class EmbeddingData:
self.image_grid_dim if i == self.part_idx else None
for i in range(self.num_parts)
]
self.error_msg = error_msg
self.error_code = error_code
def add(self, embedding_data):
assert self.req_id == embedding_data.req_id
@@ -66,7 +82,7 @@ class EmbeddingData:
return sum(self.ready_list) == self.num_parts
def __repr__(self):
return f"EmbeddingData(req_id={self.req_id}, num_parts={self.num_parts}, part_idx={self.part_idx})"
return f"EmbeddingData(req_id={self.req_id}, num_parts={self.num_parts}, part_idx={self.part_idx}) error_msg={self.error_msg}"
def copy_without_embedding(self):
new_data = EmbeddingData(
@@ -74,6 +90,8 @@ class EmbeddingData:
num_parts=self.num_parts,
part_idx=self.part_idx,
image_grid_dim=self.image_grid_dim,
error_msg=self.error_msg,
error_code=self.error_code,
)
new_data.send_time = self.send_time
new_data.dtype = self.dtype
@@ -81,6 +99,12 @@ class EmbeddingData:
return new_data
class WaitingImageRequestStatus(IntEnum):
FAIL = -1
PENDING = 0
SUCCESS = 1
# For zmq_to_scheduler
class WaitingImageRequest:
def __init__(
@@ -107,7 +131,10 @@ class WaitingImageRequest:
)
logger.info(f"Waiting for input {self.embedding_port = }")
self.recv_embedding_data = None
self.ready = False
# ok=1 pending=0 fail=-1
self.status = WaitingImageRequestStatus.PENDING
self.error_msg = None
self.error_code = None
def send_encode_request(self):
async def _send_single_request(session, url, payload):
@@ -163,7 +190,7 @@ class WaitingImageRequest:
)
def _try_recv_mm_data(self):
if self.ready:
if self.status != WaitingImageRequestStatus.PENDING:
return
while self.recv_embedding_data is None or not self.recv_embedding_data.ready:
try:
@@ -171,8 +198,17 @@ class WaitingImageRequest:
except zmq.Again:
# No data available yet, wait a bit and retry
return
recv_obj: EmbeddingData = pickle.loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None:
logger.warning(
f"Received error signal from encoder for {self.rid}: {recv_obj.error_msg} {recv_obj.error_code = }"
)
self.error_msg = recv_obj.error_msg
self.error_code = recv_obj.error_code
self.status = WaitingImageRequestStatus.FAIL
self.recv_socket.close()
return
buffer = parts[1].buffer if hasattr(parts[1], "buffer") else parts[1]
recv_obj.embedding = torch.frombuffer(buffer, dtype=recv_obj.dtype).reshape(
recv_obj.shape
@@ -191,7 +227,7 @@ class WaitingImageRequest:
)
self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = mm_inputs["input_ids"]
self.ready = True
self.status = WaitingImageRequestStatus.SUCCESS
self.recv_socket.close()
@@ -215,6 +251,7 @@ class MMReceiver:
pp_rank: Optional[int] = None,
tp_rank: Optional[int] = None,
tp_group: Optional[GroupCoordinator] = None,
scheduler: Optional["Scheduler"] = None,
):
self.context = zmq.asyncio.Context(20)
self.encoder_transfer_backend = server_args.encoder_transfer_backend
@@ -237,6 +274,7 @@ class MMReceiver:
self.nnodes = server_args.nnodes
self.hostname = get_local_ip_auto()
self.waiting_list: List[WaitingImageRequest] = []
self.scheduler = scheduler
if hf_config is not None:
transport_mode = _determine_tensor_transport_mode(server_args)
import_processors("sglang.srt.multimodal.processors")
@@ -268,6 +306,41 @@ class MMReceiver:
hf_config, server_args, _processor, transport_mode
)
def create_req(self, recv_req):
req = Req(
recv_req.rid,
recv_req.input_text,
recv_req.input_ids,
recv_req.sampling_params,
return_logprob=recv_req.return_logprob,
top_logprobs_num=recv_req.top_logprobs_num,
token_ids_logprob=recv_req.token_ids_logprob,
stream=recv_req.stream,
lora_id=recv_req.lora_id,
input_embeds=recv_req.input_embeds,
custom_logit_processor=recv_req.custom_logit_processor,
require_reasoning=recv_req.require_reasoning,
return_hidden_states=recv_req.return_hidden_states,
return_routed_experts=recv_req.return_routed_experts,
eos_token_ids=self.scheduler.model_config.hf_eos_token_id,
bootstrap_host=recv_req.bootstrap_host,
bootstrap_port=recv_req.bootstrap_port,
bootstrap_room=recv_req.bootstrap_room,
disagg_mode=self.scheduler.disaggregation_mode,
data_parallel_rank=recv_req.data_parallel_rank,
vocab_size=self.scheduler.model_config.vocab_size,
priority=recv_req.priority,
metrics_collector=(
self.scheduler.metrics_collector
if self.scheduler.enable_metrics
else None
),
http_worker_ipc=recv_req.http_worker_ipc,
dllm_config=self.scheduler.dllm_config,
)
req.tokenizer = self.scheduler.tokenizer
return req
# For zmq_to_scheduler
def process_waiting_requests(self, recv_reqs):
new_recv_reqs = []
@@ -290,12 +363,12 @@ class MMReceiver:
new_recv_reqs.append(recv_req)
if len(self.waiting_list) == 0:
return new_recv_reqs
return new_recv_reqs, []
local_status = []
for waiting_req in self.waiting_list:
waiting_req._try_recv_mm_data()
local_status.append(waiting_req.ready)
local_status.append(waiting_req.status)
local_status = torch.tensor(local_status, device="cpu", dtype=torch.int32)
@@ -306,14 +379,27 @@ class MMReceiver:
)
new_waiting = []
abort_reqs = []
for i, waiting_req in enumerate(self.waiting_list):
if local_status[i].item():
status_value = local_status[i].item()
if status_value == WaitingImageRequestStatus.SUCCESS:
new_recv_reqs.append(waiting_req.recv_req)
else:
elif status_value == WaitingImageRequestStatus.FAIL:
logger.error(
f"Waiting request {waiting_req.rid} failed: {waiting_req.error_msg} {waiting_req.error_code = }"
)
abort_reqs.append(
(
self.create_req(waiting_req.recv_req),
waiting_req.error_msg,
waiting_req.error_code,
)
)
else: # status_value == WaitingImageRequestStatus.PENDING
new_waiting.append(waiting_req)
self.waiting_list = new_waiting
return new_recv_reqs
return new_recv_reqs, abort_reqs
# For zmq_to_scheduler
def _run_encode_in_thread(
@@ -389,6 +475,16 @@ class MMReceiver:
]
responses = await asyncio.gather(*tasks)
for response in responses:
if response.status != 200:
try:
err_data = await response.json()
msg = err_data.get("message", "Unknown encoder error")
except:
msg = await response.text()
logger.error(f"Encoder returned error {response.status}: {msg}")
return
response_json_list_unsort = [
await response.json() for response in responses
]

View File

@@ -7,6 +7,7 @@ import os
import pickle
import time
import traceback
from http import HTTPStatus
from typing import Dict, List, Optional, Set, Tuple
import aiohttp
@@ -52,12 +53,30 @@ logger = logging.getLogger(__name__)
rid_lock = asyncio.Lock()
rid_to_receive_endpoint: Dict[str, List[str]] = dict()
rid_to_receive_count: Dict[str, int] = dict()
rid_to_err_msg: Dict[str, str] = dict()
use_image_processor_gpu = (
int(os.getenv("SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU", "0")) == 1
)
class MMError(Exception):
def __init__(self, message, code=HTTPStatus.INTERNAL_SERVER_ERROR):
self.message = message
self.code = code
super().__init__(self.message)
class BadRequestError(MMError):
def __init__(self, message):
super().__init__(message, code=HTTPStatus.BAD_REQUEST)
class InternalError(MMError):
def __init__(self, message):
super().__init__(message, code=HTTPStatus.INTERNAL_SERVER_ERROR)
class TensorWrapper:
"""Wrapper to keep tensor alive while exposing buffer for zero-copy."""
@@ -196,6 +215,7 @@ class MMEncoder:
)
self.embedding_to_send = dict()
self.background_tasks: Set[asyncio.Task] = set()
logger.info(f"rank {rank} init finish ")
@@ -291,53 +311,56 @@ class MMEncoder:
return await asyncio.gather(*async_futures)
async def _encode(self, mm_items) -> torch.Tensor:
images = await self._flatten_and_load_images(mm_items)
try:
images = await self._flatten_and_load_images(mm_items)
except Exception as e:
raise BadRequestError(f"Failed to load images from input: {str(e)}")
kwargs = {"device": self.device} if self.use_image_processor_gpu else {}
images_input = self.image_processor(images=images, **kwargs)
feature = images_input["pixel_values"]
mm_item = MultimodalDataItem.from_dict(
{
"modality": Modality.IMAGE,
"feature": _convert(feature),
}
)
for k, v in images_input.items():
if k == "pixel_values":
continue
mm_item.set(k, _convert(v))
try:
kwargs = {"device": self.device} if self.use_image_processor_gpu else {}
images_input = self.image_processor(images=images, **kwargs)
feature = images_input["pixel_values"]
mm_item = MultimodalDataItem.from_dict(
{
"modality": Modality.IMAGE,
"feature": _convert(feature),
}
)
for k, v in images_input.items():
if k == "pixel_values":
continue
mm_item.set(k, _convert(v))
# support mm_cache
mm_embedding = None
mm_hash = None
# support mm_cache
mm_embedding = None
mm_hash = None
start_time = time.perf_counter()
if self.server_args.enable_prefix_mm_cache:
mm_item.set_pad_value()
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
async with self.mm_cache_lock:
mm_cache = self.mm_cache.get([mm_item.hash])
if mm_cache is not None:
mm_embedding = mm_cache.embedding
if self.server_args.enable_prefix_mm_cache:
mm_item.set_pad_value()
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
async with self.mm_cache_lock:
mm_cache = self.mm_cache.get([mm_item.hash])
if mm_cache is not None:
mm_embedding = mm_cache.embedding
if mm_embedding is None:
with torch.inference_mode():
mm_embedding: torch.Tensor = self.model.get_image_feature([mm_item])
mm_embedding = mm_embedding.cpu()
if len(mm_embedding.shape) != 2:
mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1])
if mm_embedding is None:
with torch.inference_mode():
mm_embedding: torch.Tensor = self.model.get_image_feature([mm_item])
mm_embedding = mm_embedding.cpu()
if len(mm_embedding.shape) != 2:
mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1])
if self.server_args.enable_prefix_mm_cache:
async with self.mm_cache_lock:
self.mm_cache.set(mm_hash, EmbeddingResult(embedding=mm_embedding))
end_time = time.perf_counter()
logger.info(
f"Vit time : {(end_time - start_time)*1000:.2f} ms {mm_embedding.shape = }"
)
if self.profiler is not None:
self.profiler.step()
if self.server_args.enable_prefix_mm_cache:
async with self.mm_cache_lock:
self.mm_cache.set(mm_hash, EmbeddingResult(embedding=mm_embedding))
if self.profiler is not None:
self.profiler.step()
return _get_image_grid_dim(images_input), mm_embedding
return _get_image_grid_dim(images_input), mm_embedding
except BadRequestError as e:
raise BadRequestError(f"Bad request error: {str(e)}")
except Exception as e:
raise InternalError(f"Internal encoding error: {str(e)}")
async def _send(
self,
@@ -377,26 +400,47 @@ class MMEncoder:
socket.send_multipart([pickle.dumps(mm_data)])
else:
new_mm_data = mm_data.copy_without_embedding()
if new_mm_data.error_msg is not None:
socket.send_multipart([pickle.dumps(new_mm_data)])
return
embedding_tensor = TensorWrapper(mm_data.embedding)
socket.send_multipart(
[pickle.dumps(new_mm_data), embedding_tensor.__buffer__()]
)
async def encode(self, mm_items, req_id, num_parts, part_idx):
start_time = time.time()
image_grid_dim, mm_embedding = await self._encode(mm_items)
end_time = time.time()
logger.info(f"🕛 encode cost = {(end_time - start_time) * 1000:.2f}ms")
if self.rank == 0:
mm_data = EmbeddingData(
req_id,
num_parts,
part_idx,
image_grid_dim,
mm_embedding,
try:
image_grid_dim, mm_embedding = await self._encode(mm_items)
if self.rank == 0:
mm_data = EmbeddingData(
req_id, num_parts, part_idx, image_grid_dim, mm_embedding
)
self.embedding_to_send[req_id] = mm_data
return (
mm_embedding.nbytes,
mm_embedding.shape[0],
mm_embedding.shape[1],
None,
None,
)
self.embedding_to_send[mm_data.req_id] = mm_data
return mm_embedding.nbytes, mm_embedding.shape[0], mm_embedding.shape[1]
except Exception as e:
error_code = getattr(e, "code", HTTPStatus.INTERNAL_SERVER_ERROR)
error_msg = str(e)
logger.error(f"Rank {self.rank} encode failed: {error_msg} {error_code = }")
if self.rank == 0:
mm_data = EmbeddingData(
req_id,
num_parts,
part_idx,
None,
error_msg=error_msg,
error_code=error_code,
)
self.embedding_to_send[req_id] = mm_data
logger.debug(f"Created error EmbeddingData: {mm_data}")
return 0, 0, 0, error_msg, error_code
# For zmq_to_tokenizer zmq_to_scheduler and mooncake
async def send(
@@ -624,55 +668,93 @@ def launch_server(server_args: ServerArgs):
@app.post("/encode")
async def handle_encode_request(request: dict):
# broadcast request
request.update({"enter_time": time.time()})
for socket in send_sockets:
socket.send_pyobj(request)
req_id = request["req_id"]
try:
nbytes, embedding_len, embedding_dim = await encoder.encode(
mm_items=request["mm_items"],
req_id=request["req_id"],
num_parts=request["num_parts"],
part_idx=request["part_idx"],
)
if encoder.server_args.encoder_transfer_backend == "mooncake":
del request["mm_items"]
request.update(
{
"embedding_size": nbytes,
"embedding_len": embedding_len,
"embedding_dim": embedding_dim,
}
)
return ORJSONResponse(content=request)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
logger.info(f"{request['embedding_port'] = }")
if request["embedding_port"] is None:
await encoder.send_with_url(
def start_background_send(req_id):
task = asyncio.create_task(encoder.send_with_url(req_id=req_id))
encoder.background_tasks.add(task)
task.add_done_callback(encoder.background_tasks.discard)
# broadcast request
request.update({"enter_time": time.time()})
for socket in send_sockets:
socket.send_pyobj(request)
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
await encoder.encode(
mm_items=request["mm_items"],
req_id=request["req_id"],
num_parts=request["num_parts"],
part_idx=request["part_idx"],
)
else:
assert type(request["embedding_port"]) == list
tasks = []
for embedding_port in request["embedding_port"]:
tasks.append(
encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=embedding_port,
)
)
await asyncio.gather(*tasks)
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
await encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=request["embedding_port"],
)
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None)
if error_msg:
if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
if request["embedding_port"] is None:
start_background_send(req_id)
else:
for port in request["embedding_port"]:
await encoder.send(
req_id=req_id,
prefill_host=request["prefill_host"],
embedding_port=port,
)
return ORJSONResponse(
status_code=error_code,
content={"status": "error", "message": error_msg, "req_id": req_id},
)
if encoder.server_args.encoder_transfer_backend == "mooncake":
del request["mm_items"]
request.update(
{
"embedding_size": nbytes,
"embedding_len": embedding_len,
"embedding_dim": embedding_dim,
}
)
return ORJSONResponse(content=request)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
logger.info(f"{request['embedding_port'] = }")
if request["embedding_port"] is None:
await encoder.send_with_url(
req_id=request["req_id"],
)
else:
assert type(request["embedding_port"]) == list
tasks = []
for embedding_port in request["embedding_port"]:
tasks.append(
encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=embedding_port,
)
)
await asyncio.gather(*tasks)
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
await encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=request["embedding_port"],
)
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None)
except Exception as e:
error_msg = str(e)
logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}")
rid_to_err_msg[req_id] = error_msg
return ORJSONResponse(
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
content={
"status": "error",
"message": error_msg,
"req_id": req_id,
},
)
@app.post("/send")
@@ -746,7 +828,9 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None):
f"profile_id={encoder.profiler.profile_id}\n"
)
return Response(content=detail, status_code=200)
return Response(content=(msg or "Start profiling failed.\n"), status_code=400)
return Response(
content=(msg or "Start profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST
)
@app.api_route("/stop_profile", methods=["GET", "POST"])
@@ -754,11 +838,15 @@ async def stop_profile_async():
if encoder is None:
return Response(content="encoder not ready\n", status_code=503)
if encoder.profiler is None:
return Response(content="profiling not initialized\n", status_code=400)
return Response(
content="profiling not initialized\n", status_code=HTTPStatus.BAD_REQUEST
)
req = ProfileReq(ProfileReqType.STOP_PROFILE)
for socket in send_sockets:
socket.send_pyobj(req)
ok, msg = encoder.profiler.stop()
if ok:
return Response(content="Stop profiling.\n", status_code=200)
return Response(content=(msg or "Stop profiling failed.\n"), status_code=400)
return Response(
content=(msg or "Stop profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST
)

View File

@@ -959,9 +959,10 @@ class Scheduler(
self.mm_receiver = MMReceiver(
self.server_args,
hf_config=self.model_config.hf_config,
tp_rank=self.tp_rank,
pp_rank=self.pp_rank,
tp_rank=self.tp_rank,
tp_group=self.tp_group,
scheduler=self,
)
def init_overlap(self):
@@ -1260,7 +1261,16 @@ class Scheduler(
and self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
):
recv_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
for req, error_msg, error_code in abort_reqs:
status_code = (
HTTPStatus.BAD_REQUEST
if error_code == 400
else HTTPStatus.INTERNAL_SERVER_ERROR
)
prepare_abort(req, error_msg, status_code=status_code)
self.stream_output([req], req.return_logprob)
if self.enable_trace:
for req in recv_reqs:

View File

@@ -1072,7 +1072,6 @@ class SchedulerOutputProcessorMixin:
if reqs or is_idle_batch:
if self.model_config.is_multimodal_gen:
return
self.send_to_detokenizer.send_output(
BatchTokenIDOutput(
rids=rids,