[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:
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
267
python/sglang/srt/disaggregation/encode_grpc_server.py
Normal file
267
python/sglang/srt/disaggregation/encode_grpc_server.py
Normal 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)
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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("")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user