[Feature] Add FP4 KV Cache Design and support SM120 GPUs (#21601)

This commit is contained in:
Sam (Kesen Li)
2026-07-17 14:49:43 -07:00
committed by GitHub
parent 7fc3fb9657
commit ec6a3163b7
19 changed files with 1829 additions and 327 deletions
+496 -36
View File
@@ -51,6 +51,9 @@ from sglang.srt.configs.mamba_utils import BaseLinearStateParams
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import (
UnquantizedKVCacheMethod,
)
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator
from sglang.srt.mem_cache.kv_vmm_backing import KvVmmBufferOwner
@@ -73,6 +76,7 @@ from sglang.srt.utils import (
cpu_has_amx_support,
is_cpu,
is_cuda,
is_float4_e2m1fn_x2,
is_hip,
is_npu,
next_power_of_2,
@@ -1387,6 +1391,24 @@ class KVCache(abc.ABC):
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
raise NotImplementedError()
def get_kv_cache_quant_method(self) -> Any:
"""Return the concrete KV quant method, unwrapping composite KV pools."""
fallback = None
for pool in (
self,
getattr(self, "full_kv_pool", None),
getattr(self, "swa_kv_pool", None),
):
if pool is None:
continue
quant_method = getattr(pool, "quant_method", None)
if quant_method is None:
continue
if getattr(quant_method, "name", None) != "unquantized":
return quant_method
fallback = quant_method
return fallback
def maybe_get_custom_mem_pool(self):
return self.custom_mem_pool
@@ -1411,6 +1433,7 @@ class MHATokenToKVPool(KVCache):
enable_alt_stream: bool = True,
enable_kv_cache_copy: bool = False,
kv_cache_layout: Optional[str] = None,
quant_method=None,
post_capture_active: bool = False,
):
if post_capture_active:
@@ -1482,6 +1505,10 @@ class MHATokenToKVPool(KVCache):
assert self.head_dim % self._kv_vector_x == 0
assert self.v_head_dim % self._kv_vector_x == 0
self.quant_method = (
quant_method if quant_method is not None else UnquantizedKVCacheMethod()
)
self._create_buffers()
self.device_module = torch.get_device_module(self.device)
@@ -1554,12 +1581,100 @@ class MHATokenToKVPool(KVCache):
self._kv_copy_config,
)
@property
def is_quantized_kv_cache(self) -> bool:
return not isinstance(self.quant_method, UnquantizedKVCacheMethod)
def _create_buffers(self):
if self.post_capture_active:
self._alloc_post_capture_buffers()
if self.is_quantized_kv_cache:
if self.post_capture_active:
raise NotImplementedError(
"Post-capture KV backing is not supported for quantized KV cache."
)
self._create_quantized_buffers()
else:
self._create_buffers_normal()
self.k_scale_buffer = None
self.v_scale_buffer = None
self.dq_k_buffer = None
self.dq_v_buffer = None
if self.post_capture_active:
self._alloc_post_capture_buffers()
else:
self._create_buffers_normal()
self._kv_buffer_descs = self._build_kv_buffer_descs()
self._init_data_ptrs_and_strides()
def _create_quantized_buffers(self):
# Quantized recipes own packed-data, scale, and workspace shapes.
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.enable_custom_mem_pool
else nullcontext()
):
buf = self.quant_method.create_buffers(
self.size + self.page_size,
self.head_num,
self.head_dim,
self.layer_num,
self.device,
)
self.k_buffer = buf["k_buffer"]
self.v_buffer = buf["v_buffer"]
self.k_scale_buffer = buf.get("k_scale_buffer")
self.v_scale_buffer = buf.get("v_scale_buffer")
self.dq_k_buffer = buf.get("dq_k_buffer")
self.dq_v_buffer = buf.get("dq_v_buffer")
self.store_dtype = buf.get("store_dtype", torch.uint8)
self._check_quantized_buffer_access_requirements()
def _check_quantized_buffer_access_requirements(self):
expected_workspace_dtype = self.quant_method.dequant_workspace_dtype()
has_k_workspace = self.dq_k_buffer is not None
has_v_workspace = self.dq_v_buffer is not None
if has_k_workspace != has_v_workspace:
raise RuntimeError(
f"KV cache method {self.quant_method.name!r} created only one "
"dequant workspace buffer."
)
if expected_workspace_dtype is None:
if has_k_workspace:
raise RuntimeError(
f"KV cache method {self.quant_method.name!r} does not declare "
"DEQUANT_WORKSPACE access but created dequant buffers."
)
return
if not has_k_workspace:
raise RuntimeError(
f"KV cache method {self.quant_method.name!r} declares "
"DEQUANT_WORKSPACE access but did not create dequant buffers."
)
if (
self.dq_k_buffer.dtype != expected_workspace_dtype
or self.dq_v_buffer.dtype != expected_workspace_dtype
):
raise RuntimeError(
f"KV cache method {self.quant_method.name!r} declares dequant "
f"workspace dtype {expected_workspace_dtype}, but created "
f"{self.dq_k_buffer.dtype}/{self.dq_v_buffer.dtype}."
)
def _slot_move_pointer_buffers(self):
"""Buffers whose pointers/strides are used when KV slots are remapped.
FP4 KV cache stores data and per-block scales separately, so slot moves
must update both. This list feeds data_ptrs/data_strides; it does not
copy tensor contents by itself.
"""
buffers = [*self.k_buffer, *self.v_buffer]
if getattr(self, "k_scale_buffer", None) is not None:
buffers.extend([*self.k_scale_buffer, *self.v_scale_buffer])
return buffers
def _init_data_ptrs_and_strides(self):
self.k_data_ptrs = torch.tensor(
[x.data_ptr() for x in self.k_buffer],
dtype=torch.uint64,
@@ -1570,11 +1685,16 @@ class MHATokenToKVPool(KVCache):
dtype=torch.uint64,
device=self.device,
)
self.data_ptrs = torch.cat([self.k_data_ptrs, self.v_data_ptrs], dim=0)
slot_move_pointer_buffers = self._slot_move_pointer_buffers()
self.data_ptrs = torch.tensor(
[x.data_ptr() for x in slot_move_pointer_buffers],
dtype=torch.uint64,
device=self.device,
)
self.data_strides = torch.tensor(
[
np.prod(x.shape[1:]) * x.dtype.itemsize
for x in self.k_buffer + self.v_buffer
for x in slot_move_pointer_buffers
],
device=self.device,
)
@@ -1716,6 +1836,14 @@ class MHATokenToKVPool(KVCache):
def _clear_buffers(self):
del self.k_buffer
del self.v_buffer
if hasattr(self, "k_scale_buffer") and self.k_scale_buffer is not None:
del self.k_scale_buffer
if hasattr(self, "v_scale_buffer") and self.v_scale_buffer is not None:
del self.v_scale_buffer
if hasattr(self, "dq_k_buffer") and self.dq_k_buffer is not None:
del self.dq_k_buffer
if hasattr(self, "dq_v_buffer") and self.dq_v_buffer is not None:
del self.dq_v_buffer
if self._post_capture_owner is not None:
self._post_capture_owner.close()
self._post_capture_owner = None
@@ -1723,12 +1851,14 @@ class MHATokenToKVPool(KVCache):
def get_kv_size_bytes(self):
assert hasattr(self, "k_buffer")
assert hasattr(self, "v_buffer")
k_size_bytes = 0
for k_cache in self.k_buffer:
k_size_bytes += get_tensor_size_bytes(k_cache)
v_size_bytes = 0
for v_cache in self.v_buffer:
v_size_bytes += get_tensor_size_bytes(v_cache)
k_size_bytes = get_tensor_size_bytes(self.k_buffer)
v_size_bytes = get_tensor_size_bytes(self.v_buffer)
if getattr(self, "k_scale_buffer", None) is not None:
k_size_bytes += get_tensor_size_bytes(self.k_scale_buffer)
v_size_bytes += get_tensor_size_bytes(self.v_scale_buffer)
if getattr(self, "dq_k_buffer", None) is not None:
k_size_bytes += get_tensor_size_bytes(self.dq_k_buffer)
v_size_bytes += get_tensor_size_bytes(self.dq_v_buffer)
return k_size_bytes, v_size_bytes
# for disagg
@@ -1798,9 +1928,19 @@ class MHATokenToKVPool(KVCache):
def _get_key_buffer(self, layer_id: int):
# for internal use of referencing
local_layer_id = layer_id - self.start_layer
if (
self.is_quantized_kv_cache
and self.quant_method.needs_plain_kv_dequant_read()
):
return self.quant_method.dequantize_kv_tensor(
self.k_buffer[local_layer_id],
self.k_scale_buffer[local_layer_id],
layer_id,
)
if self.store_dtype != self.dtype:
return self.k_buffer[layer_id - self.start_layer].view(self.dtype)
return self.k_buffer[layer_id - self.start_layer]
return self.k_buffer[local_layer_id].view(self.dtype)
return self.k_buffer[local_layer_id]
def get_key_buffer(self, layer_id: int):
# note: get_key_buffer is hooked with synchronization for layer-wise KV cache loading
@@ -1812,9 +1952,19 @@ class MHATokenToKVPool(KVCache):
def _get_value_buffer(self, layer_id: int):
# for internal use of referencing
local_layer_id = layer_id - self.start_layer
if (
self.is_quantized_kv_cache
and self.quant_method.needs_plain_kv_dequant_read()
):
return self.quant_method.dequantize_kv_tensor(
self.v_buffer[local_layer_id],
self.v_scale_buffer[local_layer_id],
layer_id,
)
if self.store_dtype != self.dtype:
return self.v_buffer[layer_id - self.start_layer].view(self.dtype)
return self.v_buffer[layer_id - self.start_layer]
return self.v_buffer[local_layer_id].view(self.dtype)
return self.v_buffer[local_layer_id]
def get_value_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
@@ -1839,10 +1989,25 @@ class MHATokenToKVPool(KVCache):
# Catch stale slot ids here instead of as illegal-addr / silent KV
# corruption in the store_kvcache write (gated on SGLANG_ENABLE_ASYNC_ASSERT).
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA)")
if layer_id_override is not None:
layer_id = layer_id_override
else:
layer_id = layer.layer_id
layer_id = (
layer_id_override if layer_id_override is not None else layer.layer_id
)
global_layer_id = layer.layer_id if layer is not None else layer_id
if self.is_quantized_kv_cache:
if dcp_kv_mask is not None:
raise RuntimeError("dcp_kv_mask is not supported for FP4 KV cache.")
self._set_quantized_kv_buffer(
layer_id,
global_layer_id,
loc,
cache_k,
cache_v,
k_scale,
v_scale,
)
return
if cache_k.dtype != self.dtype:
if k_scale is not None:
cache_k.div_(k_scale)
@@ -1936,6 +2101,255 @@ class MHATokenToKVPool(KVCache):
same_kv_dim=self.same_kv_dim,
)
def _quantized_scales(self, global_layer_id: int, k_scale, v_scale):
if k_scale is None and hasattr(self.quant_method, "k_scales_gpu"):
k_scale = self.quant_method.k_scales_gpu[
global_layer_id : global_layer_id + 1
]
v_scale = self.quant_method.v_scales_gpu[
global_layer_id : global_layer_id + 1
]
return k_scale, v_scale
def _set_quantized_kv_buffer(
self,
layer_id: int,
global_layer_id: int,
loc_info,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
k_scale=None,
v_scale=None,
) -> None:
loc, _, _ = unwrap_write_loc(loc_info)
local_layer_id = layer_id - self.start_layer
k_scale, v_scale = self._quantized_scales(global_layer_id, k_scale, v_scale)
self.quant_method.quantize_and_store(
self.k_buffer[local_layer_id],
self.v_buffer[local_layer_id],
(
self.k_scale_buffer[local_layer_id]
if self.k_scale_buffer is not None
else None
),
(
self.v_scale_buffer[local_layer_id]
if self.v_scale_buffer is not None
else None
),
loc,
cache_k,
cache_v,
k_scale,
v_scale,
)
def get_raw_kv_buffer(
self, layer_id: int
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
local_layer_id = layer_id - self.start_layer
if self.k_scale_buffer is None or self.v_scale_buffer is None:
raise RuntimeError("Raw FP4 KV cache requested from a non-FP4 KV pool.")
k_scale = self.k_scale_buffer[local_layer_id]
v_scale = self.v_scale_buffer[local_layer_id]
scale_view_dtype = self.quant_method.scale_buffer_view_dtype()
if scale_view_dtype is not None:
k_scale = k_scale.view(scale_view_dtype)
v_scale = v_scale.view(scale_view_dtype)
return (
self.k_buffer[local_layer_id],
self.v_buffer[local_layer_id],
k_scale,
v_scale,
)
def get_dequant_workspace(self) -> tuple[torch.Tensor, torch.Tensor]:
if self.dq_k_buffer is None or self.dq_v_buffer is None:
raise RuntimeError(
"Dequant workspace requested from a KV pool without FP4 dequant buffers."
)
return self.dq_k_buffer, self.dq_v_buffer
def get_flashinfer_dequant_workspace_kv_buffer(
self,
layer: RadixAttention,
req_to_token: torch.Tensor,
req_pool_indices_cpu,
extend_prefix_lens_cpu,
extend_seq_lens_cpu,
page_size: int,
*,
prepare_workspace: bool,
use_ragged: bool,
k_cur: Optional[torch.Tensor] = None,
v_cur: Optional[torch.Tensor] = None,
layer_id_override: Optional[int] = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Return the FlashInfer FP8 KV view for a quantized KV cache.
FlashInfer prefill consumes FP8 KV. Quantized pools store packed FP4 plus
per-block scales, so the pool owns the dequant workspace and returns the
view shape expected by FlashInfer.
"""
if not self.is_quantized_kv_cache:
raise RuntimeError(
"FlashInfer quantized KV buffer requested from a non-quantized KV pool."
)
if prepare_workspace:
transfer_cur_kv = not use_ragged
k_cur_fp8 = (
k_cur.to(torch.float8_e4m3fn)
if k_cur is not None and transfer_cur_kv
else None
)
v_cur_fp8 = (
v_cur.to(torch.float8_e4m3fn)
if v_cur is not None and transfer_cur_kv
else None
)
self._prepare_dequant_extend_workspace(
layer.layer_id if layer_id_override is None else layer_id_override,
layer.layer_id,
req_to_token,
req_pool_indices_cpu,
extend_prefix_lens_cpu,
extend_seq_lens_cpu,
page_size,
k_cur_fp8=k_cur_fp8,
v_cur_fp8=v_cur_fp8,
)
k_buffer_dq, v_buffer_dq = self.get_dequant_workspace()
return (
k_buffer_dq.view(-1, layer.tp_k_head_num, layer.head_dim),
v_buffer_dq.view(-1, layer.tp_v_head_num, layer.head_dim),
)
def get_flashinfer_decode_dequant_workspace_kv_buffer(
self,
layer: RadixAttention,
req_to_token: torch.Tensor,
req_pool_indices,
seq_lens,
*,
layer_id_override: Optional[int] = None,
) -> tuple[torch.Tensor, torch.Tensor]:
if not self.is_quantized_kv_cache:
raise RuntimeError(
"FlashInfer dequant workspace requested from a non-quantized KV pool."
)
self._prepare_dequant_decode_workspace(
layer.layer_id if layer_id_override is None else layer_id_override,
layer.layer_id,
req_to_token,
req_pool_indices,
seq_lens,
)
k_buffer_dq, v_buffer_dq = self.get_dequant_workspace()
return (
k_buffer_dq.view(-1, layer.tp_k_head_num, layer.head_dim),
v_buffer_dq.view(-1, layer.tp_v_head_num, layer.head_dim),
)
@staticmethod
def _to_cpu_int_list(values) -> list[int]:
if isinstance(values, list):
return [int(value) for value in values]
if isinstance(values, torch.Tensor):
return [int(value) for value in values.cpu().tolist()]
return [int(value) for value in values]
def _prepare_dequant_extend_workspace(
self,
layer_id: int,
global_layer_id: int,
req_to_token: torch.Tensor,
req_pool_indices_cpu,
extend_prefix_lens_cpu,
extend_seq_lens_cpu,
page_size: int,
k_cur_fp8: Optional[torch.Tensor] = None,
v_cur_fp8: Optional[torch.Tensor] = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Build the shared FP8 workspace used by FlashInfer extend attention.
Cached prefix tokens are stored as packed FP4 plus per-block scales, so
paged prefill dequantizes those prefix tokens into the FP8 workspace.
The current extend chunk can already be FP8 and is copied into the same
workspace after the prefix region.
"""
k_fp4, v_fp4, k_scales, v_scales = self.get_raw_kv_buffer(layer_id)
dq_k, dq_v = self.get_dequant_workspace()
cur_batch_start_loc_cpu = 0
cur_token_idx_dq = page_size
for i in range(len(req_pool_indices_cpu)):
req_idx = int(req_pool_indices_cpu[i])
prev_len = int(extend_prefix_lens_cpu[i])
extend_len = int(extend_seq_lens_cpu[i])
if prev_len > 0:
prev_indices = req_to_token[req_idx, :prev_len]
k_prev_fp8, v_prev_fp8 = self.quant_method.dequantize_prev_kv(
k_fp4[prev_indices],
k_scales[prev_indices],
v_fp4[prev_indices],
v_scales[prev_indices],
global_layer_id,
)
dq_k[cur_token_idx_dq : cur_token_idx_dq + prev_len] = k_prev_fp8
dq_v[cur_token_idx_dq : cur_token_idx_dq + prev_len] = v_prev_fp8
if k_cur_fp8 is not None:
cur_end = cur_batch_start_loc_cpu + extend_len
dst_start = cur_token_idx_dq + prev_len
dst_end = dst_start + extend_len
dq_k[dst_start:dst_end] = k_cur_fp8[cur_batch_start_loc_cpu:cur_end]
dq_v[dst_start:dst_end] = v_cur_fp8[cur_batch_start_loc_cpu:cur_end]
cur_batch_start_loc_cpu = cur_end
workspace_len = prev_len + (extend_len if k_cur_fp8 is not None else 0)
cur_token_idx_dq = (
(cur_token_idx_dq + workspace_len + page_size - 1)
// page_size
* page_size
)
return dq_k, dq_v
def _prepare_dequant_decode_workspace(
self,
layer_id: int,
global_layer_id: int,
req_to_token: torch.Tensor,
req_pool_indices,
seq_lens,
) -> tuple[torch.Tensor, torch.Tensor]:
k_fp4, v_fp4, k_scales, v_scales = self.get_raw_kv_buffer(layer_id)
dq_k, dq_v = self.get_dequant_workspace()
req_pool_indices_cpu = self._to_cpu_int_list(req_pool_indices)
seq_lens_cpu = self._to_cpu_int_list(seq_lens)
for req_idx, seq_len in zip(req_pool_indices_cpu, seq_lens_cpu):
if seq_len <= 0:
continue
kv_indices = req_to_token[req_idx, :seq_len]
k_prev_fp8, v_prev_fp8 = self.quant_method.dequantize_prev_kv(
k_fp4[kv_indices],
k_scales[kv_indices],
v_fp4[kv_indices],
v_scales[kv_indices],
global_layer_id,
)
dq_k[kv_indices] = k_prev_fp8
dq_v[kv_indices] = v_prev_fp8
return dq_k, dq_v
def set_kv_buffer_prefix_valid(
self,
layer: RadixAttention,
@@ -2047,6 +2461,10 @@ class MHATokenToKVPool(KVCache):
# per-layer buffers here ignore page_size in move_kv_cache_native.
if self.use_native_move_kv_cache:
move_kv_cache_native(self.k_buffer, self.v_buffer, tgt_loc, src_loc)
if getattr(self, "k_scale_buffer", None) is not None:
move_kv_cache_native(
self.k_scale_buffer, self.v_scale_buffer, tgt_loc, src_loc
)
return
N = tgt_loc.numel()
@@ -2266,10 +2684,10 @@ class MHATokenToKVPoolFP4(MHATokenToKVPool):
cache_k_nope_fp4_sf = self.k_scale_buffer[layer_id - self.start_layer]
from sglang.srt.layers.quantization.kvfp4_tensor import (
BlockFP4KVQuantizeUtil,
FP4MXBlock16KVQuantizeUtil,
)
cache_k_nope_fp4_dequant = BlockFP4KVQuantizeUtil.batched_dequantize(
cache_k_nope_fp4_dequant = FP4MXBlock16KVQuantizeUtil.batched_dequantize(
cache_k_nope_fp4, cache_k_nope_fp4_sf
)
return cache_k_nope_fp4_dequant
@@ -2284,10 +2702,10 @@ class MHATokenToKVPoolFP4(MHATokenToKVPool):
cache_v_nope_fp4_sf = self.v_scale_buffer[layer_id - self.start_layer]
from sglang.srt.layers.quantization.kvfp4_tensor import (
BlockFP4KVQuantizeUtil,
FP4MXBlock16KVQuantizeUtil,
)
cache_v_nope_fp4_dequant = BlockFP4KVQuantizeUtil.batched_dequantize(
cache_v_nope_fp4_dequant = FP4MXBlock16KVQuantizeUtil.batched_dequantize(
cache_v_nope_fp4, cache_v_nope_fp4_sf
)
return cache_v_nope_fp4_dequant
@@ -2318,11 +2736,15 @@ class MHATokenToKVPoolFP4(MHATokenToKVPool):
cache_v.div_(v_scale)
from sglang.srt.layers.quantization.kvfp4_tensor import (
BlockFP4KVQuantizeUtil,
FP4MXBlock16KVQuantizeUtil,
)
cache_k, cache_k_fp4_sf = BlockFP4KVQuantizeUtil.batched_quantize(cache_k)
cache_v, cache_v_fp4_sf = BlockFP4KVQuantizeUtil.batched_quantize(cache_v)
cache_k, cache_k_fp4_sf = FP4MXBlock16KVQuantizeUtil.batched_quantize(
cache_k
)
cache_v, cache_v_fp4_sf = FP4MXBlock16KVQuantizeUtil.batched_quantize(
cache_v
)
if self.store_dtype != self.dtype:
cache_k = cache_k.view(self.store_dtype)
@@ -2526,6 +2948,7 @@ class HybridLinearKVPool(KVCache):
qk_rope_head_dim: int = None,
start_layer: Optional[int] = None,
full_kv_pool_class: Optional[type] = None,
quant_method=None,
# When provided (shared-KV-pool path), use this pool for the
# full-attention layers instead of constructing one internally.
full_kv_pool: Optional[KVCache] = None,
@@ -2551,20 +2974,28 @@ class HybridLinearKVPool(KVCache):
self.full_kv_pool = full_kv_pool
elif not use_mla:
TokenToKVPoolClass = MHATokenToKVPool
quant_method_kwarg = {"quant_method": quant_method}
if current_platform.is_out_of_tree():
TokenToKVPoolClass = current_platform.get_mha_kv_pool_cls()
quant_method_kwarg = {}
elif _is_npu:
assert not is_float4_e2m1fn_x2(
dtype
), "FP4 is not supported on NPU yet."
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMHATokenToKVPool,
)
TokenToKVPoolClass = NPUMHATokenToKVPool
quant_method_kwarg = {}
elif full_kv_pool_class is not None:
# Caller-selected MHA layout variant (e.g. the page-major
# PageMajorMHATokenToKVPool). NPU / out-of-tree classes keep
# priority since they don't understand alternate layouts.
TokenToKVPoolClass = full_kv_pool_class
else:
TokenToKVPoolClass = MHATokenToKVPool
post_capture_kwargs = (
{"post_capture_active": True} if post_capture_active else {}
@@ -2579,6 +3010,7 @@ class HybridLinearKVPool(KVCache):
device=device,
enable_memory_saver=enable_memory_saver,
enable_kv_cache_copy=enable_kv_cache_copy,
**quant_method_kwarg,
**post_capture_kwargs,
)
else:
@@ -2665,14 +3097,18 @@ class HybridLinearKVPool(KVCache):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
def get_key_buffer(self, layer_id: int):
def get_key_buffer(self, layer_id: int, scale: Optional[float] = None):
self._wait_for_layer(layer_id)
layer_id = self._transfer_full_attention_id(layer_id)
if scale is not None:
return self.full_kv_pool.get_key_buffer(layer_id, scale)
return self.full_kv_pool.get_key_buffer(layer_id)
def get_value_buffer(self, layer_id: int):
def get_value_buffer(self, layer_id: int, scale: Optional[float] = None):
self._wait_for_layer(layer_id)
layer_id = self._transfer_full_attention_id(layer_id)
if scale is not None:
return self.full_kv_pool.get_value_buffer(layer_id, scale)
return self.full_kv_pool.get_value_buffer(layer_id)
def get_kv_buffer(self, layer_id: int):
@@ -2680,6 +3116,30 @@ class HybridLinearKVPool(KVCache):
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_kv_buffer(layer_id)
def get_raw_kv_buffer(
self, layer_id: int
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
self._wait_for_layer(layer_id)
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_raw_kv_buffer(layer_id)
def get_dequant_workspace(self) -> tuple[torch.Tensor, torch.Tensor]:
return self.full_kv_pool.get_dequant_workspace()
def get_flashinfer_dequant_workspace_kv_buffer(self, layer, *args, **kwargs):
self._wait_for_layer(layer.layer_id)
local_layer_id = self._transfer_full_attention_id(layer.layer_id)
return self.full_kv_pool.get_flashinfer_dequant_workspace_kv_buffer(
layer, *args, layer_id_override=local_layer_id, **kwargs
)
def get_flashinfer_decode_dequant_workspace_kv_buffer(self, layer, *args, **kwargs):
self._wait_for_layer(layer.layer_id)
local_layer_id = self._transfer_full_attention_id(layer.layer_id)
return self.full_kv_pool.get_flashinfer_decode_dequant_workspace_kv_buffer(
layer, *args, layer_id_override=local_layer_id, **kwargs
)
@contextmanager
def _transfer_id_context(self, layer: RadixAttention):
@contextmanager
@@ -2712,7 +3172,7 @@ class HybridLinearKVPool(KVCache):
if not self.use_mla:
write_loc = full_loc if full_loc is not None else loc
self.full_kv_pool.set_kv_buffer(
None,
layer,
write_loc,
cache_k,
cache_v,
@@ -3095,10 +3555,10 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
cache_k_nope_fp4_sf = self.kv_scale_buffer[layer_id - self.start_layer]
from sglang.srt.layers.quantization.kvfp4_tensor import (
BlockFP4KVQuantizeUtil,
FP4MXBlock16KVQuantizeUtil,
)
cache_k_nope_fp4_dequant = BlockFP4KVQuantizeUtil.batched_dequantize(
cache_k_nope_fp4_dequant = FP4MXBlock16KVQuantizeUtil.batched_dequantize(
cache_k_nope_fp4, cache_k_nope_fp4_sf
)
return cache_k_nope_fp4_dequant
@@ -3119,10 +3579,10 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
assert not self.dsa_kv_cache_store_fp8
if cache_k.dtype != self.dtype:
from sglang.srt.layers.quantization.kvfp4_tensor import (
BlockFP4KVQuantizeUtil,
FP4MXBlock16KVQuantizeUtil,
)
cache_k_fp4, cache_k_fp4_sf = BlockFP4KVQuantizeUtil.batched_quantize(
cache_k_fp4, cache_k_fp4_sf = FP4MXBlock16KVQuantizeUtil.batched_quantize(
cache_k
)
@@ -3158,14 +3618,14 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
else:
if cache_k_nope.dtype != self.dtype:
from sglang.srt.layers.quantization.kvfp4_tensor import (
BlockFP4KVQuantizeUtil,
FP4MXBlock16KVQuantizeUtil,
)
cache_k_nope_fp4, cache_k_nope_fp4_sf = (
BlockFP4KVQuantizeUtil.batched_quantize(cache_k_nope)
FP4MXBlock16KVQuantizeUtil.batched_quantize(cache_k_nope)
)
cache_k_rope_fp4, cache_k_rope_fp4_sf = (
BlockFP4KVQuantizeUtil.batched_quantize(cache_k_rope)
FP4MXBlock16KVQuantizeUtil.batched_quantize(cache_k_rope)
)
if self.store_dtype != self.dtype: