[HiCache]Page head layout IO kernel (#11615)
This commit is contained in:
@@ -3,11 +3,13 @@ import torch
|
||||
from sgl_kernel.kvcacheio import (
|
||||
transfer_kv_all_layer,
|
||||
transfer_kv_all_layer_direct_lf_pf,
|
||||
transfer_kv_all_layer_lf_ph,
|
||||
transfer_kv_all_layer_mla,
|
||||
transfer_kv_direct,
|
||||
transfer_kv_per_layer,
|
||||
transfer_kv_per_layer_direct_pf_lf,
|
||||
transfer_kv_per_layer_mla,
|
||||
transfer_kv_per_layer_ph_lf,
|
||||
)
|
||||
|
||||
|
||||
@@ -30,6 +32,32 @@ def ref_copy_with_indices_pf_direct(
|
||||
][layer_id].to(dst_pool.device)
|
||||
|
||||
|
||||
def ref_copy_with_indices_page_head(
|
||||
src_pool,
|
||||
dst_pool,
|
||||
src_indices,
|
||||
dst_indices,
|
||||
page_size,
|
||||
layer_id,
|
||||
head_num,
|
||||
lf_to_ph=False,
|
||||
):
|
||||
if lf_to_ph:
|
||||
for head_id in range(head_num):
|
||||
for i in range(0, len(src_indices)):
|
||||
dst_pool[dst_indices[i] // page_size][head_id][
|
||||
dst_indices[i] % page_size
|
||||
][layer_id] = src_pool[layer_id][src_indices[i]][head_id].to(
|
||||
dst_pool.device
|
||||
)
|
||||
else:
|
||||
for head_id in range(head_num):
|
||||
for i in range(0, len(src_indices)):
|
||||
dst_pool[layer_id][dst_indices[i]][head_id] = src_pool[
|
||||
src_indices[i] // page_size
|
||||
][head_id][src_indices[i] % page_size][layer_id].to(dst_pool.device)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
|
||||
@pytest.mark.parametrize("num_items_to_transfer", [1, 128, 1024])
|
||||
@pytest.mark.parametrize("page_size", [1, 16, 64])
|
||||
@@ -481,5 +509,182 @@ def test_transfer_kv_pf_direct(
|
||||
torch.set_default_dtype(original_dtype)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
|
||||
@pytest.mark.parametrize("num_items_to_transfer", [256, 1024])
|
||||
@pytest.mark.parametrize("page_size", [16, 64, 128])
|
||||
@pytest.mark.parametrize("item_size", [1024])
|
||||
@pytest.mark.parametrize("head_num", [8, 16])
|
||||
@pytest.mark.parametrize("total_items_in_pool", [4096])
|
||||
@pytest.mark.parametrize("lf_to_ph", [False, True])
|
||||
def test_transfer_kv_page_head(
|
||||
dtype: torch.dtype,
|
||||
num_items_to_transfer: int,
|
||||
page_size: int,
|
||||
item_size: int,
|
||||
head_num: int,
|
||||
total_items_in_pool: int,
|
||||
lf_to_ph: bool,
|
||||
):
|
||||
original_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(dtype)
|
||||
device = "cuda"
|
||||
torch.cuda.manual_seed(42)
|
||||
|
||||
num_layers = 4
|
||||
|
||||
total_pages_in_pool = total_items_in_pool // page_size
|
||||
num_pages_to_transfer = num_items_to_transfer // page_size
|
||||
if num_pages_to_transfer == 0:
|
||||
torch.set_default_dtype(original_dtype)
|
||||
return
|
||||
|
||||
assert item_size % head_num == 0
|
||||
head_dim = item_size // head_num
|
||||
|
||||
page_indices = torch.randperm(total_pages_in_pool, dtype=torch.int64)
|
||||
src_indices_host = torch.cat(
|
||||
[
|
||||
torch.arange(p * page_size, (p + 1) * page_size)
|
||||
for p in page_indices[:num_pages_to_transfer]
|
||||
]
|
||||
)
|
||||
src_indices_device = src_indices_host.to(device)
|
||||
dst_indices_host = torch.cat(
|
||||
[
|
||||
torch.arange(p * page_size, (p + 1) * page_size)
|
||||
for p in page_indices[num_pages_to_transfer : 2 * num_pages_to_transfer]
|
||||
]
|
||||
)
|
||||
dst_indices_device = dst_indices_host.to(device)
|
||||
|
||||
# We will test the per-layer function on the first layer (index 0) of the pool.
|
||||
layer_idx_to_test = 0
|
||||
|
||||
if lf_to_ph:
|
||||
src_k_pool = torch.randn(
|
||||
num_layers, total_items_in_pool, head_num, head_dim
|
||||
).to(device)
|
||||
src_v_pool = torch.randn(
|
||||
num_layers, total_items_in_pool, head_num, head_dim
|
||||
).to(device)
|
||||
src_k_pool_ptrs = [src_k_pool[i] for i in range(num_layers)]
|
||||
src_k_pool_ptrs = torch.tensor(
|
||||
[x.data_ptr() for x in src_k_pool_ptrs],
|
||||
dtype=torch.uint64,
|
||||
device=device,
|
||||
)
|
||||
src_v_pool_ptrs = [src_v_pool[i] for i in range(num_layers)]
|
||||
src_v_pool_ptrs = torch.tensor(
|
||||
[x.data_ptr() for x in src_v_pool_ptrs],
|
||||
dtype=torch.uint64,
|
||||
device=device,
|
||||
)
|
||||
|
||||
dst_k_pool_ref = torch.zeros(
|
||||
total_pages_in_pool, head_num, page_size, num_layers, head_dim
|
||||
).pin_memory()
|
||||
dst_v_pool_ref = torch.zeros_like(dst_k_pool_ref).pin_memory()
|
||||
|
||||
dst_k_pool_kernel = torch.zeros_like(dst_k_pool_ref).pin_memory()
|
||||
dst_v_pool_kernel = torch.zeros_like(dst_v_pool_ref).pin_memory()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
transfer_kv_all_layer_lf_ph(
|
||||
src_k_pool_ptrs,
|
||||
dst_k_pool_kernel,
|
||||
src_v_pool_ptrs,
|
||||
dst_v_pool_kernel,
|
||||
src_indices_device,
|
||||
dst_indices_device,
|
||||
item_size * dtype.itemsize,
|
||||
item_size * num_layers * dtype.itemsize,
|
||||
num_layers,
|
||||
page_size,
|
||||
head_num,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
for i in range(num_layers):
|
||||
ref_copy_with_indices_page_head(
|
||||
src_k_pool,
|
||||
dst_k_pool_ref,
|
||||
src_indices_device,
|
||||
dst_indices_host,
|
||||
page_size,
|
||||
i,
|
||||
head_num,
|
||||
lf_to_ph=True,
|
||||
)
|
||||
ref_copy_with_indices_page_head(
|
||||
src_v_pool,
|
||||
dst_v_pool_ref,
|
||||
src_indices_device,
|
||||
dst_indices_host,
|
||||
page_size,
|
||||
i,
|
||||
head_num,
|
||||
lf_to_ph=True,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
torch.testing.assert_close(dst_k_pool_kernel, dst_k_pool_ref)
|
||||
torch.testing.assert_close(dst_v_pool_kernel, dst_v_pool_ref)
|
||||
else:
|
||||
src_k_pool = torch.randn(
|
||||
total_pages_in_pool, head_num, page_size, num_layers, head_dim
|
||||
).pin_memory()
|
||||
src_v_pool = torch.randn(
|
||||
total_pages_in_pool, head_num, page_size, num_layers, head_dim
|
||||
).pin_memory()
|
||||
|
||||
dst_k_pool_ref = torch.zeros(
|
||||
num_layers, total_items_in_pool, head_num, head_dim
|
||||
).to(device)
|
||||
dst_v_pool_ref = torch.zeros_like(dst_k_pool_ref)
|
||||
dst_k_pool_kernel = torch.zeros_like(dst_k_pool_ref)
|
||||
dst_v_pool_kernel = torch.zeros_like(dst_v_pool_ref)
|
||||
dst_k_pool_kernel_ptrs = [dst_k_pool_kernel[i] for i in range(num_layers)]
|
||||
dst_v_pool_kernel_ptrs = [dst_v_pool_kernel[i] for i in range(num_layers)]
|
||||
torch.cuda.synchronize()
|
||||
|
||||
transfer_kv_per_layer_ph_lf(
|
||||
src_k_pool,
|
||||
dst_k_pool_kernel_ptrs[layer_idx_to_test],
|
||||
src_v_pool,
|
||||
dst_v_pool_kernel_ptrs[layer_idx_to_test],
|
||||
src_indices_device,
|
||||
dst_indices_device,
|
||||
layer_idx_to_test,
|
||||
item_size * dtype.itemsize,
|
||||
item_size * num_layers * dtype.itemsize,
|
||||
page_size,
|
||||
head_num,
|
||||
)
|
||||
|
||||
ref_copy_with_indices_page_head(
|
||||
src_k_pool,
|
||||
dst_k_pool_ref,
|
||||
src_indices_host,
|
||||
dst_indices_device,
|
||||
page_size,
|
||||
layer_idx_to_test,
|
||||
head_num,
|
||||
lf_to_ph=False,
|
||||
)
|
||||
ref_copy_with_indices_page_head(
|
||||
src_v_pool,
|
||||
dst_v_pool_ref,
|
||||
src_indices_host,
|
||||
dst_indices_device,
|
||||
page_size,
|
||||
layer_idx_to_test,
|
||||
head_num,
|
||||
lf_to_ph=False,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
torch.testing.assert_close(dst_k_pool_kernel, dst_k_pool_ref)
|
||||
torch.testing.assert_close(dst_v_pool_kernel, dst_v_pool_ref)
|
||||
torch.set_default_dtype(original_dtype)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
||||
Reference in New Issue
Block a user