Support PP for zmq_to_scheduler (#15312)
This commit is contained in:
@@ -4,14 +4,16 @@ import pickle
|
|||||||
import random
|
import random
|
||||||
import threading
|
import threading
|
||||||
import uuid
|
import uuid
|
||||||
from typing import List
|
from typing import List, Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
import torch
|
import torch
|
||||||
import zmq
|
import zmq
|
||||||
import zmq.asyncio
|
import zmq.asyncio
|
||||||
|
from transformers import PretrainedConfig
|
||||||
|
|
||||||
from sglang.srt.disaggregation.mooncake.transfer_engine import MooncakeTransferEngine
|
from sglang.srt.disaggregation.mooncake.transfer_engine import MooncakeTransferEngine
|
||||||
|
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||||
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput
|
from sglang.srt.managers.io_struct import 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.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
@@ -89,7 +91,6 @@ class WaitingImageRequest:
|
|||||||
encoder_urls,
|
encoder_urls,
|
||||||
host_name,
|
host_name,
|
||||||
receive_count,
|
receive_count,
|
||||||
embedding_port=None,
|
|
||||||
):
|
):
|
||||||
self.rid = rid
|
self.rid = rid
|
||||||
self.recv_req = recv_req
|
self.recv_req = recv_req
|
||||||
@@ -209,10 +210,11 @@ class MMReceiver:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
dtype=None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
hf_config=None,
|
hf_config: Optional[PretrainedConfig] = None,
|
||||||
pp_rank=None,
|
pp_rank: Optional[int] = None,
|
||||||
tp_rank=None,
|
tp_rank: Optional[int] = None,
|
||||||
|
tp_group: Optional[GroupCoordinator] = None,
|
||||||
):
|
):
|
||||||
self.context = zmq.asyncio.Context(20)
|
self.context = zmq.asyncio.Context(20)
|
||||||
self.encoder_transfer_backend = server_args.encoder_transfer_backend
|
self.encoder_transfer_backend = server_args.encoder_transfer_backend
|
||||||
@@ -231,9 +233,9 @@ class MMReceiver:
|
|||||||
self.pp_rank = pp_rank
|
self.pp_rank = pp_rank
|
||||||
self.tp_rank = tp_rank
|
self.tp_rank = tp_rank
|
||||||
self.tp_size = server_args.tp_size
|
self.tp_size = server_args.tp_size
|
||||||
|
self.tp_group = tp_group
|
||||||
self.nnodes = server_args.nnodes
|
self.nnodes = server_args.nnodes
|
||||||
self.hostname = get_local_ip_auto()
|
self.hostname = get_local_ip_auto()
|
||||||
self.world_size = server_args.pp_size * server_args.tp_size
|
|
||||||
self.waiting_list: List[WaitingImageRequest] = []
|
self.waiting_list: List[WaitingImageRequest] = []
|
||||||
if hf_config is not None:
|
if hf_config is not None:
|
||||||
transport_mode = _determine_tensor_transport_mode(server_args)
|
transport_mode = _determine_tensor_transport_mode(server_args)
|
||||||
@@ -270,27 +272,19 @@ class MMReceiver:
|
|||||||
def process_waiting_requests(self, recv_reqs):
|
def process_waiting_requests(self, recv_reqs):
|
||||||
new_recv_reqs = []
|
new_recv_reqs = []
|
||||||
for recv_req in recv_reqs:
|
for recv_req in recv_reqs:
|
||||||
# E Disaggregation
|
|
||||||
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
|
||||||
):
|
):
|
||||||
embedding_port = None
|
|
||||||
if recv_req.embedding_ports is not None:
|
|
||||||
embedding_port = recv_req.embedding_ports[
|
|
||||||
self.tp_size * self.pp_rank + self.tp_rank
|
|
||||||
]
|
|
||||||
waiting_req = WaitingImageRequest(
|
waiting_req = WaitingImageRequest(
|
||||||
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,
|
||||||
encoder_urls=self.encode_urls,
|
encoder_urls=self.encode_urls,
|
||||||
host_name=self.hostname,
|
host_name=self.hostname,
|
||||||
receive_count=self.world_size,
|
receive_count=self.tp_size,
|
||||||
embedding_port=embedding_port,
|
|
||||||
)
|
)
|
||||||
if recv_req.embedding_ports is None:
|
waiting_req.send_encode_request()
|
||||||
waiting_req.send_encode_request()
|
|
||||||
self.waiting_list.append(waiting_req)
|
self.waiting_list.append(waiting_req)
|
||||||
else:
|
else:
|
||||||
new_recv_reqs.append(recv_req)
|
new_recv_reqs.append(recv_req)
|
||||||
@@ -303,9 +297,13 @@ class MMReceiver:
|
|||||||
waiting_req._try_recv_mm_data()
|
waiting_req._try_recv_mm_data()
|
||||||
local_status.append(waiting_req.ready)
|
local_status.append(waiting_req.ready)
|
||||||
|
|
||||||
local_status = torch.tensor(local_status, device="cuda", dtype=torch.int32)
|
local_status = torch.tensor(local_status, device="cpu", dtype=torch.int32)
|
||||||
|
|
||||||
torch.distributed.all_reduce(local_status, op=torch.distributed.ReduceOp.MIN)
|
torch.distributed.all_reduce(
|
||||||
|
local_status,
|
||||||
|
op=torch.distributed.ReduceOp.MIN,
|
||||||
|
group=self.tp_group.cpu_group,
|
||||||
|
)
|
||||||
|
|
||||||
new_waiting = []
|
new_waiting = []
|
||||||
for i, waiting_req in enumerate(self.waiting_list):
|
for i, waiting_req in enumerate(self.waiting_list):
|
||||||
@@ -465,7 +463,6 @@ class MMReceiver:
|
|||||||
obj.num_items_assigned = [
|
obj.num_items_assigned = [
|
||||||
(idx + len(image_urls)) // len(self.encode_urls) for idx in encode_idx
|
(idx + len(image_urls)) // len(self.encode_urls) for idx in encode_idx
|
||||||
]
|
]
|
||||||
obj.embedding_ports = None
|
|
||||||
encode_thread = threading.Thread(
|
encode_thread = threading.Thread(
|
||||||
target=self._run_encode_in_thread,
|
target=self._run_encode_in_thread,
|
||||||
args=(
|
args=(
|
||||||
@@ -473,7 +470,7 @@ class MMReceiver:
|
|||||||
image_urls,
|
image_urls,
|
||||||
"encode",
|
"encode",
|
||||||
obj.num_items_assigned,
|
obj.num_items_assigned,
|
||||||
obj.embedding_ports,
|
None,
|
||||||
),
|
),
|
||||||
daemon=True,
|
daemon=True,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -253,9 +253,9 @@ class GenerateReqInput(BaseReq, APIServingTimingMixin):
|
|||||||
# Whether to return entropy
|
# Whether to return entropy
|
||||||
return_entropy: bool = False
|
return_entropy: bool = False
|
||||||
|
|
||||||
|
# For EPD-disaggregated inference
|
||||||
need_wait_for_image: Optional[bool] = None
|
need_wait_for_image: Optional[bool] = None
|
||||||
num_items_assigned: Optional[List] = None
|
num_items_assigned: Optional[List] = None
|
||||||
embedding_ports: Optional[List] = None
|
|
||||||
|
|
||||||
def contains_mm_input(self) -> bool:
|
def contains_mm_input(self) -> bool:
|
||||||
return (
|
return (
|
||||||
@@ -742,7 +742,6 @@ class TokenizedGenerateReqInput(BaseReq):
|
|||||||
|
|
||||||
need_wait_for_image: bool = False
|
need_wait_for_image: bool = False
|
||||||
num_items_assigned: Optional[List] = None
|
num_items_assigned: Optional[List] = None
|
||||||
embedding_ports: Optional[List] = None
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -955,6 +955,7 @@ class Scheduler(
|
|||||||
hf_config=self.model_config.hf_config,
|
hf_config=self.model_config.hf_config,
|
||||||
tp_rank=self.tp_rank,
|
tp_rank=self.tp_rank,
|
||||||
pp_rank=self.pp_rank,
|
pp_rank=self.pp_rank,
|
||||||
|
tp_group=self.tp_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_overlap(self):
|
def init_overlap(self):
|
||||||
@@ -1239,6 +1240,14 @@ class Scheduler(
|
|||||||
src=self.tp_group.ranks[0],
|
src=self.tp_group.ranks[0],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Process MM requests under EPD-disaggregation mode
|
||||||
|
if (
|
||||||
|
self.pp_rank == 0
|
||||||
|
and self.server_args.language_only
|
||||||
|
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
|
||||||
|
):
|
||||||
|
recv_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
|
||||||
|
|
||||||
if self.enable_trace:
|
if self.enable_trace:
|
||||||
for req in recv_reqs:
|
for req in recv_reqs:
|
||||||
if isinstance(
|
if isinstance(
|
||||||
@@ -1279,12 +1288,6 @@ class Scheduler(
|
|||||||
return work_reqs, control_reqs
|
return work_reqs, control_reqs
|
||||||
|
|
||||||
def process_input_requests(self, recv_reqs: List):
|
def process_input_requests(self, recv_reqs: List):
|
||||||
# Process MM requests under EPD-disaggregation mode
|
|
||||||
if (
|
|
||||||
self.server_args.language_only
|
|
||||||
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
|
|
||||||
):
|
|
||||||
recv_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
|
|
||||||
|
|
||||||
for recv_req in recv_reqs:
|
for recv_req in recv_reqs:
|
||||||
# If it is a health check generation request and there are running requests, ignore it.
|
# If it is a health check generation request and there are running requests, ignore it.
|
||||||
|
|||||||
@@ -860,7 +860,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
extra_key=obj.extra_key,
|
extra_key=obj.extra_key,
|
||||||
need_wait_for_image=obj.need_wait_for_image,
|
need_wait_for_image=obj.need_wait_for_image,
|
||||||
num_items_assigned=obj.num_items_assigned,
|
num_items_assigned=obj.num_items_assigned,
|
||||||
embedding_ports=obj.embedding_ports,
|
|
||||||
)
|
)
|
||||||
elif isinstance(obj, EmbeddingReqInput):
|
elif isinstance(obj, EmbeddingReqInput):
|
||||||
tokenized_obj = TokenizedEmbeddingReqInput(
|
tokenized_obj = TokenizedEmbeddingReqInput(
|
||||||
|
|||||||
@@ -2173,12 +2173,6 @@ class ServerArgs:
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Cannot set --encoder-only and --disaggregation-mode prefill/decode together"
|
"Cannot set --encoder-only and --disaggregation-mode prefill/decode together"
|
||||||
)
|
)
|
||||||
if (
|
|
||||||
self.language_only
|
|
||||||
and self.encoder_transfer_backend == "zmq_to_scheduler"
|
|
||||||
and self.pp_size > 1
|
|
||||||
):
|
|
||||||
raise ValueError("zmq_to_scheduler not support pp_size > 1")
|
|
||||||
|
|
||||||
if self.language_only and len(self.encoder_urls) == 0:
|
if self.language_only and len(self.encoder_urls) == 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
Reference in New Issue
Block a user