[EPD][refactor]: introduce BaseMMReceiver for gRPC transport integration (#17921)
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user