[HiCache]: Support DeepSeek v32 cpu offloading (#17415)
Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
@@ -0,0 +1,130 @@
|
||||
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,
|
||||
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()
|
||||
Reference in New Issue
Block a user