[XPU][NIXL] Use uint64 for XPU address arithmetic in prep handle builders (#27415)
This commit is contained in:
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user