[NIXL][XPU] Use np.uint64 for pointer/length arrays in disaggregation KV transfer (#24188)
This commit is contained in:
@@ -402,6 +402,14 @@ class NixlKVManager(CommonKVManager):
|
||||
):
|
||||
"""Generic KV cache transfer supporting both MHA and MLA architectures.
|
||||
Used by both send_kvcache and maybe_send_extra."""
|
||||
# Convert pointer lists to np.uint64 arrays up front.
|
||||
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set, e.g.
|
||||
# 0xffff81ab54e01000). Casting here prevents overflow when these values
|
||||
# are later used in numpy arithmetic.
|
||||
src_data_ptrs = np.array(src_data_ptrs, dtype=np.uint64)
|
||||
dst_data_ptrs = np.array(dst_data_ptrs, dtype=np.uint64)
|
||||
item_lens = np.array(item_lens, dtype=np.uint64)
|
||||
|
||||
# group by indices
|
||||
prefill_kv_blocks, dst_kv_blocks = group_concurrent_contiguous(
|
||||
prefill_data_indices, dst_data_indices
|
||||
@@ -449,11 +457,11 @@ class NixlKVManager(CommonKVManager):
|
||||
|
||||
# Precompute block starts/lengths to reduce Python-level loops.
|
||||
prefill_starts = np.fromiter(
|
||||
(block[0] for block in prefill_kv_blocks), dtype=np.int64
|
||||
(block[0] for block in prefill_kv_blocks), dtype=np.uint64
|
||||
)
|
||||
dst_starts = np.fromiter((block[0] for block in dst_kv_blocks), dtype=np.int64)
|
||||
dst_starts = np.fromiter((block[0] for block in dst_kv_blocks), dtype=np.uint64)
|
||||
block_lens = np.fromiter(
|
||||
(len(block) for block in prefill_kv_blocks), dtype=np.int64
|
||||
(len(block) for block in prefill_kv_blocks), dtype=np.uint64
|
||||
)
|
||||
|
||||
for src_ptr, dst_ptr, item_len in layers_params:
|
||||
@@ -465,14 +473,14 @@ class NixlKVManager(CommonKVManager):
|
||||
|
||||
def make_req_array(addr_chunks, len_chunks, gpu):
|
||||
if not addr_chunks:
|
||||
return np.empty((0, 3), dtype=np.int64)
|
||||
flat_addrs = np.concatenate(addr_chunks)
|
||||
flat_lens = np.concatenate(len_chunks)
|
||||
return np.empty((0, 3), dtype=np.uint64)
|
||||
flat_addrs = np.concatenate(addr_chunks).astype(np.uint64, copy=False)
|
||||
flat_lens = np.concatenate(len_chunks).astype(np.uint64, copy=False)
|
||||
return np.column_stack(
|
||||
(
|
||||
flat_addrs,
|
||||
flat_lens,
|
||||
np.full_like(flat_addrs, gpu),
|
||||
np.full_like(flat_addrs, gpu, dtype=np.uint64),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -623,13 +631,13 @@ class NixlKVManager(CommonKVManager):
|
||||
|
||||
def make_req_array(addr_chunks, size, gpu):
|
||||
if not addr_chunks:
|
||||
return np.empty((0, 3), dtype=np.int64)
|
||||
flat_addrs = np.concatenate(addr_chunks)
|
||||
return np.empty((0, 3), dtype=np.uint64)
|
||||
flat_addrs = np.concatenate(addr_chunks).astype(np.uint64, copy=False)
|
||||
return np.column_stack(
|
||||
(
|
||||
flat_addrs,
|
||||
np.full_like(flat_addrs, size),
|
||||
np.full_like(flat_addrs, gpu),
|
||||
np.full_like(flat_addrs, size, dtype=np.uint64),
|
||||
np.full_like(flat_addrs, gpu, dtype=np.uint64),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user