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:
Tianyu Guo
2025-12-14 22:30:08 +08:00
committed by GitHub
co-authored by liusy58 ZhengWG Nicholas Shangming Cai Yuhao Yang
parent a9ce1623cd
commit 9acb21ae27
19 changed files with 1910 additions and 68 deletions
+8
View File
@@ -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):
+40 -2
View File
@@ -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,
+20
View File
@@ -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(