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
parent a9ce1623cd
commit 9acb21ae27
19 changed files with 1910 additions and 68 deletions

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,