diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index 8c84f122c..2e6fa1735 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -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 ] diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 9d347025d..6583b64a3 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -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 + ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 6c324b297..99f14e3ef 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 984025685..bf74f81dc 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -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,