[NIXL][XPU] Use np.uint64 for pointer/length arrays in disaggregation KV transfer (#24188)

This commit is contained in:
Jianhong Zhang
2026-05-06 10:09:03 +08:00
committed by GitHub
parent a965f886bf
commit c7019ff33d
3 changed files with 162 additions and 11 deletions
+19 -11
View File
@@ -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),
)
)