Interface change for kvcache io to support page first layout (#8318)
This commit is contained in:
@@ -31,21 +31,17 @@ from typing import Dict, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
from sglang.srt.utils import get_bool_env_var, is_cuda, is_npu, next_power_of_2
|
||||
from sglang.srt.utils import get_bool_env_var, is_cuda, next_power_of_2
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GB = 1024 * 1024 * 1024
|
||||
_is_cuda = is_cuda()
|
||||
_is_npu = is_npu()
|
||||
if not _is_npu:
|
||||
from sgl_kernel.kvcacheio import transfer_kv_per_layer, transfer_kv_per_layer_mla
|
||||
|
||||
|
||||
class ReqToTokenPool:
|
||||
@@ -153,18 +149,6 @@ class KVCache(abc.ABC):
|
||||
) -> None:
|
||||
raise NotImplementedError()
|
||||
|
||||
@abc.abstractmethod
|
||||
def load_from_host_per_layer(
|
||||
self, host_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
raise NotImplementedError()
|
||||
|
||||
@abc.abstractmethod
|
||||
def backup_to_host_all_layer(
|
||||
self, host_pool, host_indices, device_indices, io_backend
|
||||
):
|
||||
raise NotImplementedError()
|
||||
|
||||
def register_layer_transfer_counter(self, layer_transfer_counter):
|
||||
self.layer_transfer_counter = layer_transfer_counter
|
||||
|
||||
@@ -253,12 +237,18 @@ class MHATokenToKVPool(KVCache):
|
||||
)
|
||||
for _ in range(self.layer_num)
|
||||
]
|
||||
self.token_stride = self.head_num * self.head_dim
|
||||
self.data_ptrs = torch.tensor(
|
||||
[x.data_ptr() for x in self.k_buffer + self.v_buffer],
|
||||
|
||||
self.k_data_ptrs = torch.tensor(
|
||||
[x.data_ptr() for x in self.k_buffer],
|
||||
dtype=torch.uint64,
|
||||
device=self.device,
|
||||
)
|
||||
self.v_data_ptrs = torch.tensor(
|
||||
[x.data_ptr() for x in self.v_buffer],
|
||||
dtype=torch.uint64,
|
||||
device=self.device,
|
||||
)
|
||||
self.data_ptrs = torch.cat([self.k_data_ptrs, self.v_data_ptrs], dim=0)
|
||||
self.data_strides = torch.tensor(
|
||||
[
|
||||
np.prod(x.shape[1:]) * x.dtype.itemsize
|
||||
@@ -347,47 +337,6 @@ class MHATokenToKVPool(KVCache):
|
||||
self.v_buffer[layer_id][chunk_indices] = v_chunk
|
||||
torch.cuda.synchronize()
|
||||
|
||||
def load_from_host_per_layer(
|
||||
self,
|
||||
host_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id,
|
||||
io_backend,
|
||||
):
|
||||
transfer_kv_per_layer(
|
||||
src_k=host_pool.k_buffer[layer_id],
|
||||
dst_k=self.k_buffer[layer_id],
|
||||
src_v=host_pool.v_buffer[layer_id],
|
||||
dst_v=self.v_buffer[layer_id],
|
||||
src_indices=host_indices,
|
||||
dst_indices=device_indices,
|
||||
io_backend=io_backend,
|
||||
page_size=self.page_size,
|
||||
item_size=self.token_stride,
|
||||
)
|
||||
|
||||
def backup_to_host_all_layer(
|
||||
self, host_pool, host_indices, device_indices, io_backend
|
||||
):
|
||||
# todo: specialized all layer kernels for the layer-non-contiguous memory pool
|
||||
for layer_id in range(self.start_layer, self.start_layer + self.layer_num):
|
||||
if layer_id - self.start_layer >= len(host_pool.k_buffer):
|
||||
raise ValueError(
|
||||
f"Layer ID {layer_id} exceeds the number of layers in host pool."
|
||||
)
|
||||
transfer_kv_per_layer(
|
||||
src_k=self.k_buffer[layer_id],
|
||||
dst_k=host_pool.k_buffer[layer_id],
|
||||
src_v=self.v_buffer[layer_id],
|
||||
dst_v=host_pool.v_buffer[layer_id],
|
||||
src_indices=device_indices,
|
||||
dst_indices=host_indices,
|
||||
io_backend=io_backend,
|
||||
page_size=self.page_size,
|
||||
item_size=self.token_stride,
|
||||
)
|
||||
|
||||
def _get_key_buffer(self, layer_id: int):
|
||||
# for internal use of referencing
|
||||
if self.store_dtype != self.dtype:
|
||||
@@ -602,16 +551,6 @@ class SWAKVPool(KVCache):
|
||||
layer_id_override=layer_id_pool,
|
||||
)
|
||||
|
||||
def load_from_host_per_layer(
|
||||
self, host_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
raise NotImplementedError("HiCache not supported for SWAKVPool.")
|
||||
|
||||
def backup_to_host_all_layer(
|
||||
self, host_pool, host_indices, device_indices, io_backend
|
||||
):
|
||||
raise NotImplementedError("HiCache not supported for SWAKVPool.")
|
||||
|
||||
|
||||
class AscendTokenToKVPool(MHATokenToKVPool):
|
||||
|
||||
@@ -823,7 +762,11 @@ class MLATokenToKVPool(KVCache):
|
||||
for _ in range(layer_num)
|
||||
]
|
||||
|
||||
self.token_stride = kv_lora_rank + qk_rope_head_dim
|
||||
self.data_ptrs = torch.tensor(
|
||||
[x.data_ptr() for x in self.kv_buffer],
|
||||
dtype=torch.uint64,
|
||||
device=self.device,
|
||||
)
|
||||
self.layer_transfer_counter = None
|
||||
|
||||
kv_size = self.get_kv_size_bytes()
|
||||
@@ -909,38 +852,6 @@ class MLATokenToKVPool(KVCache):
|
||||
self.kv_buffer[layer_id], loc, cache_k_nope, cache_k_rope
|
||||
)
|
||||
|
||||
def load_from_host_per_layer(
|
||||
self, host_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
transfer_kv_per_layer_mla(
|
||||
src=host_pool.kv_buffer[layer_id],
|
||||
dst=self.kv_buffer[layer_id],
|
||||
src_indices=host_indices,
|
||||
dst_indices=device_indices,
|
||||
io_backend=io_backend,
|
||||
page_size=self.page_size,
|
||||
item_size=self.token_stride,
|
||||
)
|
||||
|
||||
def backup_to_host_all_layer(
|
||||
self, host_pool, host_indices, device_indices, io_backend
|
||||
):
|
||||
# todo: specialized all layer kernels for the layer-non-contiguous memory pool
|
||||
for layer_id in range(self.start_layer, self.start_layer + self.layer_num):
|
||||
if layer_id - self.start_layer >= len(host_pool.kv_buffer):
|
||||
raise ValueError(
|
||||
f"Layer ID {layer_id} exceeds the number of layers in host pool."
|
||||
)
|
||||
transfer_kv_per_layer_mla(
|
||||
src=self.kv_buffer[layer_id],
|
||||
dst=host_pool.kv_buffer[layer_id],
|
||||
src_indices=device_indices,
|
||||
dst_indices=host_indices,
|
||||
io_backend=io_backend,
|
||||
page_size=self.page_size,
|
||||
item_size=self.token_stride,
|
||||
)
|
||||
|
||||
def get_cpu_copy(self, indices):
|
||||
torch.cuda.synchronize()
|
||||
kv_cache_cpu = []
|
||||
@@ -1131,20 +1042,6 @@ class DoubleSparseTokenToKVPool(KVCache):
|
||||
self.v_buffer[layer_id - self.start_layer][loc] = cache_v
|
||||
self.label_buffer[layer_id - self.start_layer][loc] = cache_label
|
||||
|
||||
def load_from_host_per_layer(
|
||||
self, host_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
raise NotImplementedError(
|
||||
"HiCache not supported for DoubleSparseTokenToKVPool."
|
||||
)
|
||||
|
||||
def backup_to_host_all_layer(
|
||||
self, host_pool, host_indices, device_indices, io_backend
|
||||
):
|
||||
raise NotImplementedError(
|
||||
"HiCache not supported for DoubleSparseTokenToKVPool."
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def copy_all_layer_kv_cache(
|
||||
|
||||
Reference in New Issue
Block a user