[model gateway][0/N] router EPD support: add encoder grpc server backend support (#16552)

Co-authored-by: Zongyao Chen <ZongYao.Chen@linux.alibaba.com>
Co-authored-by: Zongyao Chen <solar1s@163.com>
This commit is contained in:
Jasonzhang517
2026-03-03 19:38:15 +08:00
committed by GitHub
parent facde4c6d3
commit d939e26585
14 changed files with 1226 additions and 175 deletions

View File

@@ -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",

View File

@@ -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",
]

View File

@@ -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",
]

View File

@@ -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",
]

View File

@@ -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",
]

View File

@@ -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

View File

@@ -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 <model> --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)

View File

@@ -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,
)

View File

@@ -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(

View File

@@ -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("")

View File

@@ -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,

View File

@@ -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,
)