Hybrid kv cache for LLaMA4 (#6563)
Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Co-authored-by: tarinkk <rt572@physics.rutger.edu> Co-authored-by: tarinkk <rt572@rutgers.physics.edu> Co-authored-by: Hanming Lu <69857889+hanming-lu@users.noreply.github.com>
This commit is contained in:
@@ -27,10 +27,11 @@ KVCache actually holds the physical kv cache.
|
||||
import abc
|
||||
import logging
|
||||
from contextlib import nullcontext
|
||||
from typing import List, Optional, Tuple, Union
|
||||
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
|
||||
|
||||
@@ -66,6 +67,7 @@ class ReqToTokenPool:
|
||||
self.req_to_token = torch.zeros(
|
||||
(size, max_context_len), dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
self.free_slots = list(range(size))
|
||||
|
||||
def write(self, indices, values):
|
||||
@@ -191,7 +193,6 @@ class MHATokenToKVPool(KVCache):
|
||||
start_layer,
|
||||
end_layer,
|
||||
)
|
||||
|
||||
self.head_num = head_num
|
||||
self.head_dim = head_dim
|
||||
|
||||
@@ -392,10 +393,14 @@ class MHATokenToKVPool(KVCache):
|
||||
cache_v: torch.Tensor,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
layer_id_override: Optional[int] = None,
|
||||
):
|
||||
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
||||
|
||||
layer_id = layer.layer_id
|
||||
if layer_id_override is not None:
|
||||
layer_id = layer_id_override
|
||||
else:
|
||||
layer_id = layer.layer_id
|
||||
if cache_k.dtype != self.dtype:
|
||||
if k_scale is not None:
|
||||
cache_k.div_(k_scale)
|
||||
@@ -431,6 +436,136 @@ class MHATokenToKVPool(KVCache):
|
||||
)
|
||||
|
||||
|
||||
class SWAKVPool(KVCache):
|
||||
"""KV cache with separate pools for full and SWA attention layers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
size_swa: int,
|
||||
dtype: torch.dtype,
|
||||
head_num: int,
|
||||
head_dim: int,
|
||||
swa_attention_layer_ids: List[int],
|
||||
full_attention_layer_ids: List[int],
|
||||
enable_kvcache_transpose: bool,
|
||||
device: str,
|
||||
):
|
||||
self.size = size
|
||||
self.size_swa = size_swa
|
||||
self.dtype = dtype
|
||||
self.device = device
|
||||
self.swa_layer_nums = len(swa_attention_layer_ids)
|
||||
self.full_layer_nums = len(full_attention_layer_ids)
|
||||
self.page_size = 1
|
||||
# TODO MHATransposedTokenToKVPool if enable_kvcache_transpose is True
|
||||
assert not enable_kvcache_transpose
|
||||
TokenToKVPoolClass = MHATokenToKVPool
|
||||
self.swa_kv_pool = TokenToKVPoolClass(
|
||||
size=size_swa,
|
||||
page_size=self.page_size,
|
||||
dtype=dtype,
|
||||
head_num=head_num,
|
||||
head_dim=head_dim,
|
||||
layer_num=self.swa_layer_nums,
|
||||
device=device,
|
||||
enable_memory_saver=False,
|
||||
)
|
||||
self.full_kv_pool = TokenToKVPoolClass(
|
||||
size=size,
|
||||
page_size=self.page_size,
|
||||
dtype=dtype,
|
||||
head_num=head_num,
|
||||
head_dim=head_dim,
|
||||
layer_num=self.full_layer_nums,
|
||||
device=device,
|
||||
enable_memory_saver=False,
|
||||
)
|
||||
self.layers_mapping: Dict[int, Tuple[int, bool]] = {}
|
||||
for full_attn_layer_id, global_layer_id in enumerate(full_attention_layer_ids):
|
||||
self.layers_mapping[global_layer_id] = (full_attn_layer_id, False)
|
||||
for swa_layer_id, global_layer_id in enumerate(swa_attention_layer_ids):
|
||||
self.layers_mapping[global_layer_id] = (swa_layer_id, True)
|
||||
self.full_to_swa_index_mapping: Optional[torch.Tensor] = None
|
||||
|
||||
def get_kv_size_bytes(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_contiguous_buf_infos(self):
|
||||
full_kv_data_ptrs, full_kv_data_lens, full_kv_item_lens = (
|
||||
self.full_kv_pool.get_contiguous_buf_infos()
|
||||
)
|
||||
swa_kv_data_ptrs, swa_kv_data_lens, swa_kv_item_lens = (
|
||||
self.swa_kv_pool.get_contiguous_buf_infos()
|
||||
)
|
||||
|
||||
kv_data_ptrs = full_kv_data_ptrs + swa_kv_data_ptrs
|
||||
kv_data_lens = full_kv_data_lens + swa_kv_data_lens
|
||||
kv_item_lens = full_kv_item_lens + swa_kv_item_lens
|
||||
|
||||
return kv_data_ptrs, kv_data_lens, kv_item_lens
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
||||
if is_swa:
|
||||
return self.swa_kv_pool.get_key_buffer(layer_id_pool)
|
||||
else:
|
||||
return self.full_kv_pool.get_key_buffer(layer_id_pool)
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
||||
if is_swa:
|
||||
return self.swa_kv_pool.get_value_buffer(layer_id_pool)
|
||||
else:
|
||||
return self.full_kv_pool.get_value_buffer(layer_id_pool)
|
||||
|
||||
def get_kv_buffer(self, layer_id: int):
|
||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
||||
if is_swa:
|
||||
return self.swa_kv_pool.get_kv_buffer(layer_id_pool)
|
||||
else:
|
||||
return self.full_kv_pool.get_kv_buffer(layer_id_pool)
|
||||
|
||||
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
|
||||
assert self.full_to_swa_index_mapping is not None
|
||||
return self.full_to_swa_index_mapping[kv_indices].to(torch.int32)
|
||||
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
k_scale: float = 1.0,
|
||||
v_scale: float = 1.0,
|
||||
):
|
||||
|
||||
layer_id = layer.layer_id
|
||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
||||
if is_swa:
|
||||
if self.full_to_swa_index_mapping is not None:
|
||||
loc = self.translate_loc_from_full_to_swa(loc)
|
||||
self.swa_kv_pool.set_kv_buffer(
|
||||
None,
|
||||
loc,
|
||||
cache_k,
|
||||
cache_v,
|
||||
k_scale,
|
||||
v_scale,
|
||||
layer_id_override=layer_id_pool,
|
||||
)
|
||||
else:
|
||||
self.full_kv_pool.set_kv_buffer(
|
||||
None,
|
||||
loc,
|
||||
cache_k,
|
||||
cache_v,
|
||||
k_scale,
|
||||
v_scale,
|
||||
layer_id_override=layer_id_pool,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def set_mla_kv_buffer_kernel(
|
||||
kv_buffer_ptr,
|
||||
|
||||
Reference in New Issue
Block a user