Files
sglang/test/registered/hicache/test_nsa_pool_host_unit.py

132 lines
4.6 KiB
Python

import unittest
import torch
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
from sglang.srt.mem_cache.memory_pool_host import (
ALLOC_MEMORY_FUNCS,
NSATokenToKVPoolHost,
alloc_with_pin_memory,
)
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=3, suite="stage-b-test-small-1-gpu")
class TestNSAHiCacheTransfer(unittest.TestCase):
def setUp(self):
if not torch.cuda.is_available():
self.skipTest("CUDA is required for NSA host transfer tests.")
if is_npu() or is_xpu():
self.skipTest("NSA host transfer tests only support CUDA/ROCm.")
if not (is_cuda() or is_hip()):
self.skipTest("CUDA/ROCm not available.")
@staticmethod
def _token_indices_for_pages(pages: torch.Tensor, page_size: int, device: str):
parts = [
torch.arange(
int(page_id) * page_size,
(int(page_id) + 1) * page_size,
device=device,
dtype=torch.int64,
)
for page_id in pages.tolist()
]
return torch.cat(parts, dim=0)
def _run_device_to_host_indexer_copy(self, io_backend: str):
page_size = 1 if is_hip() else 64
layer_num = 2
size = page_size * 4
device_pool = NSATokenToKVPool(
size=size,
page_size=page_size,
kv_lora_rank=128,
dtype=torch.bfloat16,
qk_rope_head_dim=32,
layer_num=layer_num,
device="cuda",
enable_memory_saver=False,
kv_cache_dim=576,
index_head_dim=128,
)
pin_memory = io_backend == "kernel"
original_alloc = ALLOC_MEMORY_FUNCS["cuda"]
if pin_memory:
ALLOC_MEMORY_FUNCS["cuda"] = alloc_with_pin_memory
try:
host_pool = NSATokenToKVPoolHost(
device_pool=device_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=page_size,
layout="layer_first",
pin_memory=pin_memory,
device="cpu",
)
finally:
ALLOC_MEMORY_FUNCS["cuda"] = original_alloc
for layer_id in range(layer_num):
buf = device_pool.index_k_with_scale_buffer[layer_id]
data = torch.arange(
buf.numel(), device=buf.device, dtype=torch.uint8
).view_as(buf)
buf.copy_((data + layer_id) % 256)
kv_buf = device_pool.kv_buffer[layer_id]
kv_data = torch.arange(
kv_buf.numel(), device=kv_buf.device, dtype=kv_buf.dtype
).view_as(kv_buf)
kv_buf.copy_(kv_data + layer_id)
device_pages = torch.tensor([1, 2, 3], device="cuda", dtype=torch.int64)
host_pages = torch.tensor(
[0, 1, 2],
device="cuda" if io_backend == "kernel" else "cpu",
dtype=torch.int64,
)
device_indices = self._token_indices_for_pages(
device_pages, page_size, device="cuda"
)
host_indices = self._token_indices_for_pages(
host_pages,
page_size,
device="cuda" if io_backend == "kernel" else "cpu",
)
host_pool.backup_from_device_all_layer(
device_pool, host_indices, device_indices, io_backend
)
for layer_id in range(layer_num):
for host_page, device_page in zip(
host_pages.tolist(), device_pages.tolist()
):
got = host_pool.index_k_with_scale_buffer[layer_id][host_page].cpu()
expected = device_pool.index_k_with_scale_buffer[layer_id][
device_page
].cpu()
self.assertTrue(torch.equal(got, expected))
host_start = host_page * page_size
device_start = device_page * page_size
got_kv = host_pool.kv_buffer[layer_id][
host_start : host_start + page_size
].cpu()
expected_kv = device_pool.kv_buffer[layer_id][
device_start : device_start + page_size
].cpu()
self.assertTrue(torch.equal(got_kv, expected_kv))
def test_device_to_host_indexer_kernel(self):
self._run_device_to_host_indexer_copy(io_backend="kernel")
def test_device_to_host_indexer_direct(self):
self._run_device_to_host_indexer_copy(io_backend="direct")
if __name__ == "__main__":
unittest.main()