[HiCache] refactor page_first_direct io kernel (#18113)
Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
@@ -317,6 +317,7 @@ def test_transfer_kv_pf_direct(
|
||||
torch.set_default_dtype(dtype)
|
||||
device = "cuda"
|
||||
torch.cuda.manual_seed(42)
|
||||
test_stream = torch.cuda.Stream()
|
||||
|
||||
num_layers = 4
|
||||
|
||||
@@ -356,13 +357,16 @@ def test_transfer_kv_pf_direct(
|
||||
dst_pool_direct = torch.zeros_like(dst_pool_ref)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
transfer_kv_all_layer_direct_lf_pf(
|
||||
src_pool_ptrs,
|
||||
[dst_pool_direct],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
page_size,
|
||||
)
|
||||
with torch.cuda.stream(test_stream):
|
||||
transfer_kv_all_layer_direct_lf_pf(
|
||||
src_pool_ptrs,
|
||||
[dst_pool_direct],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
page_size,
|
||||
)
|
||||
test_stream.synchronize()
|
||||
|
||||
for i in range(num_layers):
|
||||
ref_copy_with_indices_pf_direct(
|
||||
src_pool,
|
||||
@@ -393,13 +397,16 @@ def test_transfer_kv_pf_direct(
|
||||
dst_v_pool_direct = torch.zeros_like(dst_v_pool_ref)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
transfer_kv_all_layer_direct_lf_pf(
|
||||
src_k_pool_ptrs + src_v_pool_ptrs,
|
||||
[dst_k_pool_direct, dst_v_pool_direct],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
page_size,
|
||||
)
|
||||
with torch.cuda.stream(test_stream):
|
||||
transfer_kv_all_layer_direct_lf_pf(
|
||||
src_k_pool_ptrs + src_v_pool_ptrs,
|
||||
[dst_k_pool_direct, dst_v_pool_direct],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
page_size,
|
||||
)
|
||||
test_stream.synchronize()
|
||||
|
||||
for i in range(num_layers):
|
||||
ref_copy_with_indices_pf_direct(
|
||||
src_k_pool,
|
||||
@@ -435,14 +442,17 @@ def test_transfer_kv_pf_direct(
|
||||
dst_pool_direct_ptrs = [dst_pool_direct[i] for i in range(num_layers)]
|
||||
torch.cuda.synchronize()
|
||||
|
||||
transfer_kv_per_layer_direct_pf_lf(
|
||||
[src_pool],
|
||||
[dst_pool_direct_ptrs[layer_idx_to_test]],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
layer_idx_to_test,
|
||||
page_size,
|
||||
)
|
||||
with torch.cuda.stream(test_stream):
|
||||
transfer_kv_per_layer_direct_pf_lf(
|
||||
[src_pool],
|
||||
[dst_pool_direct_ptrs[layer_idx_to_test]],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
layer_idx_to_test,
|
||||
page_size,
|
||||
)
|
||||
test_stream.synchronize()
|
||||
|
||||
ref_copy_with_indices_pf_direct(
|
||||
src_pool,
|
||||
dst_pool_ref,
|
||||
@@ -473,17 +483,19 @@ def test_transfer_kv_pf_direct(
|
||||
dst_v_pool_direct_ptrs = [dst_v_pool_direct[i] for i in range(num_layers)]
|
||||
torch.cuda.synchronize()
|
||||
|
||||
transfer_kv_per_layer_direct_pf_lf(
|
||||
[src_k_pool, src_v_pool],
|
||||
[
|
||||
dst_k_pool_direct_ptrs[layer_idx_to_test],
|
||||
dst_v_pool_direct_ptrs[layer_idx_to_test],
|
||||
],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
layer_idx_to_test,
|
||||
page_size,
|
||||
)
|
||||
with torch.cuda.stream(test_stream):
|
||||
transfer_kv_per_layer_direct_pf_lf(
|
||||
[src_k_pool, src_v_pool],
|
||||
[
|
||||
dst_k_pool_direct_ptrs[layer_idx_to_test],
|
||||
dst_v_pool_direct_ptrs[layer_idx_to_test],
|
||||
],
|
||||
src_indices_host,
|
||||
dst_indices_host,
|
||||
layer_idx_to_test,
|
||||
page_size,
|
||||
)
|
||||
test_stream.synchronize()
|
||||
|
||||
ref_copy_with_indices_pf_direct(
|
||||
src_k_pool,
|
||||
|
||||
Reference in New Issue
Block a user