Fix CP page filtering by request-local position (#28718)
This commit is contained in:
@@ -867,6 +867,7 @@ class CommonKVSender(BaseKVSender):
|
|||||||
self.kv_mgr,
|
self.kv_mgr,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
index_slice,
|
index_slice,
|
||||||
|
total_pages=self.num_kv_indices,
|
||||||
)
|
)
|
||||||
elif self.kv_mgr.is_dummy_cp_rank:
|
elif self.kv_mgr.is_dummy_cp_rank:
|
||||||
if not is_last_chunk:
|
if not is_last_chunk:
|
||||||
|
|||||||
@@ -511,6 +511,16 @@ def get_kv_class(
|
|||||||
raise ValueError(f"Unsupported transfer backend: {transfer_backend}")
|
raise ValueError(f"Unsupported transfer backend: {transfer_backend}")
|
||||||
|
|
||||||
|
|
||||||
|
def _get_cp_rank_page_bounds(
|
||||||
|
total_pages: int, cp_rank: int, cp_size: int
|
||||||
|
) -> Tuple[int, int]:
|
||||||
|
base = total_pages // cp_size
|
||||||
|
rem = total_pages % cp_size
|
||||||
|
local_start = cp_rank * base + min(cp_rank, rem)
|
||||||
|
n_pages = base + (1 if cp_rank < rem else 0)
|
||||||
|
return local_start, local_start + n_pages
|
||||||
|
|
||||||
|
|
||||||
def page_indices_to_cp_rank_page_indices(
|
def page_indices_to_cp_rank_page_indices(
|
||||||
page_indices: np.ndarray,
|
page_indices: np.ndarray,
|
||||||
total_pages: int,
|
total_pages: int,
|
||||||
@@ -558,37 +568,35 @@ def page_indices_to_cp_rank_page_indices(
|
|||||||
|
|
||||||
|
|
||||||
def filter_kv_indices_for_cp_rank(
|
def filter_kv_indices_for_cp_rank(
|
||||||
kv_mgr: CommonKVManager, kv_indices: np.ndarray, index_slice: slice
|
kv_mgr: CommonKVManager,
|
||||||
|
kv_indices: np.ndarray,
|
||||||
|
index_slice: slice,
|
||||||
|
total_pages: Optional[int] = None,
|
||||||
) -> Tuple[np.ndarray, slice]:
|
) -> Tuple[np.ndarray, slice]:
|
||||||
"""Filters kv_indices and index_slice for the current CP rank."""
|
"""Filters kv_indices and index_slice for the current CP rank."""
|
||||||
total_pages = len(kv_indices)
|
if total_pages is None:
|
||||||
|
total_pages = len(kv_indices)
|
||||||
cp_rank = kv_mgr.attn_cp_rank
|
cp_rank = kv_mgr.attn_cp_rank
|
||||||
cp_size = kv_mgr.attn_cp_size
|
cp_size = kv_mgr.attn_cp_size
|
||||||
|
|
||||||
rank_page_indices = page_indices_to_cp_rank_page_indices(
|
if cp_size <= 1:
|
||||||
page_indices=kv_indices,
|
return kv_indices, index_slice
|
||||||
total_pages=total_pages,
|
|
||||||
cp_rank=cp_rank,
|
|
||||||
cp_size=cp_size,
|
|
||||||
)
|
|
||||||
|
|
||||||
if rank_page_indices.size == 0:
|
rank_start, rank_end = _get_cp_rank_page_bounds(total_pages, cp_rank, cp_size)
|
||||||
|
chunk_start = index_slice.start if index_slice.start is not None else 0
|
||||||
|
chunk_end = index_slice.stop if index_slice.stop is not None else total_pages
|
||||||
|
first_pos = max(rank_start, chunk_start) - chunk_start
|
||||||
|
last_pos = min(rank_end, chunk_end) - chunk_start
|
||||||
|
|
||||||
|
if last_pos <= first_pos:
|
||||||
new_kv_indices = kv_indices[:0]
|
new_kv_indices = kv_indices[:0]
|
||||||
new_index_slice = slice(index_slice.start, index_slice.start)
|
new_index_slice = slice(chunk_start, chunk_start)
|
||||||
else:
|
else:
|
||||||
mask = np.isin(kv_indices, rank_page_indices)
|
new_kv_indices = kv_indices[first_pos:last_pos]
|
||||||
if not mask.any():
|
new_index_slice = slice(
|
||||||
new_kv_indices = kv_indices[:0]
|
chunk_start + first_pos,
|
||||||
new_index_slice = slice(index_slice.start, index_slice.start)
|
chunk_start + last_pos,
|
||||||
else:
|
)
|
||||||
first_pos = int(mask.argmax())
|
|
||||||
last_pos = len(mask) - int(mask[::-1].argmax())
|
|
||||||
|
|
||||||
new_kv_indices = kv_indices[first_pos:last_pos]
|
|
||||||
new_index_slice = slice(
|
|
||||||
index_slice.start + first_pos,
|
|
||||||
index_slice.start + last_pos,
|
|
||||||
)
|
|
||||||
return new_kv_indices, new_index_slice
|
return new_kv_indices, new_index_slice
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user