[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:
co-authored by
Zongyao Chen
Zongyao Chen
parent
facde4c6d3
commit
d939e26585
@@ -78,3 +78,42 @@ python -m sglang_router.launch_router \
|
|||||||
--port 8000
|
--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
|
||||||
|
```
|
||||||
|
|||||||
@@ -75,7 +75,7 @@ dependencies = [
|
|||||||
"uvloop",
|
"uvloop",
|
||||||
"xgrammar==0.1.27",
|
"xgrammar==0.1.27",
|
||||||
|
|
||||||
"smg-grpc-proto>=0.3.3",
|
"smg-grpc-proto>=0.4.1",
|
||||||
"grpcio>=1.78.0",
|
"grpcio>=1.78.0",
|
||||||
"grpcio-reflection>=1.78.0",
|
"grpcio-reflection>=1.78.0",
|
||||||
"grpcio-health-checking>=1.78.0",
|
"grpcio-health-checking>=1.78.0",
|
||||||
|
|||||||
@@ -68,7 +68,7 @@ dependencies = [
|
|||||||
"uvicorn",
|
"uvicorn",
|
||||||
"uvloop",
|
"uvloop",
|
||||||
"xgrammar==0.1.27",
|
"xgrammar==0.1.27",
|
||||||
"smg-grpc-proto>=0.3.3",
|
"smg-grpc-proto>=0.4.1",
|
||||||
"grpcio>=1.78.0",
|
"grpcio>=1.78.0",
|
||||||
"grpcio-reflection>=1.78.0",
|
"grpcio-reflection>=1.78.0",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ dependencies = [
|
|||||||
"uvicorn",
|
"uvicorn",
|
||||||
"uvloop",
|
"uvloop",
|
||||||
"xgrammar==0.1.27",
|
"xgrammar==0.1.27",
|
||||||
"smg-grpc-proto>=0.3.3",
|
"smg-grpc-proto>=0.4.1",
|
||||||
"grpcio>=1.78.0",
|
"grpcio>=1.78.0",
|
||||||
"grpcio-reflection>=1.78.0",
|
"grpcio-reflection>=1.78.0",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -64,7 +64,7 @@ runtime_common = [
|
|||||||
"uvicorn",
|
"uvicorn",
|
||||||
"uvloop",
|
"uvloop",
|
||||||
"xgrammar==0.1.27",
|
"xgrammar==0.1.27",
|
||||||
"smg-grpc-proto>=0.3.3",
|
"smg-grpc-proto>=0.4.1",
|
||||||
"grpcio>=1.78.0",
|
"grpcio>=1.78.0",
|
||||||
"grpcio-reflection>=1.78.0",
|
"grpcio-reflection>=1.78.0",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -67,7 +67,7 @@ dependencies = [
|
|||||||
"uvicorn",
|
"uvicorn",
|
||||||
"uvloop",
|
"uvloop",
|
||||||
# "xgrammar==0.1.24", , xgrammar depends on CUDA PyTorch and Triton only
|
# "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>=1.78.0",
|
||||||
"grpcio-reflection>=1.78.0",
|
"grpcio-reflection>=1.78.0",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -13,14 +13,21 @@ suppress_noisy_warnings()
|
|||||||
|
|
||||||
def run_server(server_args):
|
def run_server(server_args):
|
||||||
"""Run the server based on server_args.grpc_mode and server_args.encoder_only."""
|
"""Run the server based on server_args.grpc_mode and server_args.encoder_only."""
|
||||||
|
if server_args.encoder_only:
|
||||||
if server_args.grpc_mode:
|
if server_args.grpc_mode:
|
||||||
from sglang.srt.entrypoints.grpc_server import serve_grpc
|
from sglang.srt.disaggregation.encode_grpc_server import (
|
||||||
|
serve_grpc_encoder,
|
||||||
|
)
|
||||||
|
|
||||||
asyncio.run(serve_grpc(server_args))
|
asyncio.run(serve_grpc_encoder(server_args))
|
||||||
elif server_args.encoder_only:
|
else:
|
||||||
from sglang.srt.disaggregation.encode_server import launch_server
|
from sglang.srt.disaggregation.encode_server import launch_server
|
||||||
|
|
||||||
launch_server(server_args)
|
launch_server(server_args)
|
||||||
|
elif server_args.grpc_mode:
|
||||||
|
from sglang.srt.entrypoints.grpc_server import serve_grpc
|
||||||
|
|
||||||
|
asyncio.run(serve_grpc(server_args))
|
||||||
else:
|
else:
|
||||||
# Default mode: HTTP mode.
|
# Default mode: HTTP mode.
|
||||||
from sglang.srt.entrypoints.http_server import launch_server
|
from sglang.srt.entrypoints.http_server import launch_server
|
||||||
|
|||||||
@@ -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,
|
get_mooncake_transfer_engine,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
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.multimodal_processor import get_mm_processor, import_processors
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
from sglang.srt.server_args import ServerArgs
|
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
|
from sglang.srt.utils.hf_transformers_utils import get_processor
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -34,6 +34,90 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.managers.scheduler import Scheduler
|
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:
|
class EmbeddingData:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -239,6 +323,50 @@ class WaitingImageRequest:
|
|||||||
self.recv_socket.close()
|
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):
|
def _determine_tensor_transport_mode(server_args):
|
||||||
is_cross_node = server_args.dist_init_addr
|
is_cross_node = server_args.dist_init_addr
|
||||||
|
|
||||||
@@ -250,33 +378,6 @@ def _determine_tensor_transport_mode(server_args):
|
|||||||
|
|
||||||
|
|
||||||
class MMReceiverBase(ABC):
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
@@ -341,51 +442,154 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
skip_mm_pool=True,
|
skip_mm_pool=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
def create_req(self, recv_req: TokenizedGenerateReqInput):
|
@abstractmethod
|
||||||
req = Req(
|
def process_waiting_requests(self, recv_reqs):
|
||||||
recv_req.rid,
|
pass
|
||||||
recv_req.input_text,
|
|
||||||
recv_req.input_ids,
|
async def recv_mm_data(self, img_data, mm_processor, prompt):
|
||||||
recv_req.sampling_params,
|
req_id = None
|
||||||
return_logprob=recv_req.return_logprob,
|
try:
|
||||||
top_logprobs_num=recv_req.top_logprobs_num,
|
if len(self.encode_urls) == 0:
|
||||||
token_ids_logprob=recv_req.token_ids_logprob,
|
return None
|
||||||
stream=recv_req.stream,
|
req_id = uuid.uuid4().hex
|
||||||
lora_id=recv_req.lora_id,
|
embedding_port, recv_socket = get_zmq_socket_on_host(self.context, zmq.PULL)
|
||||||
input_embeds=recv_req.input_embeds,
|
if not isinstance(img_data, list):
|
||||||
custom_logit_processor=recv_req.custom_logit_processor,
|
img_data = [img_data.url]
|
||||||
require_reasoning=recv_req.require_reasoning,
|
else:
|
||||||
return_hidden_states=recv_req.return_hidden_states,
|
img_data = [img.url for img in img_data]
|
||||||
return_routed_experts=recv_req.return_routed_experts,
|
asyncio.create_task(
|
||||||
eos_token_ids=self.scheduler.model_config.hf_eos_token_id,
|
self.encode(req_id, img_data, embedding_port, "encode", "send")
|
||||||
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 await asyncio.wait_for(
|
||||||
return req
|
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
|
# For zmq_to_scheduler
|
||||||
def process_waiting_requests(self, recv_reqs):
|
def _process_waiting_requests(self, recv_reqs, waiting_cls):
|
||||||
new_recv_reqs = []
|
new_recv_reqs = []
|
||||||
for recv_req in recv_reqs:
|
for recv_req in recv_reqs:
|
||||||
if (
|
if (
|
||||||
isinstance(recv_req, TokenizedGenerateReqInput)
|
isinstance(recv_req, TokenizedGenerateReqInput)
|
||||||
and recv_req.need_wait_for_image is True
|
and recv_req.need_wait_for_image is True
|
||||||
):
|
):
|
||||||
waiting_req = WaitingImageRequest(
|
waiting_req = waiting_cls(
|
||||||
rid=recv_req.rid,
|
rid=recv_req.rid,
|
||||||
recv_req=recv_req,
|
recv_req=recv_req,
|
||||||
mm_processor=self.mm_processor,
|
mm_processor=self.mm_processor,
|
||||||
@@ -451,7 +655,6 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
self.waiting_list = new_waiting
|
self.waiting_list = new_waiting
|
||||||
return new_recv_reqs, abort_reqs
|
return new_recv_reqs, abort_reqs
|
||||||
|
|
||||||
# For zmq_to_scheduler
|
|
||||||
def _run_encode_in_thread(
|
def _run_encode_in_thread(
|
||||||
self, req_id, img_data, endpoint_encode, num_items_assigned, embedding_port
|
self, req_id, img_data, endpoint_encode, num_items_assigned, embedding_port
|
||||||
):
|
):
|
||||||
@@ -469,6 +672,80 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Encode failed for request {req_id}: {e}", exc_info=True)
|
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(
|
async def encode(
|
||||||
self,
|
self,
|
||||||
req_id,
|
req_id,
|
||||||
@@ -579,109 +856,199 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
offset += embedding_size_list_sort[idx]
|
offset += embedding_size_list_sort[idx]
|
||||||
await asyncio.gather(*metadata_tasks)
|
await asyncio.gather(*metadata_tasks)
|
||||||
|
|
||||||
# For mooncake
|
|
||||||
async def allocate_embedding_buffer(self, req_id, embedding_length, embedding_dim):
|
class MMReceiverGrpc(MMReceiverBase):
|
||||||
embeddings = torch.zeros(
|
def __init__(
|
||||||
(embedding_length, embedding_dim),
|
self,
|
||||||
dtype=self.dtype,
|
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(),
|
def build_and_send_encode_request(self, image_urls, rid):
|
||||||
embeddings.nbytes,
|
encode_req = GenerateReqInput(
|
||||||
|
image_data=[ImageData(url=url) for url in image_urls],
|
||||||
|
rid=rid,
|
||||||
)
|
)
|
||||||
self.embeddings_buffer[req_id] = embeddings
|
self.send_encode_request(encode_req)
|
||||||
return embeddings.data_ptr()
|
return encode_req
|
||||||
|
|
||||||
# For zmq_to_scheduler
|
# For zmq_to_scheduler
|
||||||
def send_encode_request(self, obj):
|
def process_waiting_requests(self, recv_reqs):
|
||||||
if type(obj.image_data) != list:
|
return self._process_waiting_requests(recv_reqs, WaitingImageRequestGrpc)
|
||||||
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
|
|
||||||
|
|
||||||
encode_idx = list(range(len(self.encode_urls)))
|
async def encode(
|
||||||
random.shuffle(encode_idx)
|
self,
|
||||||
obj.num_items_assigned = [
|
req_id,
|
||||||
(idx + len(image_urls)) // len(self.encode_urls) for idx in encode_idx
|
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(
|
num_parts = sum(1 for x in num_items_assigned if x != 0)
|
||||||
target=self._run_encode_in_thread,
|
cum_num_items = 0
|
||||||
args=(
|
cum_idx = 0
|
||||||
obj.rid,
|
for idx, assigned_num in enumerate(num_items_assigned):
|
||||||
image_urls,
|
if assigned_num == 0:
|
||||||
"encode",
|
continue
|
||||||
obj.num_items_assigned,
|
start = cum_num_items
|
||||||
None,
|
end = cum_num_items + assigned_num
|
||||||
),
|
encode_requests.append(
|
||||||
daemon=True,
|
{
|
||||||
|
"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
|
grpc_tasks = [
|
||||||
async def recv_mm_data(self, img_data, mm_processor, prompt):
|
asyncio.to_thread(
|
||||||
try:
|
_grpc_encode_request,
|
||||||
if len(self.encode_urls) == 0:
|
_grpc_target(self.encode_urls[encode_request["encoder_idx"]]),
|
||||||
return None
|
encode_request,
|
||||||
req_id = uuid.uuid4().hex
|
)
|
||||||
embedding_port, recv_socket = get_zmq_socket_on_host(self.context, zmq.PULL)
|
for encode_request in encode_requests
|
||||||
if type(img_data) != list:
|
]
|
||||||
img_data = [img_data.url]
|
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,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
if None in response_json_unsorted:
|
||||||
|
return
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
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"]]
|
||||||
|
|
||||||
|
if grpc_metadata_tasks:
|
||||||
|
await asyncio.gather(*grpc_metadata_tasks)
|
||||||
|
|
||||||
|
|
||||||
|
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:
|
else:
|
||||||
img_data = [img.url for img in img_data]
|
return
|
||||||
asyncio.create_task(
|
|
||||||
self.encode(req_id, img_data, embedding_port, "encode", "send")
|
if any(url.startswith(invalid_prefix) for url in encoder_urls):
|
||||||
|
raise ValueError(error_msg)
|
||||||
|
|
||||||
|
|
||||||
|
_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,
|
||||||
)
|
)
|
||||||
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 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
|
|
||||||
|
|
||||||
recv_embedding = None
|
|
||||||
|
|
||||||
recv_embedding_data: EmbeddingData = None
|
|
||||||
|
|
||||||
while recv_embedding_data is None or not recv_embedding_data.ready:
|
|
||||||
parts = await recv_socket.recv_multipart(copy=False)
|
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
recv_socket.close()
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|||||||
@@ -162,6 +162,14 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer)
|
|||||||
self.scheduler_info = scheduler_info
|
self.scheduler_info = scheduler_info
|
||||||
self.start_time = time.time()
|
self.start_time = time.time()
|
||||||
self.health_servicer = health_servicer
|
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
|
# Start the request manager's event loop using auto_create_handle_loop
|
||||||
self.request_manager.auto_create_handle_loop()
|
self.request_manager.auto_create_handle_loop()
|
||||||
@@ -179,6 +187,7 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer)
|
|||||||
try:
|
try:
|
||||||
# Convert gRPC request to internal format
|
# Convert gRPC request to internal format
|
||||||
tokenized_req = self._convert_generate_request(request)
|
tokenized_req = self._convert_generate_request(request)
|
||||||
|
self._handle_epd_disaggregation_encode_request(request, tokenized_req)
|
||||||
|
|
||||||
# Submit to request manager (automatically handles n>1)
|
# Submit to request manager (automatically handles n>1)
|
||||||
response_generator = self.request_manager.generate_request(
|
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}")
|
logger.info(f"Receive embedding request: {request.request_id}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Convert request
|
|
||||||
tokenized_req = self._convert_embed_request(request)
|
tokenized_req = self._convert_embed_request(request)
|
||||||
|
|
||||||
# Submit to request manager
|
|
||||||
future = await self.request_manager.embedding_request(
|
future = await self.request_manager.embedding_request(
|
||||||
obj=tokenized_req,
|
obj=tokenized_req,
|
||||||
request_id=request.request_id,
|
request_id=request.request_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Wait for result
|
|
||||||
result = await future
|
result = await future
|
||||||
|
|
||||||
# Create response
|
|
||||||
return sglang_scheduler_pb2.EmbedResponse(
|
return sglang_scheduler_pb2.EmbedResponse(
|
||||||
request_id=request.request_id,
|
request_id=request.request_id,
|
||||||
complete=sglang_scheduler_pb2.EmbedComplete(
|
complete=sglang_scheduler_pb2.EmbedComplete(
|
||||||
@@ -536,6 +541,25 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer)
|
|||||||
aggregate=_compute_aggregate_protobuf(loads),
|
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
|
# Helper methods for request/response conversion
|
||||||
|
|
||||||
def _convert_generate_request(
|
def _convert_generate_request(
|
||||||
|
|||||||
@@ -465,6 +465,11 @@ class Envs:
|
|||||||
# Health Check
|
# Health Check
|
||||||
SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION = EnvBool(True)
|
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
|
# External models
|
||||||
SGLANG_EXTERNAL_MODEL_PACKAGE = EnvStr("")
|
SGLANG_EXTERNAL_MODEL_PACKAGE = EnvStr("")
|
||||||
SGLANG_EXTERNAL_MM_MODEL_ARCH = 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 (
|
from sglang.srt.disaggregation.decode_kvcache_offload_manager import (
|
||||||
DecodeKVCacheOffloadManager,
|
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 (
|
from sglang.srt.disaggregation.prefill import (
|
||||||
PrefillBootstrapQueue,
|
PrefillBootstrapQueue,
|
||||||
SchedulerDisaggregationPrefillMixin,
|
SchedulerDisaggregationPrefillMixin,
|
||||||
@@ -982,7 +982,7 @@ class Scheduler(
|
|||||||
self.server_args.language_only
|
self.server_args.language_only
|
||||||
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
|
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
|
||||||
):
|
):
|
||||||
self.mm_receiver = MMReceiverHTTP(
|
self.mm_receiver = create_mm_receiver(
|
||||||
self.server_args,
|
self.server_args,
|
||||||
hf_config=self.model_config.hf_config,
|
hf_config=self.model_config.hf_config,
|
||||||
pp_rank=self.pp_rank,
|
pp_rank=self.pp_rank,
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ import zmq.asyncio
|
|||||||
from fastapi import BackgroundTasks
|
from fastapi import BackgroundTasks
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
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.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.lora.lora_registry import LoRARef, LoRARegistry
|
from sglang.srt.lora.lora_registry import LoRARef, LoRARegistry
|
||||||
@@ -409,7 +409,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
|
|
||||||
# Encoder Disaggregation
|
# Encoder Disaggregation
|
||||||
if self.server_args.language_only:
|
if self.server_args.language_only:
|
||||||
self.mm_receiver = MMReceiverHTTP(
|
self.mm_receiver = create_mm_receiver(
|
||||||
self.server_args,
|
self.server_args,
|
||||||
dtype=self.model_config.dtype,
|
dtype=self.model_config.dtype,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,8 +1,14 @@
|
|||||||
import os
|
import os
|
||||||
|
import subprocess
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
import unittest
|
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.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.kits.mmmu_vlm_kit import _run_lmms_eval_with_retry
|
from sglang.test.kits.mmmu_vlm_kit import _run_lmms_eval_with_retry
|
||||||
from sglang.test.server_fixtures.disaggregation_fixture import (
|
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||||
@@ -157,7 +163,7 @@ class TestEPDDisaggregationOneEncoder(PDDisaggregationServerBase):
|
|||||||
log_suffix = "openai_compatible"
|
log_suffix = "openai_compatible"
|
||||||
os.makedirs(output_path, exist_ok=True)
|
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 = [
|
cmd = [
|
||||||
"python3",
|
"python3",
|
||||||
@@ -373,7 +379,7 @@ class TestEPDDisaggregationMultiEncoders(PDDisaggregationServerBase):
|
|||||||
log_suffix = "openai_compatible"
|
log_suffix = "openai_compatible"
|
||||||
os.makedirs(output_path, exist_ok=True)
|
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 = [
|
cmd = [
|
||||||
"python3",
|
"python3",
|
||||||
@@ -425,5 +431,341 @@ class TestEPDDisaggregationMultiEncoders(PDDisaggregationServerBase):
|
|||||||
self.assertGreater(mmmu_accuracy, 0.40)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user