[EPD] Optimize the Mooncake backend (#22587)
Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user