[EPD][refactor]: introduce BaseMMReceiver for gRPC transport integration (#17921)
This commit is contained in:
@@ -4,6 +4,7 @@ import pickle
|
||||
import random
|
||||
import threading
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from enum import IntEnum
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
@@ -241,7 +242,33 @@ def _determine_tensor_transport_mode(server_args):
|
||||
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__(
|
||||
self,
|
||||
@@ -602,7 +629,7 @@ class MMReceiver:
|
||||
|
||||
# For zmq_to_tokenizer and mooncake
|
||||
async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt):
|
||||
# Bypass MMReceiver
|
||||
# Bypass MMReceiverHTTP
|
||||
if req_id is None:
|
||||
return None
|
||||
|
||||
|
||||
Reference in New Issue
Block a user