diff --git a/docs/advanced_features/epd_disaggregation.md b/docs/advanced_features/epd_disaggregation.md index 550503dfc..c543c29bc 100644 --- a/docs/advanced_features/epd_disaggregation.md +++ b/docs/advanced_features/epd_disaggregation.md @@ -78,3 +78,42 @@ python -m sglang_router.launch_router \ --port 8000 ``` + +#### gRPC Encoder (EPD) + +You can run the encoder as a gRPC server while keeping prefill/decode as HTTP. +When using gRPC encoders, set `SGLANG_ENCODER_MM_RECEIVER_MODE=grpc` for the +prefill process so it uses the gRPC receiver. + +```bash +# gRPC encoder +python -m sglang.launch_server \ + --model-path Qwen/Qwen3-VL-8B-Instruct \ + --encoder-only \ + --grpc-mode \ + --encoder-transfer-backend zmq_to_scheduler \ + --port 30000 + +# prefill (HTTP) - tell it to use gRPC receiver +SGLANG_ENCODER_MM_RECEIVER_MODE=grpc \ +python -m sglang.launch_server \ + --model-path Qwen/Qwen3-VL-8B-Instruct \ + --disaggregation-mode prefill \ + --language-only \ + --encoder-urls grpc://127.0.0.1:30000 \ + --encoder-transfer-backend zmq_to_scheduler \ + --port 30002 + +# decode (HTTP) +python -m sglang.launch_server \ + --model-path Qwen/Qwen3-VL-8B-Instruct \ + --disaggregation-mode decode \ + --port 30003 + +# router +python -m sglang_router.launch_router \ + --pd-disaggregation \ + --prefill http://$PREFILL_HOST:30002 \ + --decode http://$DECODE_HOST:30003 \ + --port 8000 +``` diff --git a/python/pyproject.toml b/python/pyproject.toml index 2f8452f3c..65a07c88a 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -75,7 +75,7 @@ dependencies = [ "uvloop", "xgrammar==0.1.27", - "smg-grpc-proto>=0.3.3", + "smg-grpc-proto>=0.4.1", "grpcio>=1.78.0", "grpcio-reflection>=1.78.0", "grpcio-health-checking>=1.78.0", diff --git a/python/pyproject_cpu.toml b/python/pyproject_cpu.toml index ff7ac6aa0..38ae6490d 100644 --- a/python/pyproject_cpu.toml +++ b/python/pyproject_cpu.toml @@ -68,7 +68,7 @@ dependencies = [ "uvicorn", "uvloop", "xgrammar==0.1.27", - "smg-grpc-proto>=0.3.3", + "smg-grpc-proto>=0.4.1", "grpcio>=1.78.0", "grpcio-reflection>=1.78.0", ] diff --git a/python/pyproject_npu.toml b/python/pyproject_npu.toml index 78a746143..da87a936b 100644 --- a/python/pyproject_npu.toml +++ b/python/pyproject_npu.toml @@ -62,7 +62,7 @@ dependencies = [ "uvicorn", "uvloop", "xgrammar==0.1.27", - "smg-grpc-proto>=0.3.3", + "smg-grpc-proto>=0.4.1", "grpcio>=1.78.0", "grpcio-reflection>=1.78.0", ] diff --git a/python/pyproject_other.toml b/python/pyproject_other.toml index 718f7886d..6dccf38a3 100755 --- a/python/pyproject_other.toml +++ b/python/pyproject_other.toml @@ -64,7 +64,7 @@ runtime_common = [ "uvicorn", "uvloop", "xgrammar==0.1.27", - "smg-grpc-proto>=0.3.3", + "smg-grpc-proto>=0.4.1", "grpcio>=1.78.0", "grpcio-reflection>=1.78.0", ] diff --git a/python/pyproject_xpu.toml b/python/pyproject_xpu.toml index 14d1c0d3d..0f9dc3466 100644 --- a/python/pyproject_xpu.toml +++ b/python/pyproject_xpu.toml @@ -67,7 +67,7 @@ dependencies = [ "uvicorn", "uvloop", # "xgrammar==0.1.24", , xgrammar depends on CUDA PyTorch and Triton only - "smg-grpc-proto>=0.3.3", + "smg-grpc-proto>=0.4.1", "grpcio>=1.78.0", "grpcio-reflection>=1.78.0", ] diff --git a/python/sglang/launch_server.py b/python/sglang/launch_server.py index 45ea8ecfd..af4a41e62 100644 --- a/python/sglang/launch_server.py +++ b/python/sglang/launch_server.py @@ -13,14 +13,21 @@ suppress_noisy_warnings() def run_server(server_args): """Run the server based on server_args.grpc_mode and server_args.encoder_only.""" - if server_args.grpc_mode: + if server_args.encoder_only: + if server_args.grpc_mode: + from sglang.srt.disaggregation.encode_grpc_server import ( + serve_grpc_encoder, + ) + + asyncio.run(serve_grpc_encoder(server_args)) + else: + from sglang.srt.disaggregation.encode_server import launch_server + + launch_server(server_args) + elif server_args.grpc_mode: from sglang.srt.entrypoints.grpc_server import serve_grpc asyncio.run(serve_grpc(server_args)) - elif server_args.encoder_only: - from sglang.srt.disaggregation.encode_server import launch_server - - launch_server(server_args) else: # Default mode: HTTP mode. from sglang.srt.entrypoints.http_server import launch_server diff --git a/python/sglang/srt/disaggregation/encode_grpc_server.py b/python/sglang/srt/disaggregation/encode_grpc_server.py new file mode 100644 index 000000000..23778e814 --- /dev/null +++ b/python/sglang/srt/disaggregation/encode_grpc_server.py @@ -0,0 +1,267 @@ +""" +gRPC Encoder Server for SGLang EPD (Encode-Prefill-Decode) mode. + +This server provides gRPC-based encoding for multimodal inputs. + +Usage: + python -m sglang.launch_server --model-path --encoder-only --grpc-mode +""" + +import asyncio +import logging +import multiprocessing as mp +import traceback +from concurrent import futures +from typing import List + +import grpc +import zmq +import zmq.asyncio +from grpc_health.v1 import health_pb2, health_pb2_grpc +from grpc_reflection.v1alpha import reflection +from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc + +from sglang.srt.disaggregation.encode_server import ( + MMEncoder, + handle_scheduler_receive_url_request, + launch_encoder, +) +from sglang.srt.server_args import PortArgs, ServerArgs +from sglang.srt.utils import get_zmq_socket, random_uuid + +logger = logging.getLogger(__name__) +SGLangEncoderServicer = sglang_encoder_pb2_grpc.SglangEncoderServicer +add_SGLangEncoderServicer_to_server = ( + sglang_encoder_pb2_grpc.add_SglangEncoderServicer_to_server +) + + +class EncoderHealthServicer(health_pb2_grpc.HealthServicer): + """ + Standard gRPC health check service for encoder server. + Implements grpc.health.v1.Health for Kubernetes probes. + """ + + OVERALL_SERVER = "" + ENCODER_SERVICE = "sglang.grpc.encoder.SglangEncoder" + + def __init__(self): + self._serving = False + + def set_serving(self): + self._serving = True + + def set_not_serving(self): + self._serving = False + + async def Check(self, request, context) -> health_pb2.HealthCheckResponse: + if self._serving: + return health_pb2.HealthCheckResponse( + status=health_pb2.HealthCheckResponse.SERVING + ) + return health_pb2.HealthCheckResponse( + status=health_pb2.HealthCheckResponse.NOT_SERVING + ) + + async def Watch(self, request, context): + yield await self.Check(request, context) + + +class SGLangEncoderServer(SGLangEncoderServicer): + """ + gRPC service implementation for SGLang encoder. + """ + + def __init__( + self, + encoder: MMEncoder, + send_sockets: List[zmq.Socket], + server_args: ServerArgs, + ): + self.encoder = encoder + self.send_sockets = send_sockets + self.server_args = server_args + + async def Encode( + self, request: sglang_encoder_pb2.EncodeRequest, context + ) -> sglang_encoder_pb2.EncodeResponse: + try: + request_dict = { + "mm_items": list(request.mm_items), + "req_id": request.req_id, + "num_parts": request.num_parts, + "part_idx": request.part_idx, + } + for socket in self.send_sockets: + await socket.send_pyobj(request_dict) + + ( + nbytes, + embedding_len, + embedding_dim, + error_msg, + error_code, + ) = await self.encoder.encode( + mm_items=list(request.mm_items), + req_id=request.req_id, + num_parts=request.num_parts, + part_idx=request.part_idx, + ) + if error_msg is not None: + context.set_code(grpc.StatusCode.INTERNAL) + context.set_details(error_msg) + return sglang_encoder_pb2.EncodeResponse() + + if self.server_args.encoder_transfer_backend == "mooncake": + return sglang_encoder_pb2.EncodeResponse( + embedding_size=nbytes, + embedding_len=embedding_len, + embedding_dim=embedding_dim, + ) + elif self.server_args.encoder_transfer_backend == "zmq_to_scheduler": + embedding_ports = list(request.embedding_port) + logger.info(f"embedding_port = {embedding_ports}") + if not embedding_ports: + await self.encoder.send_with_url(req_id=request.req_id) + else: + tasks = [] + for embedding_port in embedding_ports: + tasks.append( + self.encoder.send( + req_id=request.req_id, + prefill_host=request.prefill_host, + embedding_port=embedding_port, + ) + ) + await asyncio.gather(*tasks) + self.encoder.embedding_to_send.pop(request.req_id, None) + return sglang_encoder_pb2.EncodeResponse() + elif self.server_args.encoder_transfer_backend == "zmq_to_tokenizer": + embedding_port = ( + request.embedding_port[0] if request.embedding_port else 0 + ) + await self.encoder.send( + req_id=request.req_id, + prefill_host=request.prefill_host, + embedding_port=embedding_port, + ) + self.encoder.embedding_to_send.pop(request.req_id, None) + return sglang_encoder_pb2.EncodeResponse() + + return sglang_encoder_pb2.EncodeResponse() + + except Exception as e: + logger.error(f"Encode error: {e}") + traceback.print_exc() + context.set_code(grpc.StatusCode.INTERNAL) + context.set_details(str(e)) + return sglang_encoder_pb2.EncodeResponse() + + async def Send( + self, request: sglang_encoder_pb2.SendRequest, context + ) -> sglang_encoder_pb2.SendResponse: + try: + await self.encoder.send( + req_id=request.req_id, + prefill_host=request.prefill_host, + embedding_port=request.embedding_port, + session_id=request.session_id if request.session_id else None, + buffer_address=( + request.buffer_address if request.buffer_address else None + ), + ) + self.encoder.embedding_to_send.pop(request.req_id, None) + return sglang_encoder_pb2.SendResponse() + + except Exception as e: + logger.error(f"Send error: {e}") + traceback.print_exc() + context.set_code(grpc.StatusCode.INTERNAL) + context.set_details(str(e)) + return sglang_encoder_pb2.SendResponse() + + async def SchedulerReceiveUrl( + self, request: sglang_encoder_pb2.SchedulerReceiveUrlRequest, context + ) -> sglang_encoder_pb2.SchedulerReceiveUrlResponse: + try: + await handle_scheduler_receive_url_request( + { + "req_id": request.req_id, + "receive_count": request.receive_count, + "receive_url": request.receive_url, + } + ) + return sglang_encoder_pb2.SchedulerReceiveUrlResponse() + + except Exception as e: + logger.error(f"SchedulerReceiveUrl error: {e}") + traceback.print_exc() + context.set_code(grpc.StatusCode.INTERNAL) + context.set_details(str(e)) + return sglang_encoder_pb2.SchedulerReceiveUrlResponse() + + +async def serve_grpc_encoder(server_args: ServerArgs): + ctx = mp.get_context("spawn") + zmq_ctx = zmq.asyncio.Context(10) + ipc_path_prefix = random_uuid() + port_args = PortArgs.init_new(server_args) + + if server_args.dist_init_addr: + dist_init_method = f"tcp://{server_args.dist_init_addr}" + else: + dist_init_method = f"tcp://127.0.0.1:{port_args.nccl_port}" + + send_sockets: List[zmq.Socket] = [] + for rank in range(1, server_args.tp_size): + schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" + send_sockets.append( + get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False) + ) + ctx.Process( + target=launch_encoder, + args=(server_args, schedule_path, dist_init_method, rank), + daemon=True, + ).start() + + encoder = MMEncoder(server_args, dist_init_method=dist_init_method) + + server = grpc.aio.server( + futures.ThreadPoolExecutor(max_workers=10), + options=[ + ("grpc.max_send_message_length", 1024 * 1024 * 256), + ("grpc.max_receive_message_length", 1024 * 1024 * 256), + ], + ) + + health_servicer = EncoderHealthServicer() + health_pb2_grpc.add_HealthServicer_to_server(health_servicer, server) + + encoder_servicer = SGLangEncoderServer( + encoder=encoder, + send_sockets=send_sockets, + server_args=server_args, + ) + add_SGLangEncoderServicer_to_server(encoder_servicer, server) + + SERVICE_NAMES = ( + sglang_encoder_pb2.DESCRIPTOR.services_by_name["SglangEncoder"].full_name, + "grpc.health.v1.Health", + reflection.SERVICE_NAME, + ) + reflection.enable_server_reflection(SERVICE_NAMES, server) + + listen_addr = f"{server_args.host}:{server_args.port}" + server.add_insecure_port(listen_addr) + + await server.start() + logger.info(f"gRPC encoder server listening on {listen_addr}") + + health_servicer.set_serving() + + try: + await server.wait_for_termination() + except KeyboardInterrupt: + logger.info("Shutting down gRPC encoder server...") + health_servicer.set_not_serving() + await server.stop(grace=5) diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index b3d7cfbeb..6af99220c 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -21,11 +21,11 @@ from sglang.srt.distributed.parallel_state import ( get_mooncake_transfer_engine, ) from sglang.srt.environ import envs -from sglang.srt.managers.io_struct import TokenizedGenerateReqInput +from sglang.srt.managers.io_struct import GenerateReqInput, 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 import ImageData, get_local_ip_auto, get_zmq_socket_on_host from sglang.srt.utils.hf_transformers_utils import get_processor logger = logging.getLogger(__name__) @@ -34,6 +34,90 @@ if TYPE_CHECKING: from sglang.srt.managers.scheduler import Scheduler +def _grpc_target(url: str) -> str: + if url.startswith("grpc://"): + return url[len("grpc://") :] + if url.startswith("grpcs://"): + raise ValueError("grpcs:// is not supported; use grpc://") + return url + + +def _normalize_embedding_ports(embedding_port): + if embedding_port is None: + return [] + if isinstance(embedding_port, list): + return embedding_port + return [embedding_port] + + +def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count): + import grpc + from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc + + timeout_secs = envs.SGLANG_ENCODER_GRPC_TIMEOUT_SECS.get() + channel = grpc.insecure_channel(target) + stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel) + try: + stub.SchedulerReceiveUrl( + sglang_encoder_pb2.SchedulerReceiveUrlRequest( + req_id=req_id, + receive_url=receive_url, + receive_count=receive_count, + ), + timeout=timeout_secs, + ) + finally: + channel.close() + + +def _grpc_encode_request(target, encode_request): + import grpc + from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc + + timeout_secs = envs.SGLANG_ENCODER_GRPC_TIMEOUT_SECS.get() + channel = grpc.insecure_channel(target) + stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel) + try: + response = stub.Encode( + sglang_encoder_pb2.EncodeRequest( + mm_items=encode_request["mm_items"], + req_id=encode_request["req_id"], + num_parts=encode_request["num_parts"], + part_idx=encode_request["part_idx"], + prefill_host=encode_request["prefill_host"], + embedding_port=_normalize_embedding_ports( + encode_request["embedding_port"] + ), + ), + timeout=timeout_secs, + ) + return response + finally: + channel.close() + + +def _grpc_send_request(target, request_json): + import grpc + from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc + + timeout_secs = envs.SGLANG_ENCODER_GRPC_TIMEOUT_SECS.get() + channel = grpc.insecure_channel(target) + stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel) + try: + stub.Send( + sglang_encoder_pb2.SendRequest( + req_id=request_json["req_id"], + prefill_host=request_json["prefill_host"], + embedding_port=request_json["embedding_port"], + session_id=request_json["session_id"], + buffer_address=request_json["buffer_address"], + ), + timeout=timeout_secs, + ) + finally: + channel.close() + + class EmbeddingData: def __init__( self, @@ -239,6 +323,50 @@ class WaitingImageRequest: self.recv_socket.close() +class WaitingImageRequestGrpc(WaitingImageRequest): + def send_encode_request(self): + async def send_embedding_port(req_id, receive_count, host_name, embedding_port): + tasks = [] + logger.info(f"{self.num_items_assigned = } ") + for idx, assigned_num in enumerate(self.num_items_assigned): + if assigned_num == 0: + continue + encoder_url = self.encoder_urls[idx] + receive_url = f"{host_name}:{embedding_port}" + target_url = f"{encoder_url}/SchedulerReceiveUrl" + logger.info(f"Preparing to send to {target_url}") + tasks.append( + asyncio.to_thread( + _grpc_scheduler_receive_url, + _grpc_target(encoder_url), + req_id, + receive_url, + receive_count, + ) + ) + + if not tasks: + logger.info("No tasks to send.") + return + logger.info(f"Concurrently sending {len(tasks)} requests...") + results = await asyncio.gather(*tasks, return_exceptions=True) + + for i, result in enumerate(results): + if isinstance(result, Exception): + logger.error(f"Request {i} failed: {result}") + else: + logger.debug(f"Request {i} succeeded.") + + asyncio.run( + send_embedding_port( + self.recv_req.rid, + self.receive_count, + self.host_name, + self.embedding_port, + ) + ) + + def _determine_tensor_transport_mode(server_args): is_cross_node = server_args.dist_init_addr @@ -250,33 +378,6 @@ def _determine_tensor_transport_mode(server_args): class MMReceiverBase(ABC): - def __init__( - self, - server_args: ServerArgs, - dtype: Optional[torch.dtype] = None, - hf_config: Optional[PretrainedConfig] = None, - pp_rank: Optional[int] = None, - tp_rank: Optional[int] = None, - tp_group: Optional[GroupCoordinator] = None, - scheduler: Optional["Scheduler"] = None, - ): - pass - - @abstractmethod - def process_waiting_requests(self, recv_reqs): - pass - - @abstractmethod - async def recv_mm_data(self, img_data, mm_processor, prompt): - pass - - @abstractmethod - def send_encode_request(self, obj): - pass - - -class MMReceiverHTTP(MMReceiverBase): - def __init__( self, server_args: ServerArgs, @@ -341,51 +442,154 @@ class MMReceiverHTTP(MMReceiverBase): skip_mm_pool=True, ) - def create_req(self, recv_req: TokenizedGenerateReqInput): - 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, - routed_dp_rank=recv_req.routed_dp_rank, - disagg_prefill_dp_rank=recv_req.disagg_prefill_dp_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 + @abstractmethod + def process_waiting_requests(self, recv_reqs): + pass + + async def recv_mm_data(self, img_data, mm_processor, prompt): + req_id = None + try: + if len(self.encode_urls) == 0: + return None + req_id = uuid.uuid4().hex + embedding_port, recv_socket = get_zmq_socket_on_host(self.context, zmq.PULL) + if not isinstance(img_data, list): + img_data = [img_data.url] + else: + img_data = [img.url for img in img_data] + asyncio.create_task( + self.encode(req_id, img_data, embedding_port, "encode", "send") + ) + return await asyncio.wait_for( + self._recv_mm_data(req_id, recv_socket, mm_processor, prompt), + timeout=20, + ) + except asyncio.TimeoutError: + logger.warning(f"Embedding recv timeout for request {req_id}") + if req_id is not None: + self._cleanup_mooncake_buffer(req_id) + return None + + def _cleanup_mooncake_buffer(self, req_id): + if self.encoder_transfer_backend != "mooncake": + return + if not hasattr(self, "embeddings_buffer"): + return + embeddings = self.embeddings_buffer.pop(req_id, None) + if embeddings is None: + return + try: + self.embeddings_engine.deregister(embeddings.data_ptr()) + except Exception: + logger.exception( + "mooncake: failed to deregister buffer for req_id=%s", req_id + ) + + async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt): + if req_id is None: + return None + + recv_embedding = None + + recv_embedding_data: EmbeddingData = None + + try: + while recv_embedding_data is None or not recv_embedding_data.ready: + parts = await recv_socket.recv_multipart(copy=False) + if not parts: + continue + recv_obj: EmbeddingData = pickle.loads(parts[0]) + if getattr(recv_obj, "error_msg", None) is not None: + logger.warning( + f"Encoder error for req_id={req_id}: {recv_obj.error_msg} " + f"error_code={getattr(recv_obj, 'error_code', None)}" + ) + self._cleanup_mooncake_buffer(req_id) + return None + logger.debug("recv_obj=%s", recv_obj) + if self.encoder_transfer_backend == "zmq_to_tokenizer": + if len(parts) < 2: + logger.error( + "zmq_to_tokenizer expected 2-part message, got %d parts", + len(parts), + ) + return None + buffer = ( + parts[1].buffer if hasattr(parts[1], "buffer") else parts[1] + ) + # Clone so we don't depend on ZMQ buffer after next recv. + recv_obj.embedding = ( + torch.frombuffer(buffer, dtype=recv_obj.dtype) + .reshape(recv_obj.shape) + .clone() + ) + if recv_embedding_data is None: + recv_obj.embedding_list[recv_obj.part_idx] = recv_obj.embedding + recv_embedding_data = recv_obj + else: + recv_embedding_data.add(recv_obj) + + if self.encoder_transfer_backend == "mooncake": + if req_id not in self.embeddings_buffer: + logger.error( + "mooncake: embeddings_buffer missing req_id=%s", req_id + ) + return None + recv_embedding = self.embeddings_buffer[req_id] + del self.embeddings_buffer[req_id] + self.embeddings_engine.deregister(recv_embedding.data_ptr()) + elif self.encoder_transfer_backend == "zmq_to_tokenizer": + recv_embedding = recv_embedding_data.get_embedding(is_concat=True) + + img_grid_thw = recv_embedding_data.get_img_grid() + mm_inputs = mm_processor.get_mm_data(prompt, recv_embedding, img_grid_thw) + return mm_inputs + finally: + recv_socket.close() + + def send_encode_request(self, obj): + self._send_encode_request(obj) + + def _send_encode_request(self, obj): + if obj.image_data is None: + image_urls = [] + elif not isinstance(obj.image_data, list): + image_urls = [obj.image_data.url] + else: + image_urls = [img.url for img in obj.image_data] + if obj.rid is None: + obj.rid = uuid.uuid4().hex + if image_urls and self.encode_urls: + logger.info(f"Processing {len(image_urls)} images for request {obj.rid}") + obj.need_wait_for_image = True + + encode_idx = list(range(len(self.encode_urls))) + random.shuffle(encode_idx) + obj.num_items_assigned = [ + (idx + len(image_urls)) // len(self.encode_urls) for idx in encode_idx + ] + encode_thread = threading.Thread( + target=self._run_encode_in_thread, + args=( + obj.rid, + image_urls, + "encode", + obj.num_items_assigned, + None, + ), + daemon=True, + ) + encode_thread.start() # For zmq_to_scheduler - def process_waiting_requests(self, recv_reqs): + def _process_waiting_requests(self, recv_reqs, waiting_cls): new_recv_reqs = [] for recv_req in recv_reqs: if ( isinstance(recv_req, TokenizedGenerateReqInput) and recv_req.need_wait_for_image is True ): - waiting_req = WaitingImageRequest( + waiting_req = waiting_cls( rid=recv_req.rid, recv_req=recv_req, mm_processor=self.mm_processor, @@ -451,7 +655,6 @@ class MMReceiverHTTP(MMReceiverBase): self.waiting_list = new_waiting return new_recv_reqs, abort_reqs - # For zmq_to_scheduler def _run_encode_in_thread( self, req_id, img_data, endpoint_encode, num_items_assigned, embedding_port ): @@ -469,6 +672,80 @@ class MMReceiverHTTP(MMReceiverBase): except Exception as e: logger.error(f"Encode failed for request {req_id}: {e}", exc_info=True) + def create_req(self, recv_req: TokenizedGenerateReqInput): + 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, + routed_dp_rank=recv_req.routed_dp_rank, + disagg_prefill_dp_rank=recv_req.disagg_prefill_dp_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 + + async def allocate_embedding_buffer(self, req_id, embedding_length, embedding_dim): + embeddings = torch.zeros( + (embedding_length, embedding_dim), + dtype=self.dtype, + ) + self.embeddings_engine.register( + embeddings.data_ptr(), + embeddings.nbytes, + ) + self.embeddings_buffer[req_id] = embeddings + return embeddings.data_ptr() + + +class MMReceiverHTTP(MMReceiverBase): + def __init__( + self, + server_args: ServerArgs, + dtype: Optional[torch.dtype] = None, + hf_config: Optional[PretrainedConfig] = None, + pp_rank: Optional[int] = None, + tp_rank: Optional[int] = None, + tp_group: Optional[GroupCoordinator] = None, + scheduler: Optional["Scheduler"] = None, + ): + super().__init__( + server_args, + dtype=dtype, + hf_config=hf_config, + pp_rank=pp_rank, + tp_rank=tp_rank, + tp_group=tp_group, + scheduler=scheduler, + ) + + # For zmq_to_scheduler + def process_waiting_requests(self, recv_reqs): + return self._process_waiting_requests(recv_reqs, WaitingImageRequest) + async def encode( self, req_id, @@ -579,109 +856,199 @@ class MMReceiverHTTP(MMReceiverBase): offset += embedding_size_list_sort[idx] await asyncio.gather(*metadata_tasks) - # For mooncake - async def allocate_embedding_buffer(self, req_id, embedding_length, embedding_dim): - embeddings = torch.zeros( - (embedding_length, embedding_dim), - dtype=self.dtype, + +class MMReceiverGrpc(MMReceiverBase): + def __init__( + self, + server_args: ServerArgs, + dtype: Optional[torch.dtype] = None, + hf_config: Optional[PretrainedConfig] = None, + pp_rank: Optional[int] = None, + tp_rank: Optional[int] = None, + tp_group: Optional[GroupCoordinator] = None, + scheduler: Optional["Scheduler"] = None, + ): + super().__init__( + server_args, + dtype=dtype, + hf_config=hf_config, + pp_rank=pp_rank, + tp_rank=tp_rank, + tp_group=tp_group, + scheduler=scheduler, ) - self.embeddings_engine.register( - embeddings.data_ptr(), - embeddings.nbytes, + + def build_and_send_encode_request(self, image_urls, rid): + encode_req = GenerateReqInput( + image_data=[ImageData(url=url) for url in image_urls], + rid=rid, ) - self.embeddings_buffer[req_id] = embeddings - return embeddings.data_ptr() + self.send_encode_request(encode_req) + return encode_req # For zmq_to_scheduler - def send_encode_request(self, obj): - if type(obj.image_data) != list: - image_urls = [obj.image_data.url] - else: - image_urls = [img.url for img in obj.image_data] - if obj.rid is None: - obj.rid = uuid.uuid4().hex - if image_urls and len(image_urls) > 0: - logger.info(f"Processing {len(image_urls)} images for request {obj.rid}") - obj.need_wait_for_image = True + def process_waiting_requests(self, recv_reqs): + return self._process_waiting_requests(recv_reqs, WaitingImageRequestGrpc) - encode_idx = list(range(len(self.encode_urls))) - random.shuffle(encode_idx) - obj.num_items_assigned = [ - (idx + len(image_urls)) // len(self.encode_urls) for idx in encode_idx + async def encode( + self, + req_id, + img_data, + embedding_port, + endpoint_encode, + endpoint_send, + num_items_assigned=None, + ): + if not img_data: + return + + encode_requests = [] + if num_items_assigned is None: + random.shuffle(self.encode_idx) + num_items_assigned = [ + (idx + len(img_data)) // len(self.encode_urls) + for idx in self.encode_idx ] - encode_thread = threading.Thread( - target=self._run_encode_in_thread, - args=( - obj.rid, - image_urls, - "encode", - obj.num_items_assigned, - None, - ), - daemon=True, + num_parts = sum(1 for x in num_items_assigned if x != 0) + cum_num_items = 0 + cum_idx = 0 + for idx, assigned_num in enumerate(num_items_assigned): + if assigned_num == 0: + continue + start = cum_num_items + end = cum_num_items + assigned_num + encode_requests.append( + { + "encoder_idx": idx, + "mm_items": img_data[start:end], + "num_parts": num_parts, + "part_idx": cum_idx, + "req_id": req_id, + "prefill_host": self.host, + "embedding_port": embedding_port, + } ) - encode_thread.start() + cum_idx += 1 + cum_num_items += assigned_num - # For zmq_to_tokenizer and mooncake - async def recv_mm_data(self, img_data, mm_processor, prompt): - try: - if len(self.encode_urls) == 0: - return None - req_id = uuid.uuid4().hex - embedding_port, recv_socket = get_zmq_socket_on_host(self.context, zmq.PULL) - if type(img_data) != list: - img_data = [img_data.url] - else: - img_data = [img.url for img in img_data] - asyncio.create_task( - self.encode(req_id, img_data, embedding_port, "encode", "send") + grpc_tasks = [ + asyncio.to_thread( + _grpc_encode_request, + _grpc_target(self.encode_urls[encode_request["encoder_idx"]]), + encode_request, ) - return await asyncio.wait_for( - self._recv_mm_data(req_id, recv_socket, mm_processor, prompt), - timeout=20, + for encode_request in encode_requests + ] + grpc_responses = await asyncio.gather(*grpc_tasks) + response_json_unsorted = [] + for encode_request, response in zip(encode_requests, grpc_responses): + if self.encoder_transfer_backend == "zmq_to_scheduler": + response_json_unsorted.append(None) + continue + response_json_unsorted.append( + { + "req_id": encode_request["req_id"], + "prefill_host": encode_request["prefill_host"], + "embedding_port": encode_request["embedding_port"], + "encoder_idx": encode_request["encoder_idx"], + "part_idx": encode_request["part_idx"], + "embedding_size": response.embedding_size, + "embedding_len": response.embedding_len, + "embedding_dim": response.embedding_dim, + } ) - except asyncio.TimeoutError: - logger.warning(f"Embedding recv timeout for request {req_id}") - if hasattr(self, "embeddings_buffer") and req_id in self.embeddings_buffer: - del self.embeddings_buffer[req_id] - return None - # For zmq_to_tokenizer and mooncake - async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt): - # Bypass MMReceiverHTTP - if req_id is None: - return None + if None in response_json_unsorted: + return - recv_embedding = None + embedding_size_by_part = [None for _ in range(num_parts)] + embedding_length_tot = 0 + response_json_sorted = [None for _ in range(num_parts)] + for response_json in response_json_unsorted: + idx = response_json["part_idx"] + embedding_size_by_part[idx] = response_json["embedding_size"] + embedding_length_tot += response_json["embedding_len"] + response_json_sorted[idx] = response_json - recv_embedding_data: EmbeddingData = None + offset = 0 + buffer_address = await self.allocate_embedding_buffer( + req_id, + embedding_length_tot, + response_json_sorted[0]["embedding_dim"], + ) + grpc_metadata_tasks = [] + for response_json in response_json_sorted: + response_json.update( + { + "session_id": self.embeddings_engine.session_id, + "buffer_address": offset + buffer_address, + } + ) + grpc_metadata_tasks.append( + asyncio.to_thread( + _grpc_send_request, + _grpc_target(self.encode_urls[response_json["encoder_idx"]]), + response_json, + ) + ) + offset += embedding_size_by_part[response_json["part_idx"]] - while recv_embedding_data is None or not recv_embedding_data.ready: - parts = await recv_socket.recv_multipart(copy=False) + if grpc_metadata_tasks: + await asyncio.gather(*grpc_metadata_tasks) - recv_obj: EmbeddingData = pickle.loads(parts[0]) - logger.info(f"{recv_obj = }") - if self.encoder_transfer_backend == "zmq_to_tokenizer": - 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) - if recv_embedding_data is None: - recv_obj.embedding_list[recv_obj.part_idx] = recv_obj.embedding - recv_embedding_data = recv_obj - else: - recv_embedding_data.add(recv_obj) - if self.encoder_transfer_backend == "mooncake": - recv_embedding = self.embeddings_buffer[req_id] - del self.embeddings_buffer[req_id] - self.embeddings_engine.deregister(recv_embedding.data_ptr()) - elif self.encoder_transfer_backend == "zmq_to_tokenizer": - recv_embedding = recv_embedding_data.get_embedding(is_concat=True) +def _validate_transport_mode(transport_mode: str, encoder_urls): + if transport_mode == "grpc": + invalid_prefix = "http://" + error_msg = ( + "EPD MMReceiver: grpc mode requires grpc:// encoder URLs. " + "Set SGLANG_ENCODER_MM_RECEIVER_MODE=http for http:// URLs." + ) + elif transport_mode == "http": + invalid_prefix = "grpc://" + error_msg = ( + "EPD MMReceiver: http mode requires http:// encoder URLs. " + "Set SGLANG_ENCODER_MM_RECEIVER_MODE=grpc for grpc:// URLs." + ) + else: + return - recv_socket.close() + if any(url.startswith(invalid_prefix) for url in encoder_urls): + raise ValueError(error_msg) - img_grid_thw = recv_embedding_data.get_img_grid() - mm_inputs = mm_processor.get_mm_data(prompt, recv_embedding, img_grid_thw) - return mm_inputs +_MM_RECEIVER_BY_MODE = { + "grpc": MMReceiverGrpc, + "http": MMReceiverHTTP, +} + + +def create_mm_receiver( + server_args: ServerArgs, + dtype: Optional[torch.dtype] = None, + hf_config: Optional[PretrainedConfig] = None, + pp_rank: Optional[int] = None, + tp_rank: Optional[int] = None, + tp_group: Optional[GroupCoordinator] = None, + scheduler: Optional["Scheduler"] = None, + transport_mode: Optional[str] = None, +): + if transport_mode is None: + transport_mode = envs.SGLANG_ENCODER_MM_RECEIVER_MODE.get() + logger.debug(f"MMReceiver transport_mode from env: {transport_mode}") + + _validate_transport_mode(transport_mode, server_args.encoder_urls) + logger.info(f"EPD MMReceiver: using transport_mode={transport_mode}") + + receiver_cls = _MM_RECEIVER_BY_MODE.get(transport_mode) + if receiver_cls is None: + raise ValueError(f"Unsupported transport_mode: {transport_mode}") + return receiver_cls( + server_args, + dtype=dtype, + hf_config=hf_config, + pp_rank=pp_rank, + tp_rank=tp_rank, + tp_group=tp_group, + scheduler=scheduler, + ) diff --git a/python/sglang/srt/entrypoints/grpc_server.py b/python/sglang/srt/entrypoints/grpc_server.py index 0e85fa30d..6e257d52d 100644 --- a/python/sglang/srt/entrypoints/grpc_server.py +++ b/python/sglang/srt/entrypoints/grpc_server.py @@ -162,6 +162,14 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer) self.scheduler_info = scheduler_info self.start_time = time.time() self.health_servicer = health_servicer + self.mm_receiver = None + if ( + self.server_args.language_only + and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" + ): + from sglang.srt.disaggregation import encode_receiver as mm_receiver + + self.mm_receiver = mm_receiver.create_mm_receiver(self.server_args) # Start the request manager's event loop using auto_create_handle_loop self.request_manager.auto_create_handle_loop() @@ -179,6 +187,7 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer) try: # Convert gRPC request to internal format tokenized_req = self._convert_generate_request(request) + self._handle_epd_disaggregation_encode_request(request, tokenized_req) # Submit to request manager (automatically handles n>1) response_generator = self.request_manager.generate_request( @@ -248,19 +257,15 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer) logger.info(f"Receive embedding request: {request.request_id}") try: - # Convert request tokenized_req = self._convert_embed_request(request) - # Submit to request manager future = await self.request_manager.embedding_request( obj=tokenized_req, request_id=request.request_id, ) - # Wait for result result = await future - # Create response return sglang_scheduler_pb2.EmbedResponse( request_id=request.request_id, complete=sglang_scheduler_pb2.EmbedComplete( @@ -536,6 +541,25 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer) aggregate=_compute_aggregate_protobuf(loads), ) + def _handle_epd_disaggregation_encode_request( + self, + grpc_req: sglang_scheduler_pb2.GenerateRequest, + tokenized_req: TokenizedGenerateReqInput, + ) -> None: + if not self.mm_receiver: + return + + image_urls = list(grpc_req.mm_inputs.image_urls) + if not image_urls: + return + + encode_req = self.mm_receiver.build_and_send_encode_request( + image_urls=image_urls, + rid=grpc_req.request_id, + ) + tokenized_req.need_wait_for_image = bool(encode_req.need_wait_for_image) + tokenized_req.num_items_assigned = encode_req.num_items_assigned + # Helper methods for request/response conversion def _convert_generate_request( diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index aa5db62d5..b1f2a8328 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -465,6 +465,11 @@ class Envs: # Health Check SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION = EnvBool(True) + # Encoder gRPC + SGLANG_ENCODER_GRPC_TIMEOUT_SECS = EnvInt(60) + # Encoder receiver selection: http|grpc (used by EPD paths). + SGLANG_ENCODER_MM_RECEIVER_MODE = EnvStr("http") + # External models SGLANG_EXTERNAL_MODEL_PACKAGE = EnvStr("") SGLANG_EXTERNAL_MM_MODEL_ARCH = EnvStr("") diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index d050b5c40..7c4806c4b 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -43,7 +43,7 @@ from sglang.srt.disaggregation.decode import ( from sglang.srt.disaggregation.decode_kvcache_offload_manager import ( DecodeKVCacheOffloadManager, ) -from sglang.srt.disaggregation.encode_receiver import MMReceiverHTTP +from sglang.srt.disaggregation.encode_receiver import create_mm_receiver from sglang.srt.disaggregation.prefill import ( PrefillBootstrapQueue, SchedulerDisaggregationPrefillMixin, @@ -982,7 +982,7 @@ class Scheduler( self.server_args.language_only and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" ): - self.mm_receiver = MMReceiverHTTP( + self.mm_receiver = create_mm_receiver( self.server_args, hf_config=self.model_config.hf_config, pp_rank=self.pp_rank, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 7beaa1d42..afca49f6f 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -38,7 +38,7 @@ import zmq.asyncio from fastapi import BackgroundTasks from sglang.srt.configs.model_config import ModelConfig -from sglang.srt.disaggregation.encode_receiver import MMReceiverHTTP +from sglang.srt.disaggregation.encode_receiver import create_mm_receiver from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.environ import envs from sglang.srt.lora.lora_registry import LoRARef, LoRARegistry @@ -409,7 +409,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi # Encoder Disaggregation if self.server_args.language_only: - self.mm_receiver = MMReceiverHTTP( + self.mm_receiver = create_mm_receiver( self.server_args, dtype=self.model_config.dtype, ) diff --git a/test/registered/distributed/test_epd_disaggregation.py b/test/registered/distributed/test_epd_disaggregation.py index 298655cfa..e5de792f6 100644 --- a/test/registered/distributed/test_epd_disaggregation.py +++ b/test/registered/distributed/test_epd_disaggregation.py @@ -1,8 +1,14 @@ import os +import subprocess import threading +import time import unittest -from sglang.srt.utils import kill_process_tree +import grpc +import zmq +from grpc_health.v1 import health_pb2, health_pb2_grpc + +from sglang.srt.utils import get_zmq_socket_on_host, kill_process_tree from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.mmmu_vlm_kit import _run_lmms_eval_with_retry from sglang.test.server_fixtures.disaggregation_fixture import ( @@ -157,7 +163,7 @@ class TestEPDDisaggregationOneEncoder(PDDisaggregationServerBase): log_suffix = "openai_compatible" os.makedirs(output_path, exist_ok=True) - model_args = f'model_version="{model_version}",' f"tp={tp}" + model_args = f'model_version="{model_version}",tp={tp}' cmd = [ "python3", @@ -373,7 +379,7 @@ class TestEPDDisaggregationMultiEncoders(PDDisaggregationServerBase): log_suffix = "openai_compatible" os.makedirs(output_path, exist_ok=True) - model_args = f'model_version="{model_version}",' f"tp={tp}" + model_args = f'model_version="{model_version}",tp={tp}' cmd = [ "python3", @@ -425,5 +431,341 @@ class TestEPDDisaggregationMultiEncoders(PDDisaggregationServerBase): self.assertGreater(mmmu_accuracy, 0.40) +@unittest.skipIf(is_in_ci(), "Skipping in CI to reduce multi-GPU runtime") +class TestEPDDisaggregationGrpcEncoderMMMU(PDDisaggregationServerBase): + """Test MMMU evaluation with gRPC encoder in EPD mode.""" + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = DEFAULT_SMALL_VLM_MODEL_NAME_FOR_TEST + cls.encode_port = f"{int(cls.lb_port) + 304}" + cls.encode_url = f"grpc://{cls.base_host}:{cls.encode_port}" + + print( + f"Setting up gRPC EPD (one encoder): encode={cls.encode_port}, " + f"prefill={cls.prefill_port}, decode={cls.decode_port}" + ) + + cls.start_encode() + prefill_thread = threading.Thread(target=cls.start_prefill) + decode_thread = threading.Thread(target=cls.start_decode) + prefill_thread.start() + decode_thread.start() + prefill_thread.join() + decode_thread.join() + + cls.wait_grpc_ready(cls.base_host, cls.encode_port, cls.process_encode) + cls.wait_server_ready(cls.prefill_url + "/health") + cls.wait_server_ready(cls.decode_url + "/health") + + cls.launch_lb() + + cls.api_key = "sk-123456" + os.environ["OPENAI_API_KEY"] = cls.api_key + os.environ["OPENAI_API_BASE"] = f"{cls.lb_url}/v1" + + @classmethod + def start_encode(cls): + encode_command = [ + "python3", + "-m", + "sglang.launch_server", + "--model-path", + cls.model, + "--host", + cls.base_host, + "--port", + cls.encode_port, + "--trust-remote-code", + "--encoder-only", + "--grpc-mode", + "--encoder-transfer-backend", + "zmq_to_scheduler", + "--tp", + "1", + "--base-gpu-id", + "0", + "--enable-prefix-mm-cache", + ] + cls.process_encode = subprocess.Popen(encode_command) + + @classmethod + def start_prefill(cls): + prefill_args = [ + "--trust-remote-code", + "--language-only", + "--encoder-urls", + cls.encode_url, + "--encoder-transfer-backend", + "zmq_to_scheduler", + "--disaggregation-mode", + "prefill", + "--tp", + "1", + "--base-gpu-id", + "1", + "--port", + cls.prefill_port, + ] + prefill_args += cls.transfer_backend + cls.rdma_devices + prefill_env = os.environ.copy() + prefill_env["SGLANG_ENCODER_MM_RECEIVER_MODE"] = "grpc" + cls.process_prefill = popen_launch_server( + cls.model, + base_url=cls.prefill_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=prefill_args, + env=prefill_env, + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--tp", + "1", + "--base-gpu-id", + "2", + "--port", + cls.decode_port, + ] + decode_args += cls.transfer_backend + cls.rdma_devices + cls.process_decode = popen_launch_server( + cls.model, + base_url=cls.decode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=decode_args, + ) + + @staticmethod + def wait_grpc_ready(host, port, process, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH): + deadline = time.time() + timeout + channel = grpc.insecure_channel(f"{host}:{port}") + stub = health_pb2_grpc.HealthStub(channel) + try: + while time.time() < deadline: + if process.poll() is not None: + raise RuntimeError( + f"gRPC encoder server exited with code {process.returncode}" + ) + try: + response = stub.Check( + health_pb2.HealthCheckRequest(service=""), timeout=2 + ) + if response.status == health_pb2.HealthCheckResponse.SERVING: + return + except grpc.RpcError: + pass + time.sleep(1) + finally: + channel.close() + + raise RuntimeError( + f"gRPC encoder server not ready at {host}:{port} within {timeout}s" + ) + + @classmethod + def tearDownClass(cls): + os.environ.pop("SGLANG_ENCODER_MM_RECEIVER_MODE", None) + os.environ.pop("OPENAI_API_KEY", None) + os.environ.pop("OPENAI_API_BASE", None) + for process in [ + cls.process_lb, + cls.process_decode, + cls.process_prefill, + cls.process_encode, + ]: + if process: + try: + kill_process_tree(process.pid) + except Exception as e: + print(f"Error killing process: {e}") + + def run_mmmu_eval(self, model_version: str, output_path: str, limit: str = "50"): + model = "openai_compatible" + tp = 1 + tasks = "mmmu_val" + batch_size = 32 + log_suffix = "openai_compatible" + os.makedirs(output_path, exist_ok=True) + + model_args = f'model_version="{model_version}",tp={tp}' + + cmd = [ + "python3", + "-m", + "lmms_eval", + "--model", + model, + "--model_args", + model_args, + "--tasks", + tasks, + "--batch_size", + str(batch_size), + "--log_samples", + "--log_samples_suffix", + log_suffix, + "--output_path", + str(output_path), + "--limit", + limit, + ] + + _run_lmms_eval_with_retry(cmd, timeout=3600) + + def test_mmmu(self): + import glob + import json + + output_path = "./logs/epd_grpc_encoder_mmmu" + self.run_mmmu_eval(self.model, output_path) + + result_files = glob.glob(f"{output_path}/**/*.json", recursive=True) + if not result_files: + result_files = glob.glob(f"{output_path}/*.json") + + if not result_files: + self.fail(f"No JSON result files found in {output_path}") + + result_file_path = result_files[0] + with open(result_file_path, "r") as f: + result = json.load(f) + print(f"MMMU result (grpc encoder): {result}") + + mmmu_accuracy = result["results"]["mmmu_val"]["mmmu_acc,none"] + print(f"MMMU accuracy (grpc encoder): {mmmu_accuracy:.4f}") + # for qwen2.5-vl-3b-instruct, the accuracy is 0.40 + self.assertGreater(mmmu_accuracy, 0.40) + + +@unittest.skipIf(is_in_ci(), "Skipping in CI to reduce multi-GPU runtime") +class TestEPDDisaggregationGrpcEncoderOnly(PDDisaggregationServerBase): + """Test gRPC encoder server integration with zmq_to_scheduler transfers.""" + + @classmethod + def setUpClass(cls): + super().setUpClass() + os.environ["SGLANG_ENCODER_MM_RECEIVER_MODE"] = "grpc" + cls.model = DEFAULT_SMALL_VLM_MODEL_NAME_FOR_TEST + cls.encode_port = f"{int(cls.lb_port) + 302}" + + print(f"Setting up gRPC EPD encoder: encode={cls.encode_port}") + + cls.start_encode() + cls.wait_grpc_ready(cls.base_host, cls.encode_port, cls.process_encode) + + @classmethod + def start_encode(cls): + encode_command = [ + "python3", + "-m", + "sglang.launch_server", + "--model-path", + cls.model, + "--host", + cls.base_host, + "--port", + cls.encode_port, + "--trust-remote-code", + "--encoder-only", + "--grpc-mode", + "--encoder-transfer-backend", + "zmq_to_scheduler", + "--tp", + "1", + "--base-gpu-id", + "0", + "--enable-prefix-mm-cache", + ] + cls.process_encode = subprocess.Popen(encode_command) + + @staticmethod + def wait_grpc_ready(host, port, process, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH): + deadline = time.time() + timeout + channel = grpc.insecure_channel(f"{host}:{port}") + stub = health_pb2_grpc.HealthStub(channel) + try: + while time.time() < deadline: + if process.poll() is not None: + raise RuntimeError( + f"gRPC encoder server exited with code {process.returncode}" + ) + try: + response = stub.Check( + health_pb2.HealthCheckRequest(service=""), timeout=2 + ) + if response.status == health_pb2.HealthCheckResponse.SERVING: + return + except grpc.RpcError: + pass + time.sleep(1) + finally: + channel.close() + + raise RuntimeError( + f"gRPC encoder server not ready at {host}:{port} within {timeout}s" + ) + + @classmethod + def tearDownClass(cls): + os.environ.pop("SGLANG_ENCODER_MM_RECEIVER_MODE", None) + if cls.process_encode: + try: + kill_process_tree(cls.process_encode.pid) + except Exception as e: + print(f"Error killing process: {e}") + super().tearDownClass() + + def test_grpc_encoder_zmq_to_scheduler(self): + from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc + + context = zmq.Context() + recv_port, recv_socket = get_zmq_socket_on_host( + context, zmq.PULL, host=self.base_host + ) + channel = grpc.insecure_channel(f"{self.base_host}:{self.encode_port}") + stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel) + req_id = f"grpc-epd-{int(time.time() * 1000)}" + image_path = os.path.abspath("examples/assets/example_image.png") + + try: + stub.SchedulerReceiveUrl( + sglang_encoder_pb2.SchedulerReceiveUrlRequest( + req_id=req_id, + receive_url=f"{self.base_host}:{recv_port}", + receive_count=1, + ), + timeout=60, + ) + stub.Encode( + sglang_encoder_pb2.EncodeRequest( + mm_items=[image_path], + req_id=req_id, + num_parts=1, + part_idx=0, + ), + timeout=300, + ) + + poller = zmq.Poller() + poller.register(recv_socket, zmq.POLLIN) + socks = dict(poller.poll(60000)) + self.assertIn( + recv_socket, + socks, + "No embedding payload received from gRPC encoder server", + ) + parts = recv_socket.recv_multipart() + self.assertTrue(parts, "Empty embedding payload from gRPC encoder server") + finally: + recv_socket.close() + context.term() + channel.close() + + if __name__ == "__main__": unittest.main()