[HiCache]: Support DeepSeek v32 cpu offloading (#17415)
Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
@@ -21,10 +21,15 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
MatchPrefixParams,
|
MatchPrefixParams,
|
||||||
MatchResult,
|
MatchResult,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, MLATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
|
MHATokenToKVPool,
|
||||||
|
MLATokenToKVPool,
|
||||||
|
NSATokenToKVPool,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
MHATokenToKVPoolHost,
|
MHATokenToKVPoolHost,
|
||||||
MLATokenToKVPoolHost,
|
MLATokenToKVPoolHost,
|
||||||
|
NSATokenToKVPoolHost,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import (
|
from sglang.srt.mem_cache.radix_cache import (
|
||||||
RadixCache,
|
RadixCache,
|
||||||
@@ -70,6 +75,15 @@ class HiRadixCache(RadixCache):
|
|||||||
server_args.hicache_mem_layout,
|
server_args.hicache_mem_layout,
|
||||||
allocator_type=server_args.hicache_storage_backend,
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
)
|
)
|
||||||
|
elif isinstance(self.kv_cache, NSATokenToKVPool):
|
||||||
|
self.token_to_kv_pool_host = NSATokenToKVPoolHost(
|
||||||
|
self.kv_cache,
|
||||||
|
server_args.hicache_ratio,
|
||||||
|
server_args.hicache_size,
|
||||||
|
self.page_size,
|
||||||
|
server_args.hicache_mem_layout,
|
||||||
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
|
)
|
||||||
elif isinstance(self.kv_cache, MLATokenToKVPool):
|
elif isinstance(self.kv_cache, MLATokenToKVPool):
|
||||||
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
|
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
|
||||||
self.kv_cache,
|
self.kv_cache,
|
||||||
|
|||||||
@@ -15,7 +15,12 @@ from sglang.jit_kernel.hicache import (
|
|||||||
from sglang.jit_kernel.hicache import (
|
from sglang.jit_kernel.hicache import (
|
||||||
transfer_hicache_one_layer as jit_transfer_hicache_one_layer,
|
transfer_hicache_one_layer as jit_transfer_hicache_one_layer,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import KVCache, MHATokenToKVPool, MLATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
|
KVCache,
|
||||||
|
MHATokenToKVPool,
|
||||||
|
MLATokenToKVPool,
|
||||||
|
NSATokenToKVPool,
|
||||||
|
)
|
||||||
from sglang.srt.utils import is_cuda, is_npu, is_xpu
|
from sglang.srt.utils import is_cuda, is_npu, is_xpu
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
@@ -689,7 +694,9 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
pin_memory: bool = True,
|
pin_memory: bool = True,
|
||||||
device: str = "cpu",
|
device: str = "cpu",
|
||||||
allocator_type: str = "default",
|
allocator_type: str = "default",
|
||||||
|
override_kv_cache_dim: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
self.override_kv_cache_dim = override_kv_cache_dim
|
||||||
super().__init__(
|
super().__init__(
|
||||||
device_pool,
|
device_pool,
|
||||||
host_to_device_ratio,
|
host_to_device_ratio,
|
||||||
@@ -711,13 +718,10 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
self.kv_lora_rank = self.device_pool.kv_lora_rank
|
self.kv_lora_rank = self.device_pool.kv_lora_rank
|
||||||
self.qk_rope_head_dim = self.device_pool.qk_rope_head_dim
|
self.qk_rope_head_dim = self.device_pool.qk_rope_head_dim
|
||||||
self.layer_num = self.device_pool.layer_num
|
self.layer_num = self.device_pool.layer_num
|
||||||
|
self.kv_cache_dim = self.override_kv_cache_dim or (
|
||||||
return (
|
self.kv_lora_rank + self.qk_rope_head_dim
|
||||||
(self.kv_lora_rank + self.qk_rope_head_dim)
|
|
||||||
* 1
|
|
||||||
* self.dtype.itemsize
|
|
||||||
* self.layer_num
|
|
||||||
)
|
)
|
||||||
|
return self.kv_cache_dim * self.dtype.itemsize * self.layer_num
|
||||||
|
|
||||||
def get_ksize_per_token(self):
|
def get_ksize_per_token(self):
|
||||||
return self.get_size_per_token()
|
return self.get_size_per_token()
|
||||||
@@ -728,14 +732,14 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
self.layer_num,
|
self.layer_num,
|
||||||
self.size,
|
self.size,
|
||||||
1,
|
1,
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
dims = (
|
dims = (
|
||||||
self.size,
|
self.size,
|
||||||
self.layer_num,
|
self.layer_num,
|
||||||
1,
|
1,
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
dims = (
|
dims = (
|
||||||
@@ -743,7 +747,7 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
self.layer_num,
|
self.layer_num,
|
||||||
self.page_size,
|
self.page_size,
|
||||||
1,
|
1,
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
# Ascend-specific: Aligns with NPUMLATokenToKVPool layout
|
# Ascend-specific: Aligns with NPUMLATokenToKVPool layout
|
||||||
# Separately allocate k_buffer and v_buffer for easier data transfer.
|
# Separately allocate k_buffer and v_buffer for easier data transfer.
|
||||||
@@ -783,9 +787,7 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
return self.k_buffer
|
return self.k_buffer
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
self.token_stride_size = (
|
self.token_stride_size = self.kv_cache_dim * self.dtype.itemsize
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim
|
|
||||||
) * self.dtype.itemsize
|
|
||||||
self.layout_dim = self.token_stride_size * self.layer_num
|
self.layout_dim = self.token_stride_size * self.layer_num
|
||||||
|
|
||||||
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
|
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
|
||||||
@@ -946,7 +948,7 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
self.layer_num,
|
self.layer_num,
|
||||||
self.page_size,
|
self.page_size,
|
||||||
1,
|
1,
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
self.kv_cache_dim,
|
||||||
),
|
),
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
@@ -959,14 +961,14 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
self.layer_num,
|
self.layer_num,
|
||||||
self.page_size,
|
self.page_size,
|
||||||
1,
|
1,
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
self.kv_buffer[index : index + self.page_size, :, :, :] = data_page.reshape(
|
self.kv_buffer[index : index + self.page_size, :, :, :] = data_page.reshape(
|
||||||
self.page_size,
|
self.page_size,
|
||||||
self.layer_num,
|
self.layer_num,
|
||||||
1,
|
1,
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
real_index = index // self.page_size
|
real_index = index // self.page_size
|
||||||
@@ -975,7 +977,7 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
self.layer_num,
|
self.layer_num,
|
||||||
self.page_size,
|
self.page_size,
|
||||||
1,
|
1,
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim,
|
self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
@@ -993,20 +995,11 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
for layer_id in range(self.layer_num):
|
for layer_id in range(self.layer_num):
|
||||||
k_ptr = (
|
k_ptr = (
|
||||||
kv_buffer_data_ptr
|
kv_buffer_data_ptr
|
||||||
+ indices[index]
|
+ indices[index] * self.kv_cache_dim * self.dtype.itemsize
|
||||||
* (self.kv_lora_rank + self.qk_rope_head_dim)
|
+ layer_id * self.size * self.kv_cache_dim * self.dtype.itemsize
|
||||||
* self.dtype.itemsize
|
|
||||||
+ layer_id
|
|
||||||
* self.size
|
|
||||||
* (self.kv_lora_rank + self.qk_rope_head_dim)
|
|
||||||
* self.dtype.itemsize
|
|
||||||
)
|
)
|
||||||
ptr_list.append(k_ptr)
|
ptr_list.append(k_ptr)
|
||||||
element_size = (
|
element_size = self.dtype.itemsize * self.page_size * self.kv_cache_dim
|
||||||
self.dtype.itemsize
|
|
||||||
* self.page_size
|
|
||||||
* (self.kv_lora_rank + self.qk_rope_head_dim)
|
|
||||||
)
|
|
||||||
element_size_list = [element_size] * len(ptr_list)
|
element_size_list = [element_size] * len(ptr_list)
|
||||||
elif self.layout in ["page_first", "page_first_direct"]:
|
elif self.layout in ["page_first", "page_first_direct"]:
|
||||||
for index in range(0, len(indices), self.page_size):
|
for index in range(0, len(indices), self.page_size):
|
||||||
@@ -1014,7 +1007,7 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
kv_buffer_data_ptr
|
kv_buffer_data_ptr
|
||||||
+ indices[index]
|
+ indices[index]
|
||||||
* self.layer_num
|
* self.layer_num
|
||||||
* (self.kv_lora_rank + self.qk_rope_head_dim)
|
* self.kv_cache_dim
|
||||||
* self.dtype.itemsize
|
* self.dtype.itemsize
|
||||||
)
|
)
|
||||||
ptr_list.append(k_ptr)
|
ptr_list.append(k_ptr)
|
||||||
@@ -1022,9 +1015,174 @@ class MLATokenToKVPoolHost(HostKVCache):
|
|||||||
self.layer_num
|
self.layer_num
|
||||||
* self.dtype.itemsize
|
* self.dtype.itemsize
|
||||||
* self.page_size
|
* self.page_size
|
||||||
* (self.kv_lora_rank + self.qk_rope_head_dim)
|
* self.kv_cache_dim
|
||||||
)
|
)
|
||||||
element_size_list = [element_size] * len(ptr_list)
|
element_size_list = [element_size] * len(ptr_list)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
return ptr_list, element_size_list
|
return ptr_list, element_size_list
|
||||||
|
|
||||||
|
|
||||||
|
class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
||||||
|
device_pool: NSATokenToKVPool
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
device_pool: NSATokenToKVPool,
|
||||||
|
host_to_device_ratio: float,
|
||||||
|
host_size: int,
|
||||||
|
page_size: int,
|
||||||
|
layout: str,
|
||||||
|
pin_memory: bool = True,
|
||||||
|
device: str = "cpu",
|
||||||
|
allocator_type: str = "default",
|
||||||
|
):
|
||||||
|
# Initialize indexer metadata before HostKVCache.__init__ calls get_size_per_token.
|
||||||
|
self.index_head_dim = device_pool.index_head_dim
|
||||||
|
self.indexer_quant_block_size = device_pool.quant_block_size
|
||||||
|
self.indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype
|
||||||
|
self.indexer_size_per_token = (
|
||||||
|
self.index_head_dim
|
||||||
|
+ self.index_head_dim // self.indexer_quant_block_size * 4
|
||||||
|
)
|
||||||
|
super().__init__(
|
||||||
|
device_pool,
|
||||||
|
host_to_device_ratio,
|
||||||
|
host_size,
|
||||||
|
page_size,
|
||||||
|
layout,
|
||||||
|
pin_memory,
|
||||||
|
device,
|
||||||
|
allocator_type,
|
||||||
|
override_kv_cache_dim=device_pool.kv_cache_dim,
|
||||||
|
)
|
||||||
|
self.indexer_page_stride_size = (
|
||||||
|
self.indexer_size_per_token * self.page_size * self.indexer_dtype.itemsize
|
||||||
|
)
|
||||||
|
self.indexer_page_num = (self.size + self.page_size + 1) // self.page_size
|
||||||
|
self._init_indexer_buffers()
|
||||||
|
logger.info(
|
||||||
|
f"NSATokenToKVPoolHost initialized with indexer page stride size: {self.indexer_page_stride_size}, page num: {self.indexer_page_num}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_size_per_token(self):
|
||||||
|
base = super().get_size_per_token()
|
||||||
|
return (
|
||||||
|
base
|
||||||
|
+ self.indexer_size_per_token * self.layer_num * self.indexer_dtype.itemsize
|
||||||
|
)
|
||||||
|
|
||||||
|
def _init_indexer_buffers(self):
|
||||||
|
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
|
||||||
|
self.index_k_with_scale_buffer = [
|
||||||
|
alloc_func(
|
||||||
|
(self.indexer_page_num, self.indexer_page_stride_size),
|
||||||
|
dtype=self.indexer_dtype,
|
||||||
|
device=self.device,
|
||||||
|
pin_memory=self.pin_memory,
|
||||||
|
allocator=self.allocator,
|
||||||
|
)
|
||||||
|
for _ in range(self.layer_num)
|
||||||
|
]
|
||||||
|
self.index_k_data_refs = [
|
||||||
|
self.index_k_with_scale_buffer[i] for i in range(self.layer_num)
|
||||||
|
]
|
||||||
|
self.index_k_data_ptrs = torch.tensor(
|
||||||
|
[x.data_ptr() for x in self.index_k_data_refs],
|
||||||
|
dtype=torch.uint64,
|
||||||
|
device=self.device_pool.device,
|
||||||
|
)
|
||||||
|
self.index_k_device_ptrs = torch.tensor(
|
||||||
|
[x.data_ptr() for x in self.device_pool.index_k_with_scale_buffer],
|
||||||
|
dtype=torch.uint64,
|
||||||
|
device=self.device_pool.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _get_indexer_page_indices(self, host_indices, device_indices):
|
||||||
|
if host_indices.numel() == 0:
|
||||||
|
return host_indices, device_indices
|
||||||
|
if host_indices.numel() % self.page_size != 0:
|
||||||
|
raise ValueError(
|
||||||
|
"Index buffer transfer expects page-aligned indices for NSA."
|
||||||
|
)
|
||||||
|
host_page_indices = (
|
||||||
|
host_indices.reshape(-1, self.page_size)[:, 0] // self.page_size
|
||||||
|
)
|
||||||
|
device_page_indices = (
|
||||||
|
device_indices.reshape(-1, self.page_size)[:, 0] // self.page_size
|
||||||
|
)
|
||||||
|
return host_page_indices, device_page_indices
|
||||||
|
|
||||||
|
def _load_indexer_to_device_per_layer(
|
||||||
|
self, device_pool, host_indices, device_indices, layer_id, io_backend
|
||||||
|
):
|
||||||
|
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||||
|
host_indices, device_indices
|
||||||
|
)
|
||||||
|
use_kernel = io_backend == "kernel" and self.indexer_page_stride_size % 8 == 0
|
||||||
|
if use_kernel:
|
||||||
|
transfer_kv_per_layer_mla(
|
||||||
|
src=self.index_k_with_scale_buffer[layer_id],
|
||||||
|
dst=device_pool.index_k_with_scale_buffer[layer_id],
|
||||||
|
src_indices=host_page_indices,
|
||||||
|
dst_indices=device_page_indices,
|
||||||
|
item_size=self.indexer_page_stride_size,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
transfer_kv_direct(
|
||||||
|
src_layers=[self.index_k_with_scale_buffer[layer_id]],
|
||||||
|
dst_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
|
||||||
|
src_indices=host_page_indices,
|
||||||
|
dst_indices=device_page_indices,
|
||||||
|
page_size=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _backup_indexer_from_device_all_layer(
|
||||||
|
self, device_pool, host_indices, device_indices, io_backend
|
||||||
|
):
|
||||||
|
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||||
|
host_indices, device_indices
|
||||||
|
)
|
||||||
|
use_kernel = io_backend == "kernel" and self.indexer_page_stride_size % 8 == 0
|
||||||
|
if use_kernel:
|
||||||
|
transfer_kv_all_layer_mla(
|
||||||
|
src_layers=self.index_k_device_ptrs,
|
||||||
|
dst_layers=self.index_k_data_ptrs,
|
||||||
|
src_indices=device_page_indices,
|
||||||
|
dst_indices=host_page_indices,
|
||||||
|
item_size=self.indexer_page_stride_size,
|
||||||
|
num_layers=self.layer_num,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
transfer_kv_direct(
|
||||||
|
src_layers=device_pool.index_k_with_scale_buffer,
|
||||||
|
dst_layers=self.index_k_with_scale_buffer,
|
||||||
|
src_indices=device_page_indices,
|
||||||
|
dst_indices=host_page_indices,
|
||||||
|
page_size=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_to_device_per_layer(
|
||||||
|
self,
|
||||||
|
device_pool,
|
||||||
|
host_indices,
|
||||||
|
device_indices,
|
||||||
|
layer_id,
|
||||||
|
io_backend,
|
||||||
|
):
|
||||||
|
super().load_to_device_per_layer(
|
||||||
|
device_pool, host_indices, device_indices, layer_id, io_backend
|
||||||
|
)
|
||||||
|
self._load_indexer_to_device_per_layer(
|
||||||
|
device_pool, host_indices, device_indices, layer_id, io_backend
|
||||||
|
)
|
||||||
|
|
||||||
|
def backup_from_device_all_layer(
|
||||||
|
self, device_pool, host_indices, device_indices, io_backend
|
||||||
|
):
|
||||||
|
super().backup_from_device_all_layer(
|
||||||
|
device_pool, host_indices, device_indices, io_backend
|
||||||
|
)
|
||||||
|
self._backup_indexer_from_device_all_layer(
|
||||||
|
device_pool, host_indices, device_indices, io_backend
|
||||||
|
)
|
||||||
|
|||||||
@@ -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