[XPU][NIXL] Use uint64 for XPU address arithmetic in prep handle builders (#27415)

This commit is contained in:
Jianhong Zhang
2026-06-09 09:18:36 +08:00
committed by GitHub
parent 317fc6a9dd
commit ab70153b62
+23 -21
View File
@@ -521,17 +521,20 @@ class NixlKVManager(CommonKVManager):
Uses prefill's kv_item_lens as stride; requires equal per-slot byte size (equal-TP or MLA). Uses prefill's kv_item_lens as stride; requires equal per-slot byte size (equal-TP or MLA).
""" """
arrays = [] arrays = []
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set).
# Convert once at entry; all downstream arithmetic stays in uint64.
kv_ptrs_u64 = np.array(kv_ptrs, dtype=np.uint64)
for base_ptr, item_len, data_len in zip( for base_ptr, item_len, data_len in zip(
kv_ptrs, self.kv_args.kv_item_lens, self.kv_args.kv_data_lens kv_ptrs_u64, self.kv_args.kv_item_lens, self.kv_args.kv_data_lens
): ):
n = num_slots if num_slots is not None else (data_len // item_len) n = num_slots if num_slots is not None else (data_len // item_len)
addrs = np.arange(n, dtype=np.int64) * item_len + base_ptr addrs = np.arange(n, dtype=np.uint64) * np.uint64(item_len) + base_ptr
arrays.append( arrays.append(
np.column_stack( np.column_stack(
[ [
addrs, addrs,
np.full(n, item_len, dtype=np.int64), np.full(n, item_len, dtype=np.uint64),
np.full(n, gpu_id, dtype=np.int64), np.full(n, gpu_id, dtype=np.uint64),
] ]
) )
) )
@@ -604,25 +607,24 @@ class NixlKVManager(CommonKVManager):
num_ptr_pairs = len(src_ptrs) num_ptr_pairs = len(src_ptrs)
num_slots = self.kv_args.kv_data_lens[0] // src_kv_item_len num_slots = self.kv_args.kv_data_lens[0] // src_kv_item_len
slots = np.arange(num_slots, dtype=np.int64) slots = np.arange(num_slots, dtype=np.uint64)
tokens = np.arange(page_size, dtype=np.int64) # reused in dst dlist below tokens = np.arange(page_size, dtype=np.uint64) # reused in dst dlist below
groups = np.arange(num_groups, dtype=np.int64) groups = np.arange(num_groups, dtype=np.uint64)
# Src dlist built once and shared. # Src dlist built once and shared.
if self.prep_handle_slice_src is None: if self.prep_handle_slice_src is None:
# (ptr, slot, token, group) → ravel; groups interleaved per token. src_ptrs_arr = np.array(src_ptrs, dtype=np.uint64)
src_ptrs_arr = np.array(src_ptrs, dtype=np.int64)
addrs = ( addrs = (
src_ptrs_arr[:, None, None, None] src_ptrs_arr[:, None, None, None]
+ slots[None, :, None, None] * src_kv_item_len + slots[None, :, None, None] * np.uint64(src_kv_item_len)
+ tokens[None, None, :, None] * bytes_per_token_src + tokens[None, None, :, None] * np.uint64(bytes_per_token_src)
+ groups[None, None, None, :] * bytes_per_token_to_send + groups[None, None, None, :] * np.uint64(bytes_per_token_to_send)
).ravel() ).ravel()
src_array = np.column_stack( src_array = np.column_stack(
[ [
addrs, addrs,
np.full(len(addrs), bytes_per_token_to_send, dtype=np.int64), np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
np.full(len(addrs), self.kv_args.gpu_id, dtype=np.int64), np.full(len(addrs), self.kv_args.gpu_id, dtype=np.uint64),
] ]
) )
src_handle = self.agent.prep_xfer_dlist("", src_array, "VRAM") src_handle = self.agent.prep_xfer_dlist("", src_array, "VRAM")
@@ -642,20 +644,20 @@ class NixlKVManager(CommonKVManager):
if decode_kv_args.dst_num_slots is not None if decode_kv_args.dst_num_slots is not None
else num_slots else num_slots
) )
dst_slots = np.arange(num_slots_dst, dtype=np.int64) dst_slots = np.arange(num_slots_dst, dtype=np.uint64)
# (ptr, slot, token) → ravel. # (ptr, slot, token) → ravel.
dst_ptrs_arr = np.array(dst_ptrs, dtype=np.int64) dst_ptrs_arr = np.array(dst_ptrs, dtype=np.uint64)
addrs = ( addrs = (
dst_ptrs_arr[:, None, None] dst_ptrs_arr[:, None, None]
+ dst_slots[None, :, None] * dst_kv_item_len + dst_slots[None, :, None] * np.uint64(dst_kv_item_len)
+ tokens[None, None, :] * bytes_per_token_dst + tokens[None, None, :] * np.uint64(bytes_per_token_dst)
+ dst_head_offset + np.uint64(dst_head_offset)
).ravel() ).ravel()
dst_array = np.column_stack( dst_array = np.column_stack(
[ [
addrs, addrs,
np.full(len(addrs), bytes_per_token_to_send, dtype=np.int64), np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
np.full(len(addrs), decode_kv_args.gpu_id, dtype=np.int64), np.full(len(addrs), decode_kv_args.gpu_id, dtype=np.uint64),
] ]
) )
dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, "VRAM") dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, "VRAM")