[Feature] Add FP4 KV Cache Design and support SM120 GPUs (#21601)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user