[EPD][refactor]: introduce BaseMMReceiver for gRPC transport integration (#17921)

This commit is contained in:
siyu
2026-02-02 11:37:32 +08:00
committed by GitHub
parent ab8b99eb23
commit f1824a957b
3 changed files with 33 additions and 6 deletions
@@ -4,6 +4,7 @@ import pickle
import random import random
import threading import threading
import uuid import uuid
from abc import ABC, abstractmethod
from enum import IntEnum from enum import IntEnum
from typing import TYPE_CHECKING, List, Optional from typing import TYPE_CHECKING, List, Optional
@@ -241,7 +242,33 @@ def _determine_tensor_transport_mode(server_args):
return "cuda_ipc" return "cuda_ipc"
class MMReceiver: 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,
@@ -602,7 +629,7 @@ class MMReceiver:
# For zmq_to_tokenizer and mooncake # For zmq_to_tokenizer and mooncake
async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt): async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt):
# Bypass MMReceiver # Bypass MMReceiverHTTP
if req_id is None: if req_id is None:
return None return None
+2 -2
View File
@@ -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 MMReceiver from sglang.srt.disaggregation.encode_receiver import MMReceiverHTTP
from sglang.srt.disaggregation.prefill import ( from sglang.srt.disaggregation.prefill import (
PrefillBootstrapQueue, PrefillBootstrapQueue,
SchedulerDisaggregationPrefillMixin, SchedulerDisaggregationPrefillMixin,
@@ -949,7 +949,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 = MMReceiver( self.mm_receiver = MMReceiverHTTP(
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,
@@ -39,7 +39,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 MMReceiver from sglang.srt.disaggregation.encode_receiver import MMReceiverHTTP
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
@@ -422,7 +422,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
# Encoder Disaggregation # Encoder Disaggregation
if self.server_args.language_only: if self.server_args.language_only:
self.mm_receiver = MMReceiver( self.mm_receiver = MMReceiverHTTP(
self.server_args, self.server_args,
dtype=self.model_config.dtype, dtype=self.model_config.dtype,
) )