feat: support EPD disaggregation (#12263)
Co-authored-by: liusy58 <liusy58@linux.alibaba.com> Co-authored-by: ZhengWG <zwg0606@gmail.com> Co-authored-by: Nicholas <45984215+liusy58@users.noreply.github.com> Co-authored-by: Shangming Cai <csmthu@gmail.com> Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com>
This commit is contained in:
co-authored by
liusy58
ZhengWG
Nicholas
Shangming Cai
Yuhao Yang
parent
a9ce1623cd
commit
9acb21ae27
@@ -245,6 +245,10 @@ class GenerateReqInput(BaseReq, APIServingTimingMixin):
|
||||
# Whether to return entropy
|
||||
return_entropy: bool = False
|
||||
|
||||
need_wait_for_image: Optional[bool] = None
|
||||
num_items_assigned: Optional[List] = None
|
||||
embedding_ports: Optional[List] = None
|
||||
|
||||
def contains_mm_input(self) -> bool:
|
||||
return (
|
||||
has_valid_data(self.image_data)
|
||||
@@ -724,6 +728,10 @@ class TokenizedGenerateReqInput(BaseReq):
|
||||
# Whether to return entropy
|
||||
return_entropy: bool = False
|
||||
|
||||
need_wait_for_image: bool = False
|
||||
num_items_assigned: Optional[List] = None
|
||||
embedding_ports: Optional[List] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchTokenizedGenerateReqInput(BaseBatchReq):
|
||||
|
||||
@@ -395,13 +395,49 @@ def get_embedding_chunk(
|
||||
|
||||
def _get_precomputed_embedding(
|
||||
items: List[MultimodalDataItem],
|
||||
prefix_length: List[int],
|
||||
extend_length: List[int],
|
||||
items_offset_list: List[List[Tuple[int, int]]],
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
If all items have precomputed_embeddings, return their concatenation.
|
||||
If some but not all have precomputed_embeddings, raise NotImplementedError.
|
||||
If none have precomputed_embeddings, return None.
|
||||
"""
|
||||
precomputed_embeddings = [item.precomputed_embeddings for item in items]
|
||||
precomputed_embeddings = []
|
||||
for idx, item in enumerate(items):
|
||||
if item.precomputed_embeddings is None:
|
||||
precomputed_embeddings.append(None)
|
||||
continue
|
||||
seq_start_idx = prefix_length[idx]
|
||||
seq_end_idx = seq_start_idx + extend_length[idx] - 1
|
||||
prefix_embedding_length = []
|
||||
extend_embedding_length = []
|
||||
for mm_start_idx, mm_end_idx in items_offset_list[idx]:
|
||||
if mm_start_idx > seq_end_idx:
|
||||
break
|
||||
if seq_start_idx > mm_start_idx:
|
||||
prefix_embedding_length.append(
|
||||
min(seq_start_idx - mm_start_idx, mm_end_idx - mm_start_idx + 1)
|
||||
)
|
||||
if mm_end_idx >= seq_start_idx:
|
||||
extend_embedding_length.append(
|
||||
min(
|
||||
mm_end_idx - seq_start_idx + 1,
|
||||
seq_end_idx - mm_start_idx + 1,
|
||||
mm_end_idx - mm_start_idx + 1,
|
||||
seq_end_idx - seq_start_idx + 1,
|
||||
)
|
||||
)
|
||||
prefix_embedding_length = int(np.sum(prefix_embedding_length))
|
||||
extend_embedding_length = int(np.sum(extend_embedding_length))
|
||||
precomputed_embeddings.append(
|
||||
item.precomputed_embeddings[
|
||||
prefix_embedding_length : prefix_embedding_length
|
||||
+ extend_embedding_length
|
||||
]
|
||||
)
|
||||
|
||||
if any(feature is not None for feature in precomputed_embeddings):
|
||||
if not all(feature is not None for feature in precomputed_embeddings):
|
||||
raise NotImplementedError(
|
||||
@@ -527,7 +563,9 @@ def get_embedding_and_mask(
|
||||
- A boolean mask tensor indicating where these embeddings should be placed
|
||||
"""
|
||||
# 1. Get embedding
|
||||
embedding = _get_precomputed_embedding(embedding_items)
|
||||
embedding = _get_precomputed_embedding(
|
||||
embedding_items, prefix_length, extend_length, items_offset_list
|
||||
)
|
||||
if embedding is None:
|
||||
embedding = _get_chunked_prefill_embedding(
|
||||
data_embedding_func,
|
||||
|
||||
@@ -47,6 +47,7 @@ from sglang.srt.disaggregation.decode import (
|
||||
from sglang.srt.disaggregation.decode_kvcache_offload_manager import (
|
||||
DecodeKVCacheOffloadManager,
|
||||
)
|
||||
from sglang.srt.disaggregation.encode_receiver import MMReceiver
|
||||
from sglang.srt.disaggregation.prefill import (
|
||||
PrefillBootstrapQueue,
|
||||
SchedulerDisaggregationPrefillMixin,
|
||||
@@ -572,6 +573,17 @@ class Scheduler(
|
||||
# Init mlp sync flag
|
||||
self.require_mlp_sync = require_mlp_sync(server_args)
|
||||
|
||||
if (
|
||||
self.server_args.language_only
|
||||
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
|
||||
):
|
||||
self.mm_receiver = MMReceiver(
|
||||
server_args,
|
||||
hf_config=self.model_config.hf_config,
|
||||
tp_rank=self.tp_rank,
|
||||
pp_rank=self.pp_rank,
|
||||
)
|
||||
|
||||
# Init request dispatcher
|
||||
self._request_dispatcher = TypeBasedDispatcher(
|
||||
[
|
||||
@@ -1214,6 +1226,14 @@ class Scheduler(
|
||||
return recv_reqs
|
||||
|
||||
def process_input_requests(self, recv_reqs: List):
|
||||
|
||||
# Process MM requests under E disaggregation
|
||||
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:
|
||||
# If it is a health check generation request and there are running requests, ignore it.
|
||||
if is_health_check_generate_req(recv_req) and (
|
||||
|
||||
@@ -38,6 +38,7 @@ import zmq.asyncio
|
||||
from fastapi import BackgroundTasks
|
||||
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.disaggregation.encode_receiver import MMReceiver
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.lora.lora_registry import LoRARef, LoRARegistry
|
||||
@@ -287,6 +288,13 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
# Make sure that each request carries the tokenizer_ipc_name for response routing
|
||||
self.send_to_scheduler = SenderWrapper(port_args, send_to_scheduler)
|
||||
|
||||
# E Disaggregation
|
||||
if self.server_args.language_only:
|
||||
self.mm_receiver = MMReceiver(
|
||||
server_args,
|
||||
dtype=self.model_config.dtype,
|
||||
)
|
||||
|
||||
# Request states
|
||||
self._chosen_loop = None
|
||||
self.rid_to_state: Dict[str, ReqState] = {}
|
||||
@@ -411,6 +419,13 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
created_time = obj.received_time if obj.received_time else time.time()
|
||||
self.auto_create_handle_loop()
|
||||
obj.normalize_batch_and_arguments()
|
||||
if (
|
||||
self.server_args.language_only
|
||||
and isinstance(obj, GenerateReqInput)
|
||||
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
|
||||
and obj.contains_mm_input()
|
||||
):
|
||||
self.mm_receiver.send_encode_request(obj)
|
||||
|
||||
if self.enable_trace:
|
||||
self._trace_request_start(obj, created_time, request)
|
||||
@@ -615,13 +630,29 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
obj.image_data = [obj.image_data]
|
||||
if obj.audio_data is not None and not isinstance(obj.audio_data, list):
|
||||
obj.audio_data = [obj.audio_data]
|
||||
mm_inputs: Dict = await self.mm_data_processor.process(
|
||||
image_data=obj.image_data,
|
||||
audio_data=obj.audio_data,
|
||||
input_text_or_ids=(input_text or input_ids),
|
||||
request_obj=obj,
|
||||
max_req_input_len=self.max_req_input_len,
|
||||
)
|
||||
|
||||
mm_inputs = None
|
||||
|
||||
if (
|
||||
not self.server_args.language_only
|
||||
or self.server_args.encoder_transfer_backend
|
||||
in ["zmq_to_tokenizer", "mooncake"]
|
||||
):
|
||||
if self.server_args.language_only:
|
||||
mm_inputs = await self.mm_receiver.recv_mm_data(
|
||||
img_data=obj.image_data,
|
||||
mm_processor=self.mm_processor,
|
||||
prompt=(input_text or input_ids),
|
||||
)
|
||||
if mm_inputs is None:
|
||||
mm_inputs: Dict = await self.mm_data_processor.process(
|
||||
image_data=obj.image_data,
|
||||
audio_data=obj.audio_data,
|
||||
input_text_or_ids=(input_text or input_ids),
|
||||
request_obj=obj,
|
||||
max_req_input_len=self.max_req_input_len,
|
||||
)
|
||||
|
||||
if mm_inputs and "input_ids" in mm_inputs:
|
||||
input_ids = mm_inputs["input_ids"]
|
||||
else:
|
||||
@@ -801,6 +832,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
data_parallel_rank=obj.data_parallel_rank,
|
||||
priority=obj.priority,
|
||||
extra_key=obj.extra_key,
|
||||
need_wait_for_image=obj.need_wait_for_image,
|
||||
num_items_assigned=obj.num_items_assigned,
|
||||
embedding_ports=obj.embedding_ports,
|
||||
)
|
||||
elif isinstance(obj, EmbeddingReqInput):
|
||||
tokenized_obj = TokenizedEmbeddingReqInput(
|
||||
|
||||
Reference in New Issue
Block a user