[EPD] Optimize the Mooncake backend (#22587)

Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
LucQueen
2026-05-29 10:42:24 +08:00
committed by GitHub
co-authored by ZhengWG
parent 569ee93357
commit 36d0a6e08e
11 changed files with 1516 additions and 159 deletions
+2
View File
@@ -257,6 +257,7 @@ class GenerateReqInput(BaseReq):
# For EPD-disaggregated inference
need_wait_for_mm_inputs: Optional[bool] = None
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
# Multimodal tiling controls (extensions)
max_dynamic_patch: Optional[int] = None
@@ -800,6 +801,7 @@ class TokenizedGenerateReqInput(BaseReq):
need_wait_for_mm_inputs: bool = False
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
+15 -1
View File
@@ -447,6 +447,16 @@ def _get_precomputed_embedding(
raise NotImplementedError(
"MM inputs where only some items are precomputed."
)
# Normalize device across chunks before concat.
target_device = next(
(t.device for t in precomputed_embeddings if t.is_cuda),
precomputed_embeddings[0].device,
)
precomputed_embeddings = [
t if t.device == target_device else t.to(target_device, non_blocking=True)
for t in precomputed_embeddings
]
result = torch.concat(precomputed_embeddings)
# some models embedding is 3-dim, reshape it to 2-dim (similar to get_embedding_chunk)
result = result.reshape(-1, result.shape[-1])
@@ -1077,7 +1087,11 @@ def offload_mm_features_to_cpu(mm_inputs_list: List[MultimodalInputs]):
item.feature = item.feature.to("cpu", non_blocking=True)
if language_only:
pe = item.precomputed_embeddings
if isinstance(pe, torch.Tensor) and pe.is_cuda:
if (
isinstance(pe, torch.Tensor)
and pe.is_cuda
and not getattr(item, "_keep_device_embedding", False)
):
item.precomputed_embeddings = pe.to("cpu", non_blocking=True)
+3 -1
View File
@@ -1108,10 +1108,12 @@ class Scheduler(
# Init mm receiver for EPD disaggregation mode
if (
self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
):
self.mm_receiver = create_mm_receiver(
self.server_args,
dtype=self.model_config.dtype,
hf_config=self.model_config.hf_config,
pp_rank=self.ps.pp_rank,
tp_rank=self.ps.tp_rank,
@@ -189,7 +189,8 @@ class SchedulerRequestReceiver:
if (
self.ps.pp_rank == 0
and self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
):
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
for req, error_msg, error_code in abort_reqs:
@@ -792,8 +792,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if (
not self.server_args.language_only
or self.server_args.encoder_transfer_backend
in ["zmq_to_tokenizer", "mooncake"]
or self.server_args.encoder_transfer_backend == "zmq_to_tokenizer"
):
if self.server_args.language_only:
mm_inputs = await self.mm_receiver.recv_mm_data(
@@ -817,10 +816,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
elif (
self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
and not obj.need_wait_for_mm_inputs
):
# In language_only mode with zmq_to_scheduler, if we didn't dispatch
# In language_only mode with zmq_to_scheduler/mooncake, if we didn't dispatch
# to encoder (e.g., only one image), process locally like non-language_only mode
mm_inputs = await self.mm_processor.process_mm_data_async(
image_data=obj.image_data,
@@ -1067,6 +1067,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
num_items_assigned=obj.num_items_assigned,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
mm_data_mooncake=obj.mm_data_mooncake,
)
elif isinstance(obj, EmbeddingReqInput):
# Resolve unresolved embed overrides now that input_ids are available
@@ -2712,7 +2713,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# This flag will be used in _tokenize_one_request to determine processing path
if should_dispatch:
obj.need_wait_for_mm_inputs = True
if self.server_args.encoder_transfer_backend == "zmq_to_scheduler":
if self.server_args.encoder_transfer_backend in [
"zmq_to_scheduler",
"mooncake",
]:
self.mm_receiver.send_encode_request(obj)
else:
obj.need_wait_for_mm_inputs = False