5706 lines
229 KiB
Python
5706 lines
229 KiB
Python
"""
|
||
Copyright 2023-2024 SGLang Team
|
||
Licensed under the Apache License, Version 2.0 (the "License");
|
||
you may not use this file except in compliance with the License.
|
||
You may obtain a copy of the License at
|
||
|
||
http://www.apache.org/licenses/LICENSE-2.0
|
||
|
||
Unless required by applicable law or agreed to in writing, software
|
||
distributed under the License is distributed on an "AS IS" BASIS,
|
||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||
See the License for the specific language governing permissions and
|
||
limitations under the License.
|
||
|
||
Memory pool.
|
||
|
||
SGLang has two levels of memory pool.
|
||
ReqToTokenPool maps a request to its token locations.
|
||
TokenToKVPoolAllocator manages the indices to kv cache data.
|
||
KVCache actually holds the physical kv cache.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import abc
|
||
import copy
|
||
import dataclasses
|
||
import logging
|
||
import math
|
||
import os
|
||
from contextlib import contextmanager, nullcontext
|
||
from dataclasses import dataclass, fields
|
||
from functools import cached_property
|
||
from typing import TYPE_CHECKING, Any, Callable, List, Optional, Tuple, Union
|
||
|
||
import numpy as np
|
||
import torch
|
||
import triton
|
||
import triton.language as tl
|
||
|
||
from sglang.kernels.ops.attention.dsa.quant_k_cache import (
|
||
quantize_k_cache,
|
||
quantize_k_cache_separate,
|
||
)
|
||
from sglang.kernels.ops.kvcache.cache_move import (
|
||
copy_all_layer_kv_cache_func,
|
||
set_kv_buffer_prefix_valid_tiled,
|
||
)
|
||
from sglang.kernels.ops.kvcache.kvcache import can_use_store_cache, store_cache
|
||
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
|
||
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.dcp.layout import maybe_dcp_kernel_indices
|
||
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.index_key_cache import IndexKeyCache
|
||
from sglang.srt.mem_cache.kv_vmm_backing import KvVmmBufferOwner
|
||
from sglang.srt.mem_cache.layout.page_major import (
|
||
build_page_major_mamba_views,
|
||
mamba_entry_bytes,
|
||
)
|
||
from sglang.srt.mem_cache.utils import (
|
||
get_mla_kv_buffer_triton,
|
||
maybe_init_custom_mem_pool,
|
||
set_mla_kv_buffer_dcp_sharded_triton,
|
||
set_mla_kv_buffer_triton,
|
||
set_mla_kv_buffer_triton_fp8_quant,
|
||
set_mla_kv_scale_buffer_triton,
|
||
)
|
||
from sglang.srt.platforms import current_platform
|
||
from sglang.srt.runtime_context import get_parallel
|
||
from sglang.srt.utils import (
|
||
cpu_has_amx_support,
|
||
is_cpu,
|
||
is_cuda,
|
||
is_float4_e2m1fn_x2,
|
||
is_hip,
|
||
is_npu,
|
||
is_xpu,
|
||
next_power_of_2,
|
||
)
|
||
from sglang.srt.utils.async_probe import (
|
||
maybe_detect_kernel_facing_loc,
|
||
maybe_detect_oob,
|
||
)
|
||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||
|
||
if TYPE_CHECKING:
|
||
from sglang.srt.managers.cache_controller import LayerDoneCounter
|
||
from sglang.srt.managers.schedule_batch import Req
|
||
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# Debug-only invariant in the Mamba slot-donation path calls tensor.item(), which
|
||
# forces a per-request cudaStreamSynchronize on the scheduler thread and can stall
|
||
# the scheduler under load. Off by default; set SGLANG_MAMBA_DEBUG_ASSERTS=1 to
|
||
# re-enable for debugging.
|
||
_MAMBA_DEBUG_ASSERTS = os.environ.get("SGLANG_MAMBA_DEBUG_ASSERTS", "0") == "1"
|
||
|
||
GB = 1024 * 1024 * 1024
|
||
_is_cuda = is_cuda()
|
||
_is_npu = is_npu()
|
||
_is_cpu = is_cpu()
|
||
_cpu_has_amx_support = cpu_has_amx_support()
|
||
_is_hip = is_hip()
|
||
_is_fp8_fnuz = is_fp8_fnuz()
|
||
# `SGLANG_AITER_KV_CACHE_LAYOUT` is only meaningful on the ROCm AITER backend
|
||
# (HIP + --enable-aiter / SGLANG_USE_AITER=1). On any other platform / backend
|
||
# the SHUFFLE 5D pool layout has no consumer kernels, so the env var is
|
||
# silently ignored and the legacy NHD layout is used.
|
||
_use_aiter = bool(envs.SGLANG_USE_AITER.get()) and _is_hip
|
||
|
||
|
||
def conv_window_dedup_enabled(
|
||
is_npu: bool, is_cpu: bool, speculative_eagle_topk: Optional[int], is_kda: bool
|
||
) -> bool:
|
||
"""Whether the deduplicated sliding-window conv-intermediate layout is safe.
|
||
|
||
It is safe for CUDA linear draft chains whose kernels consume the window raw.
|
||
Tree verify, NPU/CPU, and KDA keep dense windows: tree ancestors need independent
|
||
windows, platform kernels expect contiguous steps, and KDA transposes the window
|
||
before conv so the overlapping ``as_strided`` layout would corrupt stores.
|
||
"""
|
||
return (
|
||
not is_npu
|
||
and not is_cpu
|
||
and not is_kda
|
||
and (speculative_eagle_topk is None or speculative_eagle_topk <= 1)
|
||
)
|
||
|
||
|
||
def get_tensor_size_bytes(t: Union[torch.Tensor, List[torch.Tensor]]):
|
||
if isinstance(t, list):
|
||
return sum(get_tensor_size_bytes(x) for x in t)
|
||
return np.prod(t.shape) * t.dtype.itemsize
|
||
|
||
|
||
def _set_kv_buffer_impl(
|
||
k: torch.Tensor,
|
||
v: torch.Tensor,
|
||
k_cache: torch.Tensor,
|
||
v_cache: torch.Tensor,
|
||
indices: torch.Tensor,
|
||
row_dim: int, # head_num * head_dim
|
||
store_dtype: torch.dtype,
|
||
device_module: Any,
|
||
size_limit: int,
|
||
alt_stream: Optional[torch.cuda.Stream] = None,
|
||
v_row_dim: Optional[int] = None, # head_num * v_head_dim; defaults to row_dim
|
||
) -> None:
|
||
v_row_dim = row_dim if v_row_dim is None else v_row_dim
|
||
row_bytes = row_dim * store_dtype.itemsize
|
||
v_row_bytes = v_row_dim * store_dtype.itemsize
|
||
if (_is_cuda or _is_hip) and can_use_store_cache(row_bytes, v_row_bytes):
|
||
return store_cache(
|
||
k.view(-1, row_dim),
|
||
v.view(-1, v_row_dim),
|
||
k_cache.view(-1, row_dim),
|
||
v_cache.view(-1, v_row_dim),
|
||
indices,
|
||
row_bytes=row_bytes,
|
||
v_row_bytes=v_row_bytes,
|
||
size_limit=size_limit,
|
||
)
|
||
|
||
# store_cache_cpu takes a single row_dim for both K and V, so it only serves
|
||
# equal-width rows; asymmetric KV falls through to the naive path below.
|
||
if _is_cpu and _cpu_has_amx_support and v_row_dim == row_dim:
|
||
return torch.ops.sgl_kernel.store_cache_cpu(
|
||
k,
|
||
v,
|
||
k_cache,
|
||
v_cache,
|
||
indices,
|
||
row_dim,
|
||
)
|
||
|
||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||
|
||
if get_is_capture_mode() and alt_stream is not None:
|
||
current_stream = device_module.current_stream()
|
||
alt_stream.wait_stream(current_stream)
|
||
k_cache[indices] = k
|
||
with device_module.stream(alt_stream):
|
||
v_cache[indices] = v
|
||
current_stream.wait_stream(alt_stream)
|
||
else: # fallback to naive implementation
|
||
k_cache[indices] = k
|
||
v_cache[indices] = v
|
||
|
||
|
||
def _set_kv_buffer_prefix_valid_impl(
|
||
k: torch.Tensor,
|
||
v: torch.Tensor,
|
||
k_cache: torch.Tensor,
|
||
v_cache: torch.Tensor,
|
||
loc_2d: torch.Tensor,
|
||
commit_lens: torch.Tensor,
|
||
row_dim: int,
|
||
store_dtype: torch.dtype,
|
||
) -> None:
|
||
if k.numel() == 0 or loc_2d.numel() == 0 or commit_lens.numel() == 0:
|
||
return
|
||
|
||
if not k.is_contiguous():
|
||
k = k.contiguous()
|
||
if not v.is_contiguous():
|
||
v = v.contiguous()
|
||
if not loc_2d.is_contiguous():
|
||
loc_2d = loc_2d.contiguous()
|
||
if not commit_lens.is_contiguous():
|
||
commit_lens = commit_lens.contiguous()
|
||
|
||
row_bytes = row_dim * store_dtype.itemsize
|
||
if row_bytes <= 0:
|
||
return
|
||
|
||
if row_bytes >= 8192:
|
||
bytes_per_tile = 512
|
||
num_warps = 8
|
||
elif row_bytes >= 4096:
|
||
bytes_per_tile = 256
|
||
num_warps = 4
|
||
else:
|
||
bytes_per_tile = 128
|
||
num_warps = 4
|
||
|
||
grid = (
|
||
int(loc_2d.shape[0]),
|
||
int(loc_2d.shape[1]),
|
||
triton.cdiv(row_bytes, bytes_per_tile),
|
||
)
|
||
|
||
set_kv_buffer_prefix_valid_tiled[grid](
|
||
k,
|
||
v,
|
||
k_cache,
|
||
v_cache,
|
||
loc_2d,
|
||
commit_lens,
|
||
int(k.stride(0) * k.element_size()),
|
||
int(v.stride(0) * v.element_size()),
|
||
int(k_cache.stride(0) * k_cache.element_size()),
|
||
int(v_cache.stride(0) * v_cache.element_size()),
|
||
int(loc_2d.shape[1]),
|
||
ROW_BYTES=row_bytes,
|
||
BYTES_PER_TILE=bytes_per_tile,
|
||
num_warps=num_warps,
|
||
num_stages=2,
|
||
)
|
||
|
||
|
||
class ReqToTokenPool:
|
||
"""A memory pool that maps a request to its token locations."""
|
||
|
||
enable_mamba_extra_buffer_lazy: bool = False
|
||
# Class default: some decode pools borrow another __init__ (see
|
||
# DecodeReqToTokenPool) but inherit alloc_rows.
|
||
_on_alloc_rows: Optional[Callable[[List[int]], None]] = None
|
||
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
max_context_len: int,
|
||
device: str,
|
||
enable_memory_saver: bool,
|
||
):
|
||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||
enable=enable_memory_saver
|
||
)
|
||
|
||
self.size = size
|
||
# +1 padding row at index 0: cuda-graph padded batches default
|
||
# req_pool_indices to 0, so dummy reads/writes land here harmlessly.
|
||
self._alloc_size = size + 1
|
||
self.max_context_len = max_context_len
|
||
self.device = device
|
||
with memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||
self.req_to_token = torch.zeros(
|
||
(self._alloc_size, max_context_len), dtype=torch.int32, device=device
|
||
)
|
||
self.free_slots = list(range(1, self._alloc_size))
|
||
self.req_generation = torch.zeros(self._alloc_size, dtype=torch.int64)
|
||
self._aux_cache: Any = None
|
||
|
||
def write(self, indices, values):
|
||
self.req_to_token[indices] = values
|
||
|
||
def available_size(self):
|
||
return len(self.free_slots)
|
||
|
||
def alloc(self, reqs: list[Req]) -> Optional[List[int]]:
|
||
# Indices of reqs that already have a req_pool_idx and will reuse
|
||
# their existing slot (e.g. chunked prefill continuing across chunks).
|
||
reusing = [i for i, r in enumerate(reqs) if r.kv.holds_kv]
|
||
assert all(reqs[i].kv.kv_allocated_len > 0 for i in reusing), (
|
||
"a reused row must carry allocated KV"
|
||
)
|
||
|
||
select_index = self.alloc_rows(len(reqs) - len(reusing))
|
||
if select_index is None:
|
||
return None
|
||
offset = 0
|
||
for r in reqs:
|
||
if not r.kv.holds_kv:
|
||
r.kv.req_pool_idx = select_index[offset]
|
||
offset += 1
|
||
return [r.kv.req_pool_idx for r in reqs]
|
||
|
||
def alloc_rows(self, need_size: int) -> Optional[List[int]]:
|
||
"""Take need_size rows and bump their generation, with no Req bound to
|
||
them. alloc() layers Req binding on top; beam member rows have no Req."""
|
||
if need_size > len(self.free_slots):
|
||
return None
|
||
if need_size == 0:
|
||
# Handled separately: free_slots[-0:] is the entire list, not [].
|
||
return []
|
||
# Pop from the tail: O(need_size), unlike a prefix pop which is
|
||
# O(len(free_slots)).
|
||
select_index = self.free_slots[-need_size:]
|
||
del self.free_slots[-need_size:]
|
||
self.req_generation[select_index] += 1
|
||
if self._on_alloc_rows is not None:
|
||
self._on_alloc_rows(select_index)
|
||
return select_index
|
||
|
||
def free_rows(self, indices: List[int]) -> None:
|
||
# Per-row, so beam member rows release their aux entries too: they reach
|
||
# alloc_aux_to_lengths via the decode batch's req_pool_indices_cpu.
|
||
if self._aux_cache is not None:
|
||
for index in indices:
|
||
self._aux_cache.free(index)
|
||
self.free_slots.extend(indices)
|
||
|
||
def free(self, req: Req):
|
||
assert req.kv.holds_kv, "request must have req_pool_idx"
|
||
self.free_rows([req.kv.req_pool_idx])
|
||
req.kv.req_pool_idx = None
|
||
|
||
def clear(self):
|
||
self.free_slots = list(range(1, self._alloc_size))
|
||
self.req_generation.zero_()
|
||
if self._aux_cache is not None:
|
||
self._aux_cache.clear()
|
||
|
||
def attach_aux_cache(self, aux_cache: Any) -> None:
|
||
assert self._aux_cache is None
|
||
self._aux_cache = aux_cache
|
||
|
||
def register_on_alloc_rows(self, hook: Callable[[List[int]], None]) -> None:
|
||
assert self._on_alloc_rows is None
|
||
self._on_alloc_rows = hook
|
||
|
||
def reset_aux_cache_allocator(self) -> None:
|
||
if self._aux_cache is not None:
|
||
self._aux_cache.reset_allocator()
|
||
|
||
def schedulable_token_capacity(self, physical_capacity: int) -> int:
|
||
if self._aux_cache is None:
|
||
return physical_capacity
|
||
return self._aux_cache.dense_capacity
|
||
|
||
def alloc_aux_to_lengths(
|
||
self,
|
||
*,
|
||
req_pool_indices_cpu: torch.Tensor,
|
||
target_seq_lens_cpu: torch.Tensor,
|
||
) -> None:
|
||
if self._aux_cache is not None:
|
||
self._aux_cache.alloc_to_lengths(
|
||
req_pool_indices_cpu=req_pool_indices_cpu,
|
||
target_seq_lens_cpu=target_seq_lens_cpu,
|
||
)
|
||
|
||
|
||
class MambaPool:
|
||
# Axis of each two-dimensional conv state that represents the sliding window.
|
||
# Upstream states use (dim, K-1); subclasses may preserve another layout.
|
||
conv_window_axis = -1
|
||
|
||
# Slot-lifecycle side states (see ple_state_pool.SlotIndexedState);
|
||
# class-level default because UnifiedMambaPool skips MambaPool.__init__.
|
||
_slot_siblings: Tuple = ()
|
||
|
||
def register_slot_state(self, state) -> None:
|
||
"""Attach a state that rides along on clear / copy / host round-trip,
|
||
so a slot never changes owner with a stale sibling row attached."""
|
||
self._slot_siblings = [*self._slot_siblings, state]
|
||
|
||
@dataclass(frozen=True, kw_only=True)
|
||
class State:
|
||
conv: List[torch.Tensor]
|
||
temporal: torch.Tensor
|
||
# GDN ReplaySSM ring buffers (slice 1a). Only allocated when
|
||
# `--enable-linear-replayssm` is set; otherwise None so the legacy path is
|
||
# byte-identical. Per-layer layout: [num_layers, num_slots, ...].
|
||
# replayssm_d: [num_layers, num_slots, HV, L, V]
|
||
# replayssm_k: [num_layers, num_slots, H, L, K]
|
||
# replayssm_g: [num_layers, num_slots, HV, L] (fp32)
|
||
# Under GDN spec verify, rawv/rawk hold compact D/K low parts and beta is
|
||
# None. KDA uses the same fields for its raw-input fold window.
|
||
replayssm_d: Optional[torch.Tensor] = None
|
||
replayssm_k: Optional[torch.Tensor] = None
|
||
replayssm_g: Optional[torch.Tensor] = None
|
||
replayssm_rawv: Optional[torch.Tensor] = None
|
||
replayssm_rawk: Optional[torch.Tensor] = None
|
||
replayssm_beta: Optional[torch.Tensor] = None
|
||
|
||
def at_layer_idx(self, layer: int):
|
||
kwargs = {}
|
||
# Use fields instead of vars to avoid torch.compile graph break
|
||
for f in fields(self):
|
||
name = f.name
|
||
v = getattr(self, name)
|
||
if v is None:
|
||
kwargs[name] = None
|
||
elif name in ("conv", "intermediate_conv_window"):
|
||
kwargs[name] = [conv[layer] for conv in v]
|
||
else:
|
||
kwargs[name] = v[layer]
|
||
|
||
return type(self)(**kwargs)
|
||
|
||
def mem_usage_bytes(self):
|
||
return sum(
|
||
get_tensor_size_bytes(getattr(self, f.name))
|
||
for f in dataclasses.fields(self)
|
||
if getattr(self, f.name) is not None
|
||
)
|
||
|
||
@dataclass(frozen=True, kw_only=True)
|
||
class SpeculativeState(State):
|
||
# None under --enable-linear-replayssm-spec: the spec ring owns rollback
|
||
# (verify writes ring records, commit moves cursors), so the per-draft
|
||
# full-state snapshots are never produced or consumed.
|
||
intermediate_ssm: Optional[torch.Tensor]
|
||
intermediate_conv_window: List[torch.Tensor]
|
||
|
||
def _detect_conv_window_axis(
|
||
self, conv_state_shape: List[Tuple[int, int]], win_len: int
|
||
) -> int:
|
||
"""Prefer GDN's trailing axis when both match; mixed layer layouts cannot
|
||
share one overlapping conv-window buffer.
|
||
"""
|
||
axis = None
|
||
for conv_shape in conv_state_shape:
|
||
if conv_shape[-1] == win_len:
|
||
shape_axis = len(conv_shape) - 1
|
||
elif conv_shape[0] == win_len:
|
||
shape_axis = 0
|
||
else:
|
||
raise ValueError(
|
||
f"conv_state shape {conv_shape} has no axis of length "
|
||
f"conv_kernel-1={win_len}; cannot build the deduplicated "
|
||
"sliding-window conv-intermediate view."
|
||
)
|
||
if axis is None:
|
||
axis = shape_axis
|
||
elif axis != shape_axis:
|
||
raise ValueError(
|
||
"inconsistent conv-window axis across conv shapes "
|
||
f"{conv_state_shape}; a single conv_window_axis cannot serve "
|
||
"mixed layouts."
|
||
)
|
||
return axis
|
||
|
||
def _allocate_deduplicated_conv_window(
|
||
self,
|
||
*,
|
||
conv_shape: Tuple[int, int],
|
||
num_mamba_layers: int,
|
||
spec_state_size: int,
|
||
speculative_num_draft_tokens: int,
|
||
conv_dtype: torch.dtype,
|
||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
window_axis = self.conv_window_axis % len(conv_shape)
|
||
win = conv_shape[window_axis]
|
||
physical_conv_shape = list(conv_shape)
|
||
physical_conv_shape[window_axis] = speculative_num_draft_tokens + win - 1
|
||
phys = torch.zeros(
|
||
(
|
||
num_mamba_layers,
|
||
spec_state_size + 1,
|
||
*physical_conv_shape,
|
||
),
|
||
dtype=conv_dtype,
|
||
device=self.device,
|
||
)
|
||
physical_conv_strides = phys.stride()[2:]
|
||
window_stride = physical_conv_strides[window_axis]
|
||
view = phys.as_strided(
|
||
(
|
||
phys.shape[0],
|
||
phys.shape[1],
|
||
speculative_num_draft_tokens,
|
||
*conv_shape,
|
||
),
|
||
(
|
||
phys.stride(0),
|
||
phys.stride(1),
|
||
window_stride,
|
||
*physical_conv_strides,
|
||
),
|
||
)
|
||
return phys, view
|
||
|
||
def __init__(
|
||
self,
|
||
*,
|
||
size: int,
|
||
spec_state_size: int,
|
||
cache_params: BaseLinearStateParams,
|
||
mamba_layer_ids: List[int],
|
||
device: str,
|
||
enable_memory_saver: bool = False,
|
||
speculative_num_draft_tokens: Optional[int] = None,
|
||
speculative_eagle_topk: Optional[int] = None,
|
||
enable_linear_replayssm: bool = False,
|
||
linear_replayssm_cache_len: int = 16,
|
||
envelope_layout: bool = False,
|
||
enable_linear_replayssm_spec: bool = False,
|
||
):
|
||
conv_state_shape = cache_params.shape.conv
|
||
temporal_state_shape = cache_params.shape.temporal
|
||
conv_dtype = cache_params.dtype.conv
|
||
ssm_dtype = cache_params.dtype.temporal
|
||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||
enable=enable_memory_saver
|
||
)
|
||
num_mamba_layers = len(mamba_layer_ids)
|
||
self.mamba_layer_ids = list(mamba_layer_ids)
|
||
|
||
self.size = size
|
||
self.device = device
|
||
self.debug_memory_pool = envs.SGLANG_DEBUG_MEMORY_POOL.get()
|
||
self.enable_linear_replayssm = enable_linear_replayssm
|
||
self.linear_replayssm_cache_len = linear_replayssm_cache_len
|
||
# ReplaySSM: the decode ring (--enable-linear-replayssm) allocates the
|
||
# chunked (d, k) records + write_pos; the spec-verify flag
|
||
# (--enable-linear-replayssm-spec) uses compact replay for GDN and raw
|
||
# fold-every-commit for KDA. The shared g allocation gates on
|
||
# `_replayssm_on`.
|
||
self.enable_linear_replayssm_spec = enable_linear_replayssm_spec
|
||
self.replayssm_spec_fold = bool(
|
||
enable_linear_replayssm_spec and cache_params.is_kda
|
||
)
|
||
_replayssm_on = enable_linear_replayssm or enable_linear_replayssm_spec
|
||
|
||
# for disagg with nvlink
|
||
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
|
||
maybe_init_custom_mem_pool(device=self.device)
|
||
)
|
||
|
||
with (
|
||
self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE),
|
||
(
|
||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||
if self.enable_custom_mem_pool
|
||
else nullcontext()
|
||
),
|
||
):
|
||
if envelope_layout:
|
||
# Page-granularity envelope layout (page_size==1 for state): all
|
||
# mamba layers/slots share one contiguous byte buffer; conv and
|
||
# temporal are strided views into it (see mem_cache/layout/
|
||
# page_major.py). Only the standard CUDA Triton path is supported.
|
||
assert not _is_npu and not (_is_cpu and _cpu_has_amx_support), (
|
||
"envelope_layout mamba is only supported on the CUDA path"
|
||
)
|
||
max_slots = size + 1
|
||
entry_bytes = mamba_entry_bytes(
|
||
layer_num=num_mamba_layers,
|
||
conv_state_shapes=conv_state_shape,
|
||
conv_dtype=conv_dtype,
|
||
temporal_state_shape=temporal_state_shape,
|
||
temporal_dtype=ssm_dtype,
|
||
)
|
||
self._raw = torch.zeros(
|
||
max_slots * entry_bytes, dtype=torch.uint8, device=device
|
||
)
|
||
conv_state, temporal_state = build_page_major_mamba_views(
|
||
self._raw,
|
||
layer_num=num_mamba_layers,
|
||
conv_state_shapes=conv_state_shape,
|
||
conv_dtype=conv_dtype,
|
||
temporal_state_shape=temporal_state_shape,
|
||
temporal_dtype=ssm_dtype,
|
||
max_slots=max_slots,
|
||
)
|
||
else:
|
||
conv_state = [
|
||
torch.zeros(
|
||
size=(num_mamba_layers, size + 1) + conv_shape,
|
||
dtype=conv_dtype,
|
||
device=device,
|
||
)
|
||
for conv_shape in conv_state_shape
|
||
]
|
||
|
||
if _is_npu:
|
||
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
||
_init_npu_conv_state,
|
||
)
|
||
|
||
conv_state = _init_npu_conv_state(
|
||
conv_state[0],
|
||
conv_state_shape,
|
||
speculative_num_draft_tokens,
|
||
is_kda=cache_params.is_kda,
|
||
)
|
||
|
||
if _is_cpu and _cpu_has_amx_support:
|
||
from sglang.srt.layers.amx_utils import _init_amx_conv_state
|
||
|
||
# CPU uses a different layout of conv_state for kernel optimization
|
||
conv_state = _init_amx_conv_state(conv_state)
|
||
|
||
temporal_state = torch.zeros(
|
||
size=(num_mamba_layers, size + 1) + temporal_state_shape,
|
||
dtype=ssm_dtype,
|
||
device=device,
|
||
)
|
||
|
||
# GDN ReplaySSM ring buffers (slice 1a). Allocated only when the
|
||
# flag is on; otherwise left as None so the legacy State is
|
||
# byte-identical. temporal_state_shape == (HV, V, K). Either the decode
|
||
# ring (--enable-linear-replayssm) or the spec-verify ring
|
||
# (--enable-linear-replayssm-spec) shares this allocation.
|
||
replayssm_d = replayssm_k = replayssm_g = None
|
||
replayssm_rawv = replayssm_rawk = replayssm_beta = None
|
||
if _replayssm_on:
|
||
hv, v_dim, k_dim = temporal_state_shape
|
||
h_k = getattr(cache_params.shape, "num_k_heads_per_tp", hv)
|
||
L = linear_replayssm_cache_len
|
||
# GDN speculative replay is request-lifetime scratch. Size it by
|
||
# active requests instead of every persistent radix-cache slot.
|
||
num_slots = (
|
||
spec_state_size + 1
|
||
if enable_linear_replayssm_spec and not cache_params.is_kda
|
||
else size + 1
|
||
)
|
||
# Decode records follow the SSM dtype. Spec-verify compact d/k
|
||
# records follow the activation dtype; g stays fp32.
|
||
ring_dtype = conv_dtype if enable_linear_replayssm_spec else ssm_dtype
|
||
# Fold-every-commit: one verify window, no chunked (d, k)
|
||
# records. KDA is the exception on both counts: its window
|
||
# stays L-sized (the fused verify ring-write drops
|
||
# absorb-inflated rows past L), and d/k stay allocated --
|
||
# forward_decode routes on `replayssm_d is None` (fused vs
|
||
# decode-ring), so skipping them would flip KDA decode to the
|
||
# fused path, a behavior change needing its own validation
|
||
# (memory follow-up).
|
||
if self.replayssm_spec_fold and not cache_params.is_kda:
|
||
record_len = (
|
||
speculative_num_draft_tokens
|
||
if speculative_num_draft_tokens is not None
|
||
else L
|
||
)
|
||
else:
|
||
record_len = L
|
||
if not self.replayssm_spec_fold or cache_params.is_kda:
|
||
replayssm_d = torch.zeros(
|
||
size=(num_mamba_layers, num_slots, hv, L, v_dim),
|
||
dtype=ring_dtype,
|
||
device=device,
|
||
)
|
||
replayssm_k = torch.zeros(
|
||
size=(num_mamba_layers, num_slots, h_k, L, k_dim),
|
||
dtype=ring_dtype,
|
||
device=device,
|
||
)
|
||
# The log-decay gate ring (fp32): per-head SCALAR for the GDN
|
||
# gate -> [.., record_len]; per-K VECTOR for the KDA gate ->
|
||
# [.., record_len, K] (k_dim == temporal_state_shape[-1] for both).
|
||
g_shape = (
|
||
(num_mamba_layers, num_slots, hv, record_len, k_dim)
|
||
if cache_params.is_kda
|
||
else (num_mamba_layers, num_slots, hv, record_len)
|
||
)
|
||
replayssm_g = torch.zeros(
|
||
size=g_shape,
|
||
dtype=torch.float32,
|
||
device=device,
|
||
)
|
||
# KDA still uses raw-input fold-every-commit. GDN materializes
|
||
# its compact d/k/g history directly and needs no duplicate ring.
|
||
if enable_linear_replayssm_spec and cache_params.is_kda:
|
||
if cache_params.is_kda or not self.replayssm_spec_fold:
|
||
# Backstop for the KDA ring invariants; this pool is
|
||
# sized with the final adaptive-aware draft maximum.
|
||
if L & (L - 1) != 0:
|
||
raise ValueError(
|
||
f"spec-verify ring length must be a power of two, got {L}"
|
||
)
|
||
if (
|
||
speculative_num_draft_tokens is not None
|
||
and L < 2 * speculative_num_draft_tokens
|
||
):
|
||
raise ValueError(
|
||
f"spec-verify ring too small: {L} < "
|
||
f"2 * {speculative_num_draft_tokens} (early-flush margin)"
|
||
)
|
||
replayssm_rawv = torch.zeros(
|
||
size=(num_mamba_layers, num_slots, hv, record_len, v_dim),
|
||
dtype=conv_dtype,
|
||
device=device,
|
||
)
|
||
replayssm_rawk = torch.zeros(
|
||
size=(num_mamba_layers, num_slots, h_k, record_len, k_dim),
|
||
dtype=conv_dtype,
|
||
device=device,
|
||
)
|
||
replayssm_beta = torch.zeros(
|
||
size=(num_mamba_layers, num_slots, hv, record_len),
|
||
dtype=torch.float32,
|
||
device=device,
|
||
)
|
||
elif enable_linear_replayssm_spec and ring_dtype != torch.float32:
|
||
# Low parts of compact D and normalized K. The rings follow
|
||
# the activation dtype regardless of checkpoint dtype, so
|
||
# materialization always needs both parts.
|
||
replayssm_rawv = torch.zeros(
|
||
size=(num_mamba_layers, num_slots, hv, record_len, v_dim),
|
||
dtype=conv_dtype,
|
||
device=device,
|
||
)
|
||
replayssm_rawk = torch.zeros(
|
||
size=(num_mamba_layers, num_slots, h_k, record_len, k_dim),
|
||
dtype=conv_dtype,
|
||
device=device,
|
||
)
|
||
|
||
if speculative_num_draft_tokens is not None:
|
||
if _is_npu:
|
||
temporal_state = temporal_state.transpose(-1, -2)
|
||
temporal_state_shape = (
|
||
*temporal_state_shape[:-2],
|
||
temporal_state_shape[-1],
|
||
temporal_state_shape[-2],
|
||
)
|
||
# Cache intermediate SSM states per draft token during target verify
|
||
# Shape: [num_layers, size + 1, speculative_num_draft_tokens, HV, K, V]
|
||
#
|
||
# ReplaySSM spec-verify owns rollback via the ring + cursors (the
|
||
# verify kernel never writes per-draft snapshots; the commit never
|
||
# reads them), so this buffer -- the dominant spec scratch, ~46x
|
||
# the conv state -- is dead weight there and is skipped. The conv
|
||
# intermediate windows below STAY (conv rollback consumes them).
|
||
# The recurrent-verify fallback cannot be reached under the flag
|
||
# (GDN + linear chain + triton enforced in server_args; the
|
||
# backend asserts loudly if it ever is).
|
||
# ReplaySSM skips this dominant scratch (~9GB @ K3 dspark γ=7): the
|
||
# KDA verify kernel takes intermediate_states_buffer=None (skips the
|
||
# per-step write, CACHE_INTERMEDIATE_STATES=False) and the commit
|
||
# replays the ring into the checkpoint instead. This is the memory win.
|
||
if enable_linear_replayssm_spec:
|
||
intermediate_ssm_state_cache = None
|
||
else:
|
||
intermediate_ssm_state_cache = torch.zeros(
|
||
size=(
|
||
num_mamba_layers,
|
||
spec_state_size + 1,
|
||
speculative_num_draft_tokens,
|
||
temporal_state_shape[0],
|
||
temporal_state_shape[1],
|
||
temporal_state_shape[2],
|
||
),
|
||
dtype=ssm_dtype,
|
||
device=device,
|
||
)
|
||
# Cache intermediate conv windows (last K-1 inputs) per draft token
|
||
# during target verify.
|
||
#
|
||
# On CUDA (Triton conv kernel + Triton scatter) we use a
|
||
# *deduplicated sliding-window* layout: consecutive draft tokens'
|
||
# (K-1)-wide windows overlap by (K-2), so instead of D separate
|
||
# [dim, K-1] windows we store one shared [dim, D+K-2] buffer per
|
||
# (layer, slot) and expose an overlapping `as_strided` view of
|
||
# logical shape [num_layers, size+1, draft_tokens, dim, K-1] where
|
||
# step `t`'s window is the slice shared[..., :, t:t+K-1]. This
|
||
# halves the conv-intermediate footprint (D*(K-1) -> D+K-2 columns)
|
||
# with no numerical change: both the conv kernel write (idempotent
|
||
# overlapping stores) and `fused_conv_window_scatter_with_mask`
|
||
# consume the view through its strides.
|
||
#
|
||
# Dedup the sliding-window conv-intermediate only when it is safe:
|
||
# CUDA + a linear draft chain (topk <= 1). NPU/CPU and EAGLE tree
|
||
# verify (topk > 1) keep the dense layout -- see
|
||
# `conv_window_dedup_enabled` for the full rationale. The
|
||
# `fused_conv_window_scatter_with_mask` scatter is layout-agnostic,
|
||
# so the dense fallback reads correctly through the same code path.
|
||
dedup_conv_window = (
|
||
not cache_params.shape.disable_conv_window_dedup
|
||
and conv_window_dedup_enabled(
|
||
_is_npu, _is_cpu, speculative_eagle_topk, cache_params.is_kda
|
||
)
|
||
)
|
||
self._intermediate_conv_window_phys = []
|
||
if dedup_conv_window:
|
||
win_len = cache_params.shape.conv_kernel - 1
|
||
self.conv_window_axis = self._detect_conv_window_axis(
|
||
conv_state_shape, win_len
|
||
)
|
||
intermediate_conv_window_cache = []
|
||
for conv_shape in conv_state_shape:
|
||
phys, view = self._allocate_deduplicated_conv_window(
|
||
conv_shape=conv_shape,
|
||
num_mamba_layers=num_mamba_layers,
|
||
spec_state_size=spec_state_size,
|
||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||
conv_dtype=conv_dtype,
|
||
)
|
||
self._intermediate_conv_window_phys.append(phys)
|
||
intermediate_conv_window_cache.append(view)
|
||
else:
|
||
# Original dense layout (NPU/CPU, or EAGLE tree verify): one
|
||
# [dim, K-1] window per draft token.
|
||
# Shape: [num_layers, size+1, draft_tokens, dim, K-1]
|
||
dense_conv_shapes = [
|
||
(
|
||
(conv_shape[1], conv_shape[0])
|
||
if _is_npu and cache_params.is_kda
|
||
else conv_shape
|
||
)
|
||
for conv_shape in conv_state_shape
|
||
]
|
||
intermediate_conv_window_cache = [
|
||
torch.zeros(
|
||
size=(
|
||
num_mamba_layers,
|
||
spec_state_size + 1,
|
||
speculative_num_draft_tokens,
|
||
conv_shape[0],
|
||
conv_shape[1],
|
||
),
|
||
dtype=conv_dtype,
|
||
device=device,
|
||
)
|
||
for conv_shape in dense_conv_shapes
|
||
]
|
||
self._intermediate_conv_window_phys = intermediate_conv_window_cache
|
||
self.mamba_cache = self.SpeculativeState(
|
||
conv=conv_state,
|
||
temporal=temporal_state,
|
||
intermediate_ssm=intermediate_ssm_state_cache,
|
||
intermediate_conv_window=intermediate_conv_window_cache,
|
||
replayssm_d=replayssm_d,
|
||
replayssm_k=replayssm_k,
|
||
replayssm_g=replayssm_g,
|
||
replayssm_rawv=replayssm_rawv,
|
||
replayssm_rawk=replayssm_rawk,
|
||
replayssm_beta=replayssm_beta,
|
||
)
|
||
intermediate_ssm_gb = (
|
||
get_tensor_size_bytes(intermediate_ssm_state_cache) / GB
|
||
if intermediate_ssm_state_cache is not None
|
||
else 0.0
|
||
)
|
||
logger.info(
|
||
f"Mamba Cache is allocated. "
|
||
f"max_mamba_cache_size: {size}, "
|
||
f"conv_state size: {get_tensor_size_bytes(conv_state) / GB:.2f}GB, "
|
||
f"ssm_state size: {get_tensor_size_bytes(temporal_state) / GB:.2f}GB "
|
||
f"intermediate_ssm_state_cache size: {intermediate_ssm_gb:.2f}GB "
|
||
# Report the deduplicated PHYSICAL conv-window buffers (the view
|
||
# over-reports its logical, un-deduplicated size).
|
||
f"intermediate_conv_window_cache size: {get_tensor_size_bytes(self._intermediate_conv_window_phys) / GB:.2f}GB "
|
||
)
|
||
else:
|
||
self.mamba_cache = self.State(
|
||
conv=conv_state,
|
||
temporal=temporal_state,
|
||
replayssm_d=replayssm_d,
|
||
replayssm_k=replayssm_k,
|
||
replayssm_g=replayssm_g,
|
||
replayssm_rawv=replayssm_rawv,
|
||
replayssm_rawk=replayssm_rawk,
|
||
replayssm_beta=replayssm_beta,
|
||
)
|
||
logger.info(
|
||
f"Mamba Cache is allocated. "
|
||
f"max_mamba_cache_size: {size}, "
|
||
f"conv_state size: {get_tensor_size_bytes(conv_state) / GB:.2f}GB, "
|
||
f"ssm_state size: {get_tensor_size_bytes(temporal_state) / GB:.2f}GB "
|
||
)
|
||
if _replayssm_on:
|
||
logger.info(
|
||
f"GDN ReplaySSM ring buffers allocated "
|
||
f"(record_len={record_len}, fold={self.replayssm_spec_fold}): "
|
||
f"d={get_tensor_size_bytes(replayssm_d) / GB if replayssm_d is not None else 0.0:.3f}GB, "
|
||
f"k={get_tensor_size_bytes(replayssm_k) / GB if replayssm_k is not None else 0.0:.3f}GB, "
|
||
f"g={get_tensor_size_bytes(replayssm_g) / GB:.3f}GB "
|
||
+ (
|
||
f"rawv={get_tensor_size_bytes(replayssm_rawv) / GB:.3f}GB, "
|
||
f"rawk={get_tensor_size_bytes(replayssm_rawk) / GB:.3f}GB, "
|
||
+ (
|
||
f"beta={get_tensor_size_bytes(replayssm_beta) / GB:.3f}GB "
|
||
if replayssm_beta is not None
|
||
else ""
|
||
)
|
||
if replayssm_rawv is not None
|
||
else ""
|
||
)
|
||
)
|
||
# Gate granularity of the linear-attn layers (drives the kernel's
|
||
# IS_KDA path + the g_cache layout). Read by the backend metadata to
|
||
# decide the per-K (KDA) vs scalar (GDN) flush/advance handling.
|
||
self.replayssm_is_kda = bool(_replayssm_on and cache_params.is_kda)
|
||
# Decode ReplaySSM remains keyed by persistent mamba slots.
|
||
self.replayssm_write_pos = (
|
||
torch.zeros((size + 1,), dtype=torch.int32, device=device)
|
||
if enable_linear_replayssm
|
||
else None
|
||
)
|
||
self.replayssm_spec_write_pos = (
|
||
torch.zeros((spec_state_size + 1,), dtype=torch.int32, device=device)
|
||
if enable_linear_replayssm_spec and not self.replayssm_spec_fold
|
||
else None
|
||
)
|
||
self.replayssm_cache_base = (
|
||
torch.zeros((spec_state_size + 1,), dtype=torch.int32, device=device)
|
||
if enable_linear_replayssm_spec and not self.replayssm_spec_fold
|
||
else None
|
||
)
|
||
self.replayssm_is_flush = (
|
||
torch.zeros((spec_state_size + 1,), dtype=torch.int8, device=device)
|
||
if enable_linear_replayssm_spec and not self.replayssm_spec_fold
|
||
else None
|
||
)
|
||
mem_usage_bytes = self.mamba_cache.mem_usage_bytes()
|
||
if isinstance(self.mamba_cache, self.SpeculativeState):
|
||
# `intermediate_conv_window` is an as_strided view whose logical
|
||
# shape over-reports its real footprint; charge the physical buffers
|
||
# instead. No-op for the dense layout, where the view and the
|
||
# physical tensors coincide.
|
||
mem_usage_bytes -= get_tensor_size_bytes(
|
||
self.mamba_cache.intermediate_conv_window
|
||
)
|
||
mem_usage_bytes += get_tensor_size_bytes(
|
||
self._intermediate_conv_window_phys
|
||
)
|
||
self.mem_usage = mem_usage_bytes / GB
|
||
self.num_mamba_layers = num_mamba_layers
|
||
# Full (unsharded) conv sub-block dims for PD transfer across different
|
||
# attn_tp_size (GDN: [key_dim, key_dim, value_dim]); None otherwise.
|
||
self.conv_shard_groups = getattr(cache_params.shape, "conv_shard_groups", None)
|
||
self.conv_slice_axis = getattr(cache_params.shape, "conv_slice_axis", 0)
|
||
|
||
def get_speculative_mamba2_params_all_layers(self) -> SpeculativeState:
|
||
assert isinstance(self.mamba_cache, self.SpeculativeState)
|
||
return self.mamba_cache
|
||
|
||
def mamba2_layer_cache(self, layer_id: int):
|
||
return self.mamba_cache.at_layer_idx(layer_id)
|
||
|
||
# Pool-stable (conv tensors don't move after allocation) so cached per instance
|
||
# on first use. A cached_property rather than set in __init__ because
|
||
# UnifiedMambaPool skips super().__init__.
|
||
@cached_property
|
||
def _conv_fuse_ok(self) -> bool:
|
||
"""Whether clear/copy may use the fused kernel: CUDA bf16 contiguous conv.
|
||
Strided (page-major / unified envelope) or non-bf16 conv fall back to the
|
||
per-tensor Python loop."""
|
||
convs = self.mamba_cache.conv
|
||
return (
|
||
not _is_npu
|
||
and len(convs) > 0
|
||
and convs[0].shape[0] > 0
|
||
and convs[0].is_cuda
|
||
and all(c.dtype == torch.bfloat16 and c.is_contiguous() for c in convs)
|
||
)
|
||
|
||
@cached_property
|
||
def _conv_slot_desc(self):
|
||
from sglang.srt.mem_cache.mamba_slot_fused import build_conv_slot_descriptor
|
||
|
||
return build_conv_slot_descriptor(self.mamba_cache.conv)
|
||
|
||
def _should_fuse_slot_ops(self) -> bool:
|
||
return self._conv_fuse_ok and not envs.SGLANG_DISABLE_FUSED_MAMBA_SLOT_OPS.get()
|
||
|
||
def clear_slots(self, indices: torch.Tensor):
|
||
"""Zero out mamba state at the given pool indices. Must run on forward stream."""
|
||
for sibling in self._slot_siblings:
|
||
sibling.reset_slots(indices)
|
||
if self._should_fuse_slot_ops():
|
||
from sglang.srt.mem_cache.mamba_slot_fused import fused_clear_conv_slots
|
||
|
||
fused_clear_conv_slots(self._conv_slot_desc, indices)
|
||
temporal = self.mamba_cache.temporal
|
||
if temporal.numel() > 0:
|
||
temporal[:, indices] = 0
|
||
return
|
||
if not _is_npu:
|
||
need_size = len(indices)
|
||
for i in range(len(self.mamba_cache.conv)):
|
||
t = self.mamba_cache.conv[i]
|
||
z = torch.zeros(1, dtype=t.dtype, device=t.device).expand(
|
||
t.shape[0], need_size, *t.shape[2:]
|
||
)
|
||
t[:, indices] = z
|
||
t = self.mamba_cache.temporal
|
||
z = torch.zeros(1, dtype=t.dtype, device=t.device).expand(
|
||
t.shape[0], need_size, *t.shape[2:]
|
||
)
|
||
t[:, indices] = z
|
||
else:
|
||
for i in range(len(self.mamba_cache.conv)):
|
||
t = self.mamba_cache.conv[i]
|
||
t[:, indices] = 0
|
||
t = self.mamba_cache.temporal
|
||
t[:, indices] = 0
|
||
|
||
def copy_from(self, src_indices: torch.Tensor, dst_indices: torch.Tensor):
|
||
"""Clone mamba state (conv + temporal) from src slots into dst slots.
|
||
|
||
ReplaySSM invariant: the SOURCE must be a fully-flushed checkpoint
|
||
(``write_pos[src] == 0``). Only ``temporal`` is copied, not the ring, so
|
||
an un-flushed source would drop its last ``write_pos`` updates. Callers
|
||
comply: COW copies radix checkpoints; ``cache_unfinished_req`` copies an
|
||
active slot only during prefill (ring empty); ``cache_finished_req``
|
||
caps the donate to the last flush boundary. The dst cursor is reset to 0
|
||
(the copied checkpoint has no pending ring entries).
|
||
"""
|
||
if self.replayssm_write_pos is not None and self.debug_memory_pool:
|
||
# Debug-only (syncs): catch any copy of an active, un-flushed slot.
|
||
src_wp = self.replayssm_write_pos[src_indices]
|
||
assert bool((src_wp == 0).all().item()), (
|
||
"copy_from requires a fully-flushed ReplaySSM source "
|
||
f"(write_pos==0), got {src_wp.tolist()} for src "
|
||
f"{src_indices.tolist()}"
|
||
)
|
||
if self._should_fuse_slot_ops():
|
||
from sglang.srt.mem_cache.mamba_slot_fused import fused_copy_conv_slots
|
||
|
||
if envs.SGLANG_DEBUG_MEMORY_POOL.get():
|
||
overlap = set(src_indices.tolist()) & set(dst_indices.tolist())
|
||
assert not overlap, (
|
||
"fused copy_from requires disjoint src/dst slots; "
|
||
f"overlap={sorted(overlap)}"
|
||
)
|
||
fused_copy_conv_slots(self._conv_slot_desc, src_indices, dst_indices)
|
||
temporal = self.mamba_cache.temporal
|
||
if temporal.numel() > 0:
|
||
temporal[:, dst_indices] = temporal[:, src_indices]
|
||
else:
|
||
for i in range(len(self.mamba_cache.conv)):
|
||
self.mamba_cache.conv[i][:, dst_indices] = self.mamba_cache.conv[i][
|
||
:, src_indices
|
||
]
|
||
self.mamba_cache.temporal[:, dst_indices] = self.mamba_cache.temporal[
|
||
:, src_indices
|
||
]
|
||
if self.replayssm_write_pos is not None:
|
||
self.replayssm_write_pos[dst_indices] = 0
|
||
for sibling in self._slot_siblings:
|
||
sibling.copy_slots(src_indices, dst_indices)
|
||
|
||
def get_cpu_copy(self, indices):
|
||
current_platform.synchronize()
|
||
conv_cpu = [
|
||
conv[:, indices].to("cpu", non_blocking=True)
|
||
for conv in self.mamba_cache.conv
|
||
]
|
||
temporal_cpu = self.mamba_cache.temporal[:, indices].to(
|
||
"cpu", non_blocking=True
|
||
)
|
||
siblings_cpu = [s.get_cpu_slots(indices) for s in self._slot_siblings]
|
||
current_platform.synchronize()
|
||
if self._slot_siblings:
|
||
return conv_cpu, temporal_cpu, siblings_cpu
|
||
return conv_cpu, temporal_cpu
|
||
|
||
def load_cpu_copy(self, mamba_cache_cpu, indices):
|
||
# The trailing element exists exactly when this instance registered siblings:
|
||
# the pool that saved the copy is the pool that loads it.
|
||
siblings_cpu = None
|
||
if self._slot_siblings:
|
||
siblings_cpu = mamba_cache_cpu[-1]
|
||
mamba_cache_cpu = mamba_cache_cpu[:-1]
|
||
# Accept historical 3-tuples, but request-keyed replay scratch is not
|
||
# restored with a physical checkpoint slot.
|
||
if len(mamba_cache_cpu) == 3:
|
||
conv_cpu, temporal_cpu, _ = mamba_cache_cpu
|
||
else:
|
||
conv_cpu, temporal_cpu = mamba_cache_cpu
|
||
current_platform.synchronize()
|
||
for i, conv in enumerate(self.mamba_cache.conv):
|
||
conv[:, indices] = conv_cpu[i].to(conv.device, non_blocking=True)
|
||
self.mamba_cache.temporal[:, indices] = temporal_cpu.to(
|
||
self.mamba_cache.temporal.device, non_blocking=True
|
||
)
|
||
if siblings_cpu is not None:
|
||
for sibling, data in zip(self._slot_siblings, siblings_cpu):
|
||
sibling.load_cpu_slots(data, indices)
|
||
current_platform.synchronize()
|
||
|
||
_NON_TRANSFER_STATE_FIELDS = frozenset(
|
||
{
|
||
"intermediate_ssm",
|
||
"intermediate_conv_window",
|
||
"replayssm_d",
|
||
"replayssm_k",
|
||
"replayssm_g",
|
||
"replayssm_rawv",
|
||
"replayssm_rawk",
|
||
"replayssm_beta",
|
||
}
|
||
)
|
||
|
||
def _iter_transfer_state_entries(self):
|
||
"""Yield ``[slot, ...]`` state entries and their transfer metadata."""
|
||
for field, value in vars(self.mamba_cache).items():
|
||
if field in self._NON_TRANSFER_STATE_FIELDS or value is None:
|
||
continue
|
||
tensors = value if isinstance(value, list) else [value]
|
||
slice_axis = self.conv_slice_axis if field == "conv" else 0
|
||
for state_tensor in tensors:
|
||
# A ShortConv layer has no temporal state, so that buffer is
|
||
# empty. Advertising it fails the whole batch registration.
|
||
if state_tensor.numel() == 0:
|
||
continue
|
||
for layer_index, layer_id in enumerate(self.mamba_layer_ids):
|
||
yield field, state_tensor[layer_index], slice_axis, layer_id
|
||
|
||
for sibling in self._slot_siblings:
|
||
yield from sibling.iter_transfer_state_entries()
|
||
|
||
def get_contiguous_buf_infos(self):
|
||
"""Get transferable state buffer information for RDMA registration."""
|
||
data_ptrs, data_lens, item_lens = [], [], []
|
||
|
||
for _, state_tensor, _, _ in self._iter_transfer_state_entries():
|
||
data_ptrs.append(state_tensor.data_ptr())
|
||
data_lens.append(state_tensor.nbytes)
|
||
item_lens.append(state_tensor[0].nbytes)
|
||
return data_ptrs, data_lens, item_lens
|
||
|
||
def get_state_dim_per_tensor(self):
|
||
"""Get the sliceable dimension size for each state tensor.
|
||
|
||
The slice axis is tensor-specific: normally the first per-slot axis,
|
||
while Kimi conv state uses the second per-slot axis.
|
||
"""
|
||
dim_per_tensor = []
|
||
for _, state_tensor, slice_axis, _ in self._iter_transfer_state_entries():
|
||
# Zero is a protocol marker for request state replicated across the
|
||
# attention-TP group. Heterogeneous PD copies the whole item from one
|
||
# elected source rank instead of slicing it as a TP-sharded tensor.
|
||
if slice_axis is None:
|
||
dim_per_tensor.append(0)
|
||
continue
|
||
# state_tensor shape: [size+1, sliceable_dim, ...]. Kimi conv state
|
||
# transposes the two per-slot axes to [K-1, dim].
|
||
axis = 1 + slice_axis
|
||
dim_per_tensor.append(state_tensor.shape[axis])
|
||
return dim_per_tensor
|
||
|
||
def get_state_layer_ids(self):
|
||
"""Global model-layer id for each RDMA state entry.
|
||
|
||
Aligned element-wise with get_contiguous_buf_infos(), which flattens
|
||
the state list tensor-major x layer. Lets PD transfer match entries
|
||
by layer id when prefill (PP stage) holds a subset of the mamba layers.
|
||
"""
|
||
return [layer_id for _, _, _, layer_id in self._iter_transfer_state_entries()]
|
||
|
||
def get_state_slice_outer_counts(self):
|
||
"""Get the number of rows preceding each tensor's TP slice axis."""
|
||
outer_counts = []
|
||
for _, state_tensor, slice_axis, _ in self._iter_transfer_state_entries():
|
||
outer_count = (
|
||
1
|
||
if slice_axis is None
|
||
else math.prod(state_tensor.shape[1 : 1 + slice_axis])
|
||
)
|
||
outer_counts.append(outer_count)
|
||
return outer_counts
|
||
|
||
def get_state_conv_shard_groups(self):
|
||
"""Per-tensor conv sub-block dims, aligned element-wise with
|
||
get_state_dim_per_tensor().
|
||
|
||
For GDN, conv_state's sliceable axis is cat([query, key, value]) with
|
||
each sub-block head-sharded independently across attn-TP; the full
|
||
(unsharded) sub-block dims are returned so PD transfer across different
|
||
attn_tp_size can slice each sub-block. Returns None for temporal_state
|
||
(single head-sharded axis) and whenever no descriptor is available, so
|
||
those tensors keep the single contiguous slice.
|
||
"""
|
||
subdims_per_tensor = []
|
||
for field, _, _, _ in self._iter_transfer_state_entries():
|
||
# Only conv_state carries a q/k/v decomposition.
|
||
subdims = (
|
||
list(self.conv_shard_groups)
|
||
if field == "conv" and self.conv_shard_groups is not None
|
||
else None
|
||
)
|
||
subdims_per_tensor.append(subdims)
|
||
return subdims_per_tensor
|
||
|
||
def get_kv_size_bytes(self):
|
||
return self.mamba_cache.mem_usage_bytes()
|
||
|
||
|
||
class HybridReqToTokenPool(ReqToTokenPool):
|
||
"""A memory pool that maps a request to its token locations."""
|
||
|
||
mamba_pool_cls = MambaPool
|
||
|
||
def __init__(
|
||
self,
|
||
*,
|
||
size: int,
|
||
mamba_size: int,
|
||
mamba_spec_state_size: int,
|
||
max_context_len: int,
|
||
device: str,
|
||
enable_memory_saver: bool,
|
||
cache_params: BaseLinearStateParams,
|
||
mamba_layer_ids: List[int],
|
||
enable_mamba_extra_buffer: bool,
|
||
enable_mamba_extra_buffer_lazy: bool = False,
|
||
speculative_num_draft_tokens: int = None,
|
||
speculative_eagle_topk: Optional[int] = None,
|
||
enable_overlap_schedule: bool = True,
|
||
start_layer: Optional[int] = None,
|
||
enable_linear_replayssm: bool = False,
|
||
linear_replayssm_cache_len: int = 16,
|
||
mamba_envelope_layout: bool = False,
|
||
enable_linear_replayssm_spec: bool = False,
|
||
short_conv_layer_ids: Optional[List[int]] = None,
|
||
short_conv_state_shape: Optional[Tuple[int, int]] = None,
|
||
ngram_context_len: int = 0,
|
||
ngram_eos_token_id: int = 0,
|
||
):
|
||
super().__init__(
|
||
size=size,
|
||
max_context_len=max_context_len,
|
||
device=device,
|
||
enable_memory_saver=enable_memory_saver,
|
||
)
|
||
|
||
self.mamba_ping_pong_track_buffer_size = 2 if enable_overlap_schedule else 1
|
||
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
|
||
self.enable_mamba_extra_buffer_lazy = enable_mamba_extra_buffer_lazy
|
||
self.enable_memory_saver = enable_memory_saver
|
||
self.start_layer = start_layer if start_layer is not None else 0
|
||
self.layer_transfer_counter = None
|
||
self.ple_window_cache = None
|
||
self._init_mamba_pool(
|
||
mamba_size=mamba_size,
|
||
mamba_spec_state_size=mamba_spec_state_size,
|
||
cache_params=cache_params,
|
||
mamba_layer_ids=mamba_layer_ids,
|
||
device=device,
|
||
enable_mamba_extra_buffer=enable_mamba_extra_buffer,
|
||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||
speculative_eagle_topk=speculative_eagle_topk,
|
||
enable_linear_replayssm=enable_linear_replayssm,
|
||
linear_replayssm_cache_len=linear_replayssm_cache_len,
|
||
mamba_envelope_layout=mamba_envelope_layout,
|
||
enable_linear_replayssm_spec=enable_linear_replayssm_spec,
|
||
short_conv_layer_ids=short_conv_layer_ids,
|
||
short_conv_state_shape=short_conv_state_shape,
|
||
ngram_context_len=ngram_context_len,
|
||
ngram_eos_token_id=ngram_eos_token_id,
|
||
)
|
||
|
||
def _init_mamba_pool(
|
||
self,
|
||
mamba_size: int,
|
||
mamba_spec_state_size: int,
|
||
cache_params: BaseLinearStateParams,
|
||
mamba_layer_ids: List[int],
|
||
device: str,
|
||
enable_mamba_extra_buffer: bool,
|
||
speculative_num_draft_tokens: int = None,
|
||
speculative_eagle_topk: Optional[int] = None,
|
||
enable_linear_replayssm: bool = False,
|
||
linear_replayssm_cache_len: int = 16,
|
||
mamba_envelope_layout: bool = False,
|
||
enable_linear_replayssm_spec: bool = False,
|
||
short_conv_layer_ids: Optional[List[int]] = None,
|
||
short_conv_state_shape: Optional[Tuple[int, int]] = None,
|
||
ngram_context_len: int = 0,
|
||
ngram_eos_token_id: int = 0,
|
||
):
|
||
self.mamba_pool = self.mamba_pool_cls(
|
||
size=mamba_size,
|
||
spec_state_size=mamba_spec_state_size,
|
||
cache_params=cache_params,
|
||
mamba_layer_ids=mamba_layer_ids,
|
||
device=device,
|
||
enable_memory_saver=self.enable_memory_saver,
|
||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||
speculative_eagle_topk=speculative_eagle_topk,
|
||
enable_linear_replayssm=enable_linear_replayssm,
|
||
linear_replayssm_cache_len=linear_replayssm_cache_len,
|
||
envelope_layout=mamba_envelope_layout,
|
||
enable_linear_replayssm_spec=enable_linear_replayssm_spec,
|
||
)
|
||
self.mamba_allocator = MambaSlotAllocator(
|
||
size=mamba_size,
|
||
device=device,
|
||
)
|
||
self.mamba_map = {layer_id: i for i, layer_id in enumerate(mamba_layer_ids)}
|
||
|
||
# Qwen4-Exp PLE side states; built disabled rather than None without a config,
|
||
# so every hybrid model has both attributes.
|
||
from sglang.srt.mem_cache.ple_state_pool import NGramPool, ShortConvPool
|
||
|
||
self.short_conv_pool = ShortConvPool(
|
||
size=mamba_size,
|
||
spec_state_size=mamba_spec_state_size,
|
||
state_shape=short_conv_state_shape,
|
||
layer_ids=short_conv_layer_ids or [],
|
||
dtype=cache_params.dtype.conv,
|
||
device=device,
|
||
enable_memory_saver=self.enable_memory_saver,
|
||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||
)
|
||
self.ngram_pool = NGramPool(
|
||
size=mamba_size,
|
||
spec_state_size=mamba_spec_state_size,
|
||
context_len=ngram_context_len,
|
||
eos_token_id=ngram_eos_token_id,
|
||
device=device,
|
||
enable_memory_saver=self.enable_memory_saver,
|
||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||
)
|
||
# Disabled pools stay off the sibling list so the host-offload payload
|
||
# keeps its legacy shape for every non-PLE hybrid model.
|
||
if self.short_conv_pool.enabled:
|
||
self.mamba_pool.register_slot_state(self.short_conv_pool)
|
||
if self.ngram_pool.enabled:
|
||
self.mamba_pool.register_slot_state(self.ngram_pool)
|
||
|
||
# Optional int8 checkpoint pool: the radix caches states here (int8) instead
|
||
# of holding them in the active bf16 pool -> ~2x cached-prefix capacity at
|
||
# fixed memory. Strategy-agnostic (no_buffer / extra_buffer / spec).
|
||
from sglang.srt.mem_cache.mamba_checkpoint_pool import (
|
||
maybe_init_int8_mamba_checkpoint_pool,
|
||
)
|
||
|
||
self.mamba_ckpt_pool = maybe_init_int8_mamba_checkpoint_pool(
|
||
mamba_size=mamba_size,
|
||
cache_params=cache_params,
|
||
mamba_layer_ids=mamba_layer_ids,
|
||
device=device,
|
||
)
|
||
if self.mamba_ckpt_pool is not None and (
|
||
self.short_conv_pool.enabled or self.ngram_pool.enabled
|
||
):
|
||
# The int8 checkpoint pool frees the bf16 slot after donating its state,
|
||
# taking the bf16-slot-indexed PLE side states with it.
|
||
raise ValueError(
|
||
"--enable-int8-mamba-checkpoint is incompatible with Qwen4-Exp "
|
||
"PLE side states"
|
||
)
|
||
|
||
self.device = device
|
||
req_pool_size = self.req_to_token.shape[0]
|
||
self.req_index_to_mamba_index_mapping: torch.Tensor = torch.zeros(
|
||
req_pool_size, dtype=torch.int32, device=self.device
|
||
)
|
||
if enable_mamba_extra_buffer:
|
||
self.req_index_to_mamba_ping_pong_track_buffer_mapping: torch.Tensor = (
|
||
torch.zeros(
|
||
(req_pool_size, self.mamba_ping_pong_track_buffer_size),
|
||
dtype=torch.int64,
|
||
device=self.device,
|
||
)
|
||
)
|
||
|
||
def clone_with_new_mamba(
|
||
self,
|
||
*,
|
||
mamba_size: int,
|
||
mamba_spec_state_size: int,
|
||
cache_params: BaseLinearStateParams,
|
||
device: str,
|
||
enable_mamba_extra_buffer: bool,
|
||
draft_model_idx: int,
|
||
speculative_num_draft_tokens: int = None,
|
||
speculative_eagle_topk: Optional[int] = None,
|
||
) -> HybridReqToTokenPool:
|
||
"""Shallow copy that shares the req_to_token mapping but owns a fresh mamba
|
||
pool keyed on a single draft layer. Used by multi-layer EAGLE draft workers:
|
||
each draft head shares the target's request-to-token mapping but needs its
|
||
own sconv/mamba cache at layer_id=draft_model_idx.
|
||
"""
|
||
clone = copy.copy(self)
|
||
clone._init_mamba_pool(
|
||
mamba_size=mamba_size,
|
||
mamba_spec_state_size=mamba_spec_state_size,
|
||
cache_params=cache_params,
|
||
mamba_layer_ids=[draft_model_idx],
|
||
device=device,
|
||
enable_mamba_extra_buffer=enable_mamba_extra_buffer,
|
||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||
speculative_eagle_topk=speculative_eagle_topk,
|
||
)
|
||
clone.req_index_to_mamba_index_mapping = self.req_index_to_mamba_index_mapping
|
||
if enable_mamba_extra_buffer:
|
||
clone.req_index_to_mamba_ping_pong_track_buffer_mapping = (
|
||
self.req_index_to_mamba_ping_pong_track_buffer_mapping
|
||
)
|
||
return clone
|
||
|
||
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
|
||
self.layer_transfer_counter = layer_transfer_counter
|
||
|
||
# For chunk prefill req, we do not need to allocate mamba cache,
|
||
# We could use allocated mamba cache instead.
|
||
def alloc(self, reqs: List[Req]) -> Optional[List[int]]:
|
||
fresh_req_rows = [req.kv.req_pool_idx is None for req in reqs]
|
||
select_index = super().alloc(reqs)
|
||
if select_index is None:
|
||
return None
|
||
|
||
spec_write_pos = getattr(self.mamba_pool, "replayssm_spec_write_pos", None)
|
||
if spec_write_pos is not None:
|
||
fresh_indices = [
|
||
idx for idx, fresh in zip(select_index, fresh_req_rows) if fresh
|
||
]
|
||
if fresh_indices:
|
||
spec_write_pos[fresh_indices] = 0
|
||
self.mamba_pool.replayssm_cache_base[fresh_indices] = 0
|
||
self.mamba_pool.replayssm_is_flush[fresh_indices] = 0
|
||
|
||
mamba_indices: list[torch.Tensor] = []
|
||
mamba_ping_pong_track_buffers: list[torch.Tensor] = []
|
||
for req in reqs:
|
||
if req.kv.holds_mamba: # for radix cache / continuing chunked
|
||
pass
|
||
else:
|
||
mid = self.mamba_allocator.alloc(1)
|
||
assert mid is not None, (
|
||
f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_allocator.available_size()=}, {len(reqs)=}"
|
||
)
|
||
req.kv.mamba_pool_idx = mid[0]
|
||
req.kv.mamba_needs_clear = True
|
||
# GDN ReplaySSM: a freshly (re)assigned slot starts an empty
|
||
# ring. write_pos=0 means "ring empty", so the decode kernel
|
||
# ignores ring contents and reads only the checkpoint state
|
||
# (the post-prefill state that prefill wrote into this slot).
|
||
if self.mamba_pool.replayssm_write_pos is not None:
|
||
self.mamba_pool.replayssm_write_pos[req.kv.mamba_pool_idx] = 0
|
||
mamba_indices.append(req.kv.mamba_pool_idx)
|
||
if self.enable_mamba_extra_buffer:
|
||
if req.kv.mamba_ping_pong_track_buffer is None:
|
||
self._alloc_ping_pong_buffer(req)
|
||
mamba_ping_pong_track_buffers.append(
|
||
req.kv.mamba_ping_pong_track_buffer
|
||
)
|
||
assert len(select_index) == len(mamba_indices), (
|
||
"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
|
||
)
|
||
if self.enable_mamba_extra_buffer:
|
||
assert len(select_index) == len(mamba_ping_pong_track_buffers), (
|
||
"Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
|
||
)
|
||
mamba_index_tensor = torch.stack(mamba_indices).to(dtype=torch.int32)
|
||
self.req_index_to_mamba_index_mapping[select_index] = mamba_index_tensor
|
||
if self.enable_mamba_extra_buffer:
|
||
ping_pong_tensor = torch.stack(mamba_ping_pong_track_buffers)
|
||
self.req_index_to_mamba_ping_pong_track_buffer_mapping[select_index] = (
|
||
ping_pong_tensor
|
||
)
|
||
return select_index
|
||
|
||
def get_mamba_indices(self, req_indices: torch.Tensor) -> torch.Tensor:
|
||
return self.req_index_to_mamba_index_mapping[req_indices]
|
||
|
||
@property
|
||
def mamba_v2p_table(self) -> Optional[torch.Tensor]:
|
||
"""The mamba virtual->physical slot table, or None when the ids this
|
||
pool hands out are already physical."""
|
||
return None
|
||
|
||
@property
|
||
def mamba_translate_is_fusable(self) -> bool:
|
||
"""Whether `fused_replay_state_indices` can reproduce this pool's
|
||
`translate_mamba_indices` in its own launch.
|
||
|
||
The kernel expresses exactly two shapes: the identity, and one gather
|
||
through `mamba_v2p_table`. A subclass that replaces the translate with
|
||
anything else is excluded here rather than silently mis-served.
|
||
"""
|
||
if self.mamba_v2p_table is not None:
|
||
return True
|
||
return (
|
||
type(self).translate_mamba_indices
|
||
is HybridReqToTokenPool.translate_mamba_indices
|
||
)
|
||
|
||
def translate_mamba_indices(self, mamba_indices: torch.Tensor) -> torch.Tensor:
|
||
"""Virtual->physical mamba-slot translate. Identity for a static pool
|
||
(slots are physical); UnifiedHybridReqToTokenPool overrides it for the
|
||
unified memory pool, where mamba slot ids are virtual. Callers translate
|
||
before calling the pool's physical-id state ops (copy_from / clear_slots
|
||
/ get_cpu_copy / load_cpu_copy)."""
|
||
return mamba_indices
|
||
|
||
def mamba2_layer_index(self, layer_id: int) -> int:
|
||
"""Pool-side index of ``layer_id``'s state, gated on its HiCache transfer.
|
||
|
||
For a caller that wants one specific state tensor: it indexes the pool
|
||
tensor itself instead of taking a ``State`` sliced over every field.
|
||
"""
|
||
assert layer_id in self.mamba_map
|
||
if self.layer_transfer_counter is not None:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
return self.mamba_map[layer_id]
|
||
|
||
def mamba2_layer_cache(self, layer_id: int):
|
||
return self.mamba_pool.mamba2_layer_cache(self.mamba2_layer_index(layer_id))
|
||
|
||
def short_conv_layer_cache(self, layer_id: int) -> torch.Tensor:
|
||
if self.layer_transfer_counter is not None:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
return self.short_conv_pool.layer_cache(layer_id)
|
||
|
||
def short_conv_layer_intermediate_cache(
|
||
self, layer_id: int
|
||
) -> Optional[torch.Tensor]:
|
||
return self.short_conv_pool.layer_intermediate_cache(layer_id)
|
||
|
||
def get_ngram_context(self, ngram_indices: torch.Tensor) -> torch.Tensor:
|
||
return self.ngram_pool.get_context(ngram_indices)
|
||
|
||
def set_ngram_context(
|
||
self, ngram_indices: torch.Tensor, context: torch.Tensor
|
||
) -> None:
|
||
self.ngram_pool.set_context(ngram_indices, context)
|
||
|
||
def set_ngram_intermediate_context(self, context: torch.Tensor) -> None:
|
||
self.ngram_pool.set_intermediate_context(context)
|
||
|
||
def copy_mamba_state(
|
||
self, src_index: torch.Tensor, dst_index: torch.Tensor
|
||
) -> None:
|
||
if src_index.numel() == 0:
|
||
return
|
||
if (
|
||
self.layer_transfer_counter is not None
|
||
and self.layer_transfer_counter.consumer_index >= 0
|
||
):
|
||
last_mamba_layer = max(self.mamba_map)
|
||
self.layer_transfer_counter.wait_until(last_mamba_layer - self.start_layer)
|
||
self.mamba_pool.copy_from(src_index, dst_index)
|
||
|
||
def get_speculative_mamba2_params_all_layers(self) -> MambaPool.SpeculativeState:
|
||
return self.mamba_pool.get_speculative_mamba2_params_all_layers()
|
||
|
||
def get_state_buf_infos(self):
|
||
return self.mamba_pool.get_contiguous_buf_infos()
|
||
|
||
def get_state_dim_per_tensor(self):
|
||
return self.mamba_pool.get_state_dim_per_tensor()
|
||
|
||
def get_state_slice_outer_counts(self):
|
||
return self.mamba_pool.get_state_slice_outer_counts()
|
||
|
||
def get_state_conv_shard_groups(self):
|
||
return self.mamba_pool.get_state_conv_shard_groups()
|
||
|
||
def get_mamba_ping_pong_other_idx(self, mamba_next_track_idx: int) -> int:
|
||
if self.mamba_ping_pong_track_buffer_size == 2:
|
||
return 1 - mamba_next_track_idx
|
||
else:
|
||
return mamba_next_track_idx
|
||
|
||
def get_mamba_ping_pong_keep_idx(self, req: Req) -> int:
|
||
"""Return the ping-pong index holding the most recent tracked state."""
|
||
return req.kv.mamba_last_track_idx
|
||
|
||
def _alloc_ping_pong_buffer(self, req: Req):
|
||
"""Allocate the ping-pong track buffer for a new request.
|
||
|
||
Lazy mode allocates 1 slot with the second set to -1 (allocated
|
||
on demand at track boundaries). Normal mode allocates all slots upfront.
|
||
"""
|
||
n = (
|
||
1
|
||
if self.enable_mamba_extra_buffer_lazy
|
||
else self.mamba_ping_pong_track_buffer_size
|
||
)
|
||
slots = self.mamba_allocator.alloc(n)
|
||
assert slots is not None, (
|
||
"Not enough space for mamba ping pong idx, "
|
||
"try to increase --mamba-full-memory-ratio."
|
||
)
|
||
buf = torch.full(
|
||
(self.mamba_ping_pong_track_buffer_size,),
|
||
-1,
|
||
dtype=slots.dtype,
|
||
device=slots.device,
|
||
)
|
||
buf[:n] = slots
|
||
req.kv.mamba_ping_pong_track_buffer = buf
|
||
req.kv.mamba_next_track_idx = 0
|
||
req.kv.mamba_last_track_idx = (
|
||
0
|
||
if self.enable_mamba_extra_buffer_lazy
|
||
else self.get_mamba_ping_pong_other_idx(0)
|
||
)
|
||
|
||
def set_mamba_ping_pong_slot(self, req: Req, idx: int, value):
|
||
"""Update a ping-pong slot value and sync the device-side mapping.
|
||
|
||
The req holds the authoritative buffer; this keeps the
|
||
req_index_to_mamba_ping_pong_track_buffer_mapping in sync so that
|
||
set_mamba_track_indices_from_reqs reads correct slot indices.
|
||
"""
|
||
req.kv.mamba_ping_pong_track_buffer[idx] = value
|
||
self.req_index_to_mamba_ping_pong_track_buffer_mapping[req.kv.req_pool_idx] = (
|
||
req.kv.mamba_ping_pong_track_buffer
|
||
)
|
||
|
||
def donate_mamba_ping_pong_slot(
|
||
self, req: Req, new_slot: torch.Tensor
|
||
) -> torch.Tensor:
|
||
"""Donate the tracked-state ping-pong slot to the radix cache.
|
||
|
||
Returns the old slot index (shape [1]) for cache insertion and
|
||
replaces it with new_slot so the request can continue tracking.
|
||
"""
|
||
donate_idx = self.get_mamba_ping_pong_keep_idx(req)
|
||
mamba_value_donated = (
|
||
req.kv.mamba_ping_pong_track_buffer[donate_idx].unsqueeze(-1).clone()
|
||
)
|
||
if _MAMBA_DEBUG_ASSERTS:
|
||
# .item() forces a cudaStreamSynchronize; only pay it when debugging.
|
||
assert mamba_value_donated.item() != -1, (
|
||
f"Donated mamba slot is -1: donate_idx={donate_idx}, "
|
||
f"buf={req.kv.mamba_ping_pong_track_buffer.tolist()}, "
|
||
f"next_track_idx={req.kv.mamba_next_track_idx}, "
|
||
f"rid={req.rid}"
|
||
)
|
||
self.set_mamba_ping_pong_slot(req, donate_idx, new_slot[0])
|
||
return mamba_value_donated
|
||
|
||
def free_mamba_cache(
|
||
self, req: Req, mamba_ping_pong_track_buffer_to_keep: Optional[int] = None
|
||
):
|
||
mamba_index = req.kv.mamba_pool_idx
|
||
assert mamba_index is not None, "double free? mamba_index is None"
|
||
self.mamba_allocator.free(mamba_index.unsqueeze(0))
|
||
req.kv.mamba_pool_idx = None
|
||
|
||
if self.enable_mamba_extra_buffer:
|
||
mamba_ping_pong_track_buffer_to_free = (
|
||
self.req_index_to_mamba_ping_pong_track_buffer_mapping[
|
||
req.kv.req_pool_idx
|
||
]
|
||
)
|
||
if mamba_ping_pong_track_buffer_to_keep is not None:
|
||
assert mamba_ping_pong_track_buffer_to_keep in [
|
||
0,
|
||
1,
|
||
], (
|
||
f"mamba_ping_pong_track_buffer_to_keep must be 0 or 1, {mamba_ping_pong_track_buffer_to_keep=}"
|
||
)
|
||
# Avoid Python-list advanced indexing on a device tensor.
|
||
# The ping-pong buffer size is either 2 (normal) or 1 (spec decode).
|
||
if self.mamba_ping_pong_track_buffer_size == 2:
|
||
idx_to_free = 1 - mamba_ping_pong_track_buffer_to_keep
|
||
mamba_ping_pong_track_buffer_to_free = (
|
||
mamba_ping_pong_track_buffer_to_free[
|
||
idx_to_free : idx_to_free + 1
|
||
]
|
||
)
|
||
else:
|
||
assert self.mamba_ping_pong_track_buffer_size == 1, (
|
||
f"Unexpected mamba_ping_pong_track_buffer_size="
|
||
f"{self.mamba_ping_pong_track_buffer_size}"
|
||
)
|
||
assert mamba_ping_pong_track_buffer_to_keep == 0, (
|
||
"mamba_ping_pong_track_buffer_to_keep must be 0 when "
|
||
"mamba_ping_pong_track_buffer_size is 1"
|
||
)
|
||
# Keep the only slot, so free nothing.
|
||
mamba_ping_pong_track_buffer_to_free = (
|
||
mamba_ping_pong_track_buffer_to_free[0:0]
|
||
)
|
||
if self.enable_mamba_extra_buffer_lazy:
|
||
mamba_ping_pong_track_buffer_to_free = (
|
||
mamba_ping_pong_track_buffer_to_free[
|
||
mamba_ping_pong_track_buffer_to_free != -1
|
||
]
|
||
)
|
||
self.mamba_allocator.free(mamba_ping_pong_track_buffer_to_free)
|
||
# Match the req.kv.mamba_pool_idx=None clear above so the next
|
||
# alloc() doesn't see a stale ping-pong reference on the req
|
||
# and skip allocation (which would silently reuse a freed
|
||
# tensor on the req side while the new pool slot leaks).
|
||
req.kv.mamba_ping_pong_track_buffer = None
|
||
req.kv.mamba_next_track_idx = None
|
||
req.kv.mamba_last_track_idx = None
|
||
req.kv.mamba_last_track_seqlen = None
|
||
req.kv.mamba_cow_src_index = None
|
||
req.kv.mamba_needs_clear = False
|
||
|
||
def clear(self):
|
||
logger.info("Reset HybridReqToTokenPool")
|
||
super().clear()
|
||
self.mamba_allocator.clear()
|
||
self.short_conv_pool.clear()
|
||
self.ngram_pool.clear()
|
||
# The int8 checkpoint pool holds radix-cached states in its own slots; a
|
||
# flush/reset drops the radix tree, so its slots must be released too,
|
||
# otherwise the (now unreferenced) slots leak and break the int8-pool
|
||
# invariant (int8_available + radix_cached != int8_total).
|
||
if self.mamba_ckpt_pool is not None:
|
||
self.mamba_ckpt_pool.clear()
|
||
self.req_index_to_mamba_index_mapping.zero_()
|
||
if self.enable_mamba_extra_buffer:
|
||
self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_()
|
||
|
||
|
||
@dataclass
|
||
class KVWriteLoc:
|
||
"""Write target(s) for ``KVCache.set_kv_buffer``.
|
||
|
||
All location info lives here (in the attention metadata), NOT in the pool:
|
||
- ``loc``: the generic per-token write location (``out_cache_loc``).
|
||
KERNEL-FACING on every pool: physical by allocation on non-unified
|
||
pools, rebound at ForwardBatch construction (``rebind_write_loc``) on
|
||
the unified pool.
|
||
- ``swa_loc``: the SWA-sub-pool location for hybrid SWA pools (``None``
|
||
otherwise); under the unified pool the translator derives it from the
|
||
same rebound loc (``sliding_window_write_loc_for``).
|
||
- ``full_loc``: OPTIONAL full-attention-sub-pool location. Since the
|
||
construction-time rebind it is the SAME id space as ``loc``, so pools
|
||
fall back to ``loc`` when it is ``None`` -- only triton's captured path
|
||
still passes its capture-stable
|
||
``ForwardMetadata.out_cache_loc_full_physical`` buffer here (a
|
||
same-space alias slated for collapse).
|
||
|
||
``swa_loc`` and ``full_loc`` are the parallel pair (each a pre-resolved
|
||
loc into its sub-pool, mirroring ``swa_kv_pool`` / ``full_kv_pool``);
|
||
``loc`` is the generic fallback. Bundling them lets a backend issue one
|
||
``set_kv_buffer`` call regardless of pool type.
|
||
"""
|
||
|
||
loc: torch.Tensor
|
||
swa_loc: Optional[torch.Tensor] = None
|
||
full_loc: Optional[torch.Tensor] = None
|
||
|
||
def __post_init__(self):
|
||
# swa_loc / full_loc are resolved once at metadata-init from the full
|
||
# (padded) out_cache_loc; piecewise/DP-padded paths later narrow loc per
|
||
# layer, so slice these pre-resolved locs to match (same per-token order).
|
||
if self.swa_loc is not None and self.swa_loc.shape[0] != self.loc.shape[0]:
|
||
self.swa_loc = self.swa_loc[: self.loc.shape[0]]
|
||
if self.full_loc is not None and self.full_loc.shape[0] != self.loc.shape[0]:
|
||
self.full_loc = self.full_loc[: self.loc.shape[0]]
|
||
|
||
|
||
def unwrap_write_loc(loc_info):
|
||
"""Return ``(loc, swa_loc, full_loc)`` from a ``KVWriteLoc`` or a bare loc."""
|
||
if isinstance(loc_info, KVWriteLoc):
|
||
return loc_info.loc, loc_info.swa_loc, loc_info.full_loc
|
||
return loc_info, None, None
|
||
|
||
|
||
class KvBufferDesc:
|
||
"""Byte-span math for one KV buffer laid out as rows of ``row_bytes`` holding
|
||
``tokens_per_row`` tokens each (a row = one token slot, or one whole page)."""
|
||
|
||
__slots__ = ("name", "shape", "row_bytes", "tokens_per_row")
|
||
|
||
def __init__(self, name: str, shape: tuple, *, row_bytes: int, tokens_per_row: int):
|
||
self.name = name
|
||
self.shape = tuple(shape)
|
||
self.row_bytes = int(row_bytes)
|
||
self.tokens_per_row = int(tokens_per_row)
|
||
|
||
def _rows(self, num_tokens: int) -> int:
|
||
n = max(int(num_tokens), 0)
|
||
return (n + self.tokens_per_row - 1) // self.tokens_per_row
|
||
|
||
def reserved_span_bytes(self, itemsize: int) -> int:
|
||
"""Full upper-bound byte size of the buffer (its whole tensor)."""
|
||
return math.prod(self.shape) * itemsize
|
||
|
||
def prefix_span_bytes(self, num_tokens: int, page_size: int) -> int:
|
||
"""Bytes to back to make the first ``num_tokens`` tokens usable."""
|
||
return self._rows(num_tokens) * self.row_bytes
|
||
|
||
def final_span_bytes(self, num_tokens: int, page_size: int) -> int:
|
||
"""Bytes of the final advertised span (adds the padded page). CEIL, not floor:
|
||
an unaligned count must still cover its partial last page (e.g. n=17, page=16
|
||
-> 3 pages, not 2)."""
|
||
return self._rows(max(int(num_tokens), 0) + page_size) * self.row_bytes
|
||
|
||
def item_len_bytes(self, page_size: int) -> int:
|
||
"""Per-page transfer chunk (one page's worth of this buffer)."""
|
||
return (page_size // self.tokens_per_row) * self.row_bytes
|
||
|
||
|
||
class KVCache(abc.ABC):
|
||
layer_shard_enabled: bool = False
|
||
post_capture_active: bool = False
|
||
# Whether get_cpu_copy/load_cpu_copy carry the recurrent state. False when the
|
||
# state lives on the request pool instead, and the caller has to move it.
|
||
cpu_copy_carries_mamba: bool = False
|
||
|
||
@abc.abstractmethod
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
page_size: int,
|
||
dtype: torch.dtype,
|
||
layer_num: int,
|
||
device: str,
|
||
enable_memory_saver: bool,
|
||
start_layer: Optional[int] = None,
|
||
end_layer: Optional[int] = None,
|
||
allocation_label: Optional[str] = None,
|
||
):
|
||
self.size = size
|
||
self.page_size = page_size
|
||
# Row-blocks one page holds in this pool's kernel-facing id space; >1
|
||
# only for the unified pool's per-layer views, and then a write loc must
|
||
# have been translated into that space first.
|
||
self.kernel_page_blocks = 1
|
||
self.dtype = dtype
|
||
self.device = device
|
||
if dtype in (torch.float8_e5m2, torch.float8_e4m3fn, torch.float8_e4m3fnuz):
|
||
# NOTE: Store as torch.uint8 because Tensor.index_put is not implemented for torch.float8_e5m2
|
||
self.store_dtype = torch.uint8
|
||
else:
|
||
self.store_dtype = dtype
|
||
self.layer_num = layer_num
|
||
self.start_layer = start_layer or 0
|
||
self.end_layer = end_layer or layer_num - 1
|
||
self.allocation_label = allocation_label
|
||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||
enable=enable_memory_saver
|
||
)
|
||
self.mem_usage = 0
|
||
|
||
# used for chunked cpu-offloading
|
||
self.cpu_offloading_chunk_size = 8192
|
||
|
||
# default state for optional layer-wise transfer control
|
||
self.layer_transfer_counter = None
|
||
|
||
# for disagg with nvlink
|
||
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
|
||
maybe_init_custom_mem_pool(device=self.device)
|
||
)
|
||
|
||
def _finalize_allocation_log(self, num_tokens: int):
|
||
"""Common logging and mem_usage computation for KV cache allocation.
|
||
Supports both tuple (K, V) size returns and single KV size returns.
|
||
"""
|
||
cache_name = (
|
||
f"{self.allocation_label} KV Cache"
|
||
if self.allocation_label is not None
|
||
else "KV Cache"
|
||
)
|
||
kv_size_bytes = self.get_kv_size_bytes()
|
||
if isinstance(kv_size_bytes, tuple):
|
||
k_size, v_size = kv_size_bytes
|
||
k_size_GB = k_size / GB
|
||
v_size_GB = v_size / GB
|
||
logger.info(
|
||
f"{cache_name} {'VA upper bound' if self.post_capture_active else 'is allocated'}. dtype: {self.dtype}, "
|
||
f"#tokens: {num_tokens}, K size: {k_size_GB:.2f} GB, "
|
||
f"V size: {v_size_GB:.2f} GB"
|
||
)
|
||
self.mem_usage = k_size_GB + v_size_GB
|
||
else:
|
||
kv_size_GB = kv_size_bytes / GB
|
||
logger.info(
|
||
f"{cache_name} {'VA upper bound' if self.post_capture_active else 'is allocated'}. dtype: {self.dtype}, "
|
||
f"#tokens: {num_tokens}, KV size: {kv_size_GB:.2f} GB"
|
||
)
|
||
self.mem_usage = kv_size_GB
|
||
|
||
def get_kv_buffer_shape(self) -> Tuple[torch.Size, torch.Size]:
|
||
k_buffer, v_buffer = self.get_kv_buffer(self.start_layer)
|
||
return k_buffer.shape, v_buffer.shape
|
||
|
||
@abc.abstractmethod
|
||
def get_key_buffer(self, layer_id: int) -> torch.Tensor:
|
||
raise NotImplementedError()
|
||
|
||
@abc.abstractmethod
|
||
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
|
||
raise NotImplementedError()
|
||
|
||
@abc.abstractmethod
|
||
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
raise NotImplementedError()
|
||
|
||
@abc.abstractmethod
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
) -> None:
|
||
raise NotImplementedError()
|
||
|
||
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
|
||
self.layer_transfer_counter = layer_transfer_counter
|
||
|
||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||
raise NotImplementedError()
|
||
|
||
def load_cpu_copy(
|
||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=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
|
||
|
||
|
||
class MHATokenToKVPool(KVCache):
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
page_size: int,
|
||
dtype: torch.dtype,
|
||
head_num: int,
|
||
head_dim: int,
|
||
layer_num: int,
|
||
device: str,
|
||
enable_memory_saver: bool,
|
||
v_head_dim: Optional[int] = None,
|
||
swa_head_num: Optional[int] = None,
|
||
swa_head_dim: Optional[int] = None,
|
||
swa_v_head_dim: Optional[int] = None,
|
||
start_layer: Optional[int] = None,
|
||
end_layer: Optional[int] = None,
|
||
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,
|
||
allocation_label: Optional[str] = None,
|
||
):
|
||
self.k_buffer = None
|
||
self.v_buffer = None
|
||
if post_capture_active:
|
||
# Reserved upper bound only (unbacked VA): page-align UP so
|
||
# (size + page_size) % page_size == 0 holds for paged layouts.
|
||
size = (size + page_size - 1) // page_size * page_size
|
||
super().__init__(
|
||
size,
|
||
page_size,
|
||
dtype,
|
||
layer_num,
|
||
device,
|
||
enable_memory_saver,
|
||
start_layer,
|
||
end_layer,
|
||
allocation_label,
|
||
)
|
||
self.post_capture_active = post_capture_active
|
||
self._post_capture_owner = None
|
||
self.head_num = swa_head_num if swa_head_num is not None else head_num
|
||
self.head_dim = swa_head_dim if swa_head_dim is not None else head_dim
|
||
self.v_head_dim = (
|
||
swa_v_head_dim
|
||
if swa_v_head_dim is not None
|
||
else v_head_dim
|
||
if v_head_dim is not None
|
||
else head_dim
|
||
)
|
||
|
||
# Layout: NHD (default) | HND (SGLANG_USE_HND_KVCACHE) | vectorized_5d (ROCm AITER).
|
||
# HND folds (page, head) into one paged index for per-kv-head sparse page tables
|
||
# (paged backends like trtllm_mha consume directly). vectorized_5d SHUFFLE 5D:
|
||
# K: (num_blocks, H, D_k // X, page, X) V: (num_blocks, H, page // X, D_v, X),
|
||
# X = 16 / dtype_bytes — AITER-only (ignored elsewhere, no consumer kernel).
|
||
# HND and vectorized_5d are mutually exclusive; HND takes precedence.
|
||
self.use_hnd = envs.SGLANG_USE_HND_KVCACHE.get()
|
||
self.use_native_move_kv_cache = envs.SGLANG_NATIVE_MOVE_KV_CACHE.get()
|
||
if kv_cache_layout is not None:
|
||
# Explicit physical-layout selector wins over the platform default.
|
||
# This is a label only; layouts that change buffer identity (e.g. the
|
||
# page-granularity envelope) live in a dedicated pool subclass
|
||
# (PageMajorMHATokenToKVPool) rather than in branches here.
|
||
self.use_hnd = False
|
||
self.kv_cache_layout = kv_cache_layout
|
||
elif self.use_hnd:
|
||
total_slots = self.size + self.page_size
|
||
assert total_slots % self.page_size == 0, (
|
||
f"HND KV cache needs (size+page_size) divisible by page_size, got "
|
||
f"size={self.size}, page_size={self.page_size}"
|
||
)
|
||
self.num_pages = total_slots // self.page_size
|
||
self.kv_cache_layout = "hnd"
|
||
else:
|
||
self.kv_cache_layout = "nhd"
|
||
if _use_aiter:
|
||
layout = envs.SGLANG_AITER_KV_CACHE_LAYOUT.get().lower()
|
||
if layout not in ("nhd", "vectorized_5d"):
|
||
raise ValueError(
|
||
f"Unsupported SGLANG_AITER_KV_CACHE_LAYOUT={layout!r}; "
|
||
"expected 'nhd' or 'vectorized_5d'."
|
||
)
|
||
self.kv_cache_layout = layout
|
||
if layout == "vectorized_5d":
|
||
# X = 16 / storage itemsize: sized by the STORAGE dtype (not compute
|
||
# dtype) since it tiles the 16-byte on-pool vector.
|
||
self._kv_vector_x = 16 // self.store_dtype.itemsize
|
||
assert (self.size + self.page_size) % self.page_size == 0
|
||
assert self.page_size % self._kv_vector_x == 0, (
|
||
f"page_size={self.page_size} must be divisible by "
|
||
f"X={self._kv_vector_x} for vectorized_5d layout"
|
||
)
|
||
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)
|
||
|
||
_use_alt_stream = _is_cuda or current_platform.is_cuda_alike()
|
||
self.alt_stream = (
|
||
self.device_module.Stream()
|
||
if _use_alt_stream and enable_alt_stream
|
||
else None
|
||
)
|
||
|
||
if enable_kv_cache_copy and not self.use_hnd:
|
||
# The tiled byte copy assumes NHD slot-rows; HND uses a (page, off)
|
||
# gather in move_kv_cache instead, so skip the slot-row copy config.
|
||
self._init_kv_copy_and_warmup()
|
||
else:
|
||
self._kv_copy_config = None
|
||
|
||
self._finalize_allocation_log(size)
|
||
|
||
# for store_cache JIT kernel
|
||
self.row_dim = self.head_num * self.head_dim
|
||
self.v_row_dim = self.head_num * self.v_head_dim
|
||
|
||
def _init_kv_copy_and_warmup(self):
|
||
# Zero-layer pool (e.g. all-SWA model's full sub-pool) has no buffers.
|
||
if self.layer_num == 0:
|
||
self._kv_copy_config = None
|
||
return
|
||
|
||
# Heuristics for KV copy tiling
|
||
_KV_COPY_STRIDE_THRESHOLD_LARGE = 8192
|
||
_KV_COPY_STRIDE_THRESHOLD_MEDIUM = 4096
|
||
_KV_COPY_TILE_SIZE_LARGE = 512
|
||
_KV_COPY_TILE_SIZE_MEDIUM = 256
|
||
_KV_COPY_TILE_SIZE_SMALL = 128
|
||
_KV_COPY_NUM_WARPS_LARGE_TILE = 8
|
||
_KV_COPY_NUM_WARPS_SMALL_TILE = 4
|
||
|
||
stride_bytes = int(self.data_strides[0].item())
|
||
if stride_bytes >= _KV_COPY_STRIDE_THRESHOLD_LARGE:
|
||
bytes_per_tile = _KV_COPY_TILE_SIZE_LARGE
|
||
elif stride_bytes >= _KV_COPY_STRIDE_THRESHOLD_MEDIUM:
|
||
bytes_per_tile = _KV_COPY_TILE_SIZE_MEDIUM
|
||
else:
|
||
bytes_per_tile = _KV_COPY_TILE_SIZE_SMALL
|
||
|
||
# Calculate num_locs_upper to avoid large Triton specialization (e.g. 8192)
|
||
chunk_upper = 128 if bytes_per_tile >= _KV_COPY_TILE_SIZE_LARGE else 256
|
||
|
||
self._kv_copy_config = {
|
||
"bytes_per_tile": bytes_per_tile,
|
||
"byte_tiles": (stride_bytes + bytes_per_tile - 1) // bytes_per_tile,
|
||
"num_warps": (
|
||
_KV_COPY_NUM_WARPS_SMALL_TILE
|
||
if bytes_per_tile <= _KV_COPY_TILE_SIZE_MEDIUM
|
||
else _KV_COPY_NUM_WARPS_LARGE_TILE
|
||
),
|
||
"num_locs_upper": chunk_upper,
|
||
}
|
||
|
||
dummy_loc = torch.zeros(chunk_upper, dtype=torch.int64, device=self.device)
|
||
copy_all_layer_kv_cache_func(
|
||
self.data_ptrs,
|
||
self.data_strides,
|
||
dummy_loc,
|
||
dummy_loc,
|
||
1,
|
||
chunk_upper,
|
||
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.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.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,
|
||
device=self.device,
|
||
)
|
||
self.v_data_ptrs = torch.tensor(
|
||
[x.data_ptr() for x in self.v_buffer],
|
||
dtype=torch.uint64,
|
||
device=self.device,
|
||
)
|
||
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 slot_move_pointer_buffers
|
||
],
|
||
device=self.device,
|
||
)
|
||
|
||
def _kv_buffer_shapes(self):
|
||
"""(k_shape, v_shape)"""
|
||
if self.use_hnd:
|
||
return (
|
||
(self.num_pages, self.head_num, self.page_size, self.head_dim),
|
||
(self.num_pages, self.head_num, self.page_size, self.v_head_dim),
|
||
)
|
||
rows = self.size + self.page_size
|
||
return (
|
||
(rows, self.head_num, self.head_dim),
|
||
(rows, self.head_num, self.v_head_dim),
|
||
)
|
||
|
||
def _create_buffers_normal(self):
|
||
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()
|
||
):
|
||
# The padded page (slot 0's page) absorbs dummy padded-token writes.
|
||
if self.kv_cache_layout == "vectorized_5d":
|
||
total_slots = self.size + self.page_size
|
||
num_blocks = total_slots // self.page_size
|
||
x = self._kv_vector_x
|
||
# K: (num_blocks, H, D_k // X, page, X)
|
||
self.k_buffer = [
|
||
torch.zeros(
|
||
(
|
||
num_blocks,
|
||
self.head_num,
|
||
self.head_dim // x,
|
||
self.page_size,
|
||
x,
|
||
),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
# V: (num_blocks, H, page // X, D_v, X)
|
||
self.v_buffer = [
|
||
torch.zeros(
|
||
(
|
||
num_blocks,
|
||
self.head_num,
|
||
self.page_size // x,
|
||
self.v_head_dim,
|
||
x,
|
||
),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
else:
|
||
k_shape, v_shape = self._kv_buffer_shapes()
|
||
self.k_buffer = [
|
||
torch.zeros(k_shape, dtype=self.store_dtype, device=self.device)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
self.v_buffer = [
|
||
torch.zeros(v_shape, dtype=self.store_dtype, device=self.device)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
# -- post-capture VA backing (opt-in; overridable per layout) --------------
|
||
|
||
def _build_kv_buffer_descs(self):
|
||
"""Per-buffer layout descriptors, k0..k(L-1) then v0..v(L-1). Drives both the
|
||
CUDA-VMM post-capture backing and PD-transfer registration
|
||
(get_contiguous_buf_infos). Override per layout."""
|
||
itemsize = self.store_dtype.itemsize
|
||
# Derive from the real buffers when they exist (covers arbitrary layouts,
|
||
# e.g. vectorized_5d); fall back to _kv_buffer_shapes for the pre-allocation
|
||
# post-capture call, which only runs for NHD/HND.
|
||
if self.k_buffer and self.v_buffer:
|
||
k_shape = tuple(self.k_buffer[0].shape)
|
||
v_shape = tuple(self.v_buffer[0].shape)
|
||
else:
|
||
k_shape, v_shape = self._kv_buffer_shapes()
|
||
# A row is a whole page when the leading dim is pages (hnd, vectorized_5d),
|
||
# a single token slot for the plain NHD [slots, ...] layout.
|
||
num_slots = self.size + self.page_size
|
||
tokens_per_row = (
|
||
self.page_size if k_shape[0] * self.page_size == num_slots else 1
|
||
)
|
||
descs = []
|
||
for prefix, shape in (("k", k_shape), ("v", v_shape)):
|
||
row_bytes = int(np.prod(shape[1:])) * itemsize
|
||
for layer in range(self.layer_num):
|
||
descs.append(
|
||
KvBufferDesc(
|
||
f"{prefix}{layer}",
|
||
shape,
|
||
row_bytes=row_bytes,
|
||
tokens_per_row=tokens_per_row,
|
||
)
|
||
)
|
||
return descs
|
||
|
||
def _assign_post_capture_tensors(self, tensors):
|
||
"""Map owner tensors (in ``_build_kv_buffer_descs`` order) to k/v_buffer."""
|
||
self.k_buffer = tensors[: self.layer_num]
|
||
self.v_buffer = tensors[self.layer_num :]
|
||
|
||
def _alloc_post_capture_buffers(self):
|
||
dev = torch.device(self.device)
|
||
device_id = dev.index if dev.index is not None else torch.cuda.current_device()
|
||
self._post_capture_owner = KvVmmBufferOwner(
|
||
device=self.device,
|
||
device_id=device_id,
|
||
store_dtype=self.store_dtype,
|
||
page_size=self.page_size,
|
||
reserved_num_tokens=self.size,
|
||
buffer_descs=self._build_kv_buffer_descs(),
|
||
)
|
||
self._assign_post_capture_tensors(self._post_capture_owner.tensors)
|
||
|
||
def finalize_backing(self, config) -> None:
|
||
"""After capture+sizing: back the final span and set serving capacity.
|
||
``config`` is a MemoryPoolConfig (duck-typed); each pool family reads the
|
||
fields it needs, so the finalizer stays pool-agnostic."""
|
||
self._finalize_backing_tokens(config.max_total_num_tokens)
|
||
|
||
def _finalize_backing_tokens(self, final_num_tokens: int) -> None:
|
||
"""Token-count primitive shared by composite pools (e.g. SWA sub-pools)."""
|
||
self._post_capture_owner.finalize(final_num_tokens)
|
||
self.size = int(final_num_tokens)
|
||
|
||
@property
|
||
def post_capture_backed_bytes(self) -> int:
|
||
return self._post_capture_owner.backed_bytes if self._post_capture_owner else 0
|
||
|
||
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
|
||
|
||
def get_kv_size_bytes(self):
|
||
assert hasattr(self, "k_buffer")
|
||
assert hasattr(self, "v_buffer")
|
||
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
|
||
def _pd_registerable_tensors(self):
|
||
"""Buffers to register for PD KV transfer, in ``_kv_buffer_descs`` order.
|
||
Override when the registerable storage differs from k/v_buffer."""
|
||
return self.k_buffer + self.v_buffer
|
||
|
||
def get_contiguous_buf_infos(self):
|
||
"""(ptrs, lens, item_lens) for PD KV transfer, derived from the descriptors.
|
||
``lens`` is the final span at the CURRENT serving size -- for a post-capture
|
||
pool that is the physically-backed span, not the reserved VA upper bound."""
|
||
assert not self.use_hnd, (
|
||
"PD-disaggregation KV transfer assumes NHD slot-row layout; "
|
||
"HND KV cache (SGLANG_USE_HND_KVCACHE) is not supported with disagg yet."
|
||
)
|
||
tensors = self._pd_registerable_tensors()
|
||
ptrs = [t.data_ptr() for t in tensors]
|
||
lens = [
|
||
d.final_span_bytes(self.size, self.page_size) for d in self._kv_buffer_descs
|
||
]
|
||
item_lens = [d.item_len_bytes(self.page_size) for d in self._kv_buffer_descs]
|
||
return ptrs, lens, item_lens
|
||
|
||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||
assert not self.use_hnd, (
|
||
"CPU KV offload indexes by slot (NHD); HND KV cache "
|
||
"(SGLANG_USE_HND_KVCACHE) is not supported with CPU offload yet."
|
||
)
|
||
current_platform.synchronize()
|
||
kv_cache_cpu = []
|
||
chunk_size = self.cpu_offloading_chunk_size
|
||
for layer_id in range(self.layer_num):
|
||
kv_cache_cpu.append([])
|
||
for i in range(0, len(indices), chunk_size):
|
||
chunk_indices = indices[i : i + chunk_size]
|
||
k_cpu = self.k_buffer[layer_id][chunk_indices].to(
|
||
"cpu", non_blocking=True
|
||
)
|
||
v_cpu = self.v_buffer[layer_id][chunk_indices].to(
|
||
"cpu", non_blocking=True
|
||
)
|
||
kv_cache_cpu[-1].append([k_cpu, v_cpu])
|
||
current_platform.synchronize()
|
||
return kv_cache_cpu
|
||
|
||
def load_cpu_copy(
|
||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||
):
|
||
assert not self.use_hnd, (
|
||
"CPU KV offload indexes by slot (NHD); HND KV cache "
|
||
"(SGLANG_USE_HND_KVCACHE) is not supported with CPU offload yet."
|
||
)
|
||
current_platform.synchronize()
|
||
chunk_size = self.cpu_offloading_chunk_size
|
||
for layer_id in range(self.layer_num):
|
||
for i in range(0, len(indices), chunk_size):
|
||
chunk_indices = indices[i : i + chunk_size]
|
||
k_cpu, v_cpu = (
|
||
kv_cache_cpu[layer_id][i // chunk_size][0],
|
||
kv_cache_cpu[layer_id][i // chunk_size][1],
|
||
)
|
||
assert k_cpu.shape[0] == v_cpu.shape[0] == len(chunk_indices)
|
||
k_chunk = k_cpu.to(self.k_buffer[0].device, non_blocking=True)
|
||
v_chunk = v_cpu.to(self.v_buffer[0].device, non_blocking=True)
|
||
self.k_buffer[layer_id][chunk_indices] = k_chunk
|
||
self.v_buffer[layer_id][chunk_indices] = v_chunk
|
||
current_platform.synchronize()
|
||
|
||
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[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
|
||
# it is supposed to be used only by attention backend not for information purpose
|
||
# same applies to get_value_buffer and get_kv_buffer
|
||
if self.layer_transfer_counter is not None:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
return self._get_key_buffer(layer_id)
|
||
|
||
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[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:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
return self._get_value_buffer(layer_id)
|
||
|
||
def get_v_head_dim(self):
|
||
# Every layer in this pool is full-attention, so the value head dim is
|
||
# uniform and known at construction. Mirrors HybridLinearKVPool's
|
||
# get_v_head_dim() so the TritonAttnBackend mambaish branch works when a
|
||
# mamba2 config is served by a plain MHA pool (no per-linear-layer split).
|
||
return self.v_head_dim
|
||
|
||
def get_kv_buffer(self, layer_id: int):
|
||
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
|
||
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc_info,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
k_scale: Optional[float] = None,
|
||
v_scale: Optional[float] = None,
|
||
layer_id_override: Optional[int] = None,
|
||
dcp_kv_mask: Optional[torch.Tensor] = None,
|
||
):
|
||
loc, _, _ = unwrap_write_loc(loc_info)
|
||
# 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)")
|
||
maybe_detect_kernel_facing_loc(
|
||
loc, self.page_size, self.kernel_page_blocks, "set_kv_buffer (MHA)"
|
||
)
|
||
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)
|
||
if v_scale is not None:
|
||
cache_v.div_(v_scale)
|
||
cache_k = cache_k.to(self.dtype)
|
||
cache_v = cache_v.to(self.dtype)
|
||
|
||
if self.store_dtype != self.dtype:
|
||
cache_k = cache_k.view(self.store_dtype)
|
||
cache_v = cache_v.view(self.store_dtype)
|
||
|
||
if dcp_kv_mask is not None:
|
||
N, H, D = cache_k.shape
|
||
masked_set_kv_buffer_kernel[(N,)](
|
||
cache_k,
|
||
cache_v,
|
||
self.k_buffer[layer_id - self.start_layer],
|
||
self.v_buffer[layer_id - self.start_layer],
|
||
loc,
|
||
dcp_kv_mask,
|
||
N,
|
||
H,
|
||
D,
|
||
128,
|
||
cache_k.stride(0),
|
||
cache_k.stride(1),
|
||
cache_v.stride(0),
|
||
cache_v.stride(1),
|
||
)
|
||
return
|
||
|
||
if self.use_hnd:
|
||
# A slot is [page, :, off, :] (not a contiguous row), so scatter by (page, off).
|
||
k_buf = self.k_buffer[layer_id - self.start_layer]
|
||
v_buf = self.v_buffer[layer_id - self.start_layer]
|
||
pages = loc // self.page_size
|
||
offs = loc % self.page_size
|
||
k_buf[pages, :, offs, :] = cache_k
|
||
v_buf[pages, :, offs, :] = cache_v
|
||
return
|
||
|
||
self._store_kv_layer(layer_id - self.start_layer, loc, cache_k, cache_v)
|
||
|
||
def _store_kv_layer(
|
||
self,
|
||
layer_idx: int,
|
||
loc: torch.Tensor,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
):
|
||
# Per-layer physical write into K/V buffer ``layer_idx``. Override for
|
||
# layouts that change buffer identity (e.g. PageMajorMHATokenToKVPool's
|
||
# 4-D strided views). ``loc`` and the cache tensors are already dtype-cast
|
||
# and viewed as ``store_dtype`` by ``set_kv_buffer``.
|
||
if self.kv_cache_layout == "vectorized_5d":
|
||
# Late-import to keep the NHD path import-clean.
|
||
from sglang.kernels.ops.attention.utils import (
|
||
launch_reshape_and_cache_shuffle_5d,
|
||
)
|
||
|
||
# The writer kernel uses key.stride(0) directly as the source
|
||
# token stride; head/dim are assumed contiguous within each
|
||
# token (stride(1)=head_size, stride(2)=1). Both hold for K/V
|
||
# produced by QKV split + RoPE in upstream attention even when
|
||
# the outer per-token stride is non-canonical, so we skip the
|
||
# protective .contiguous() copies that would otherwise fire
|
||
# large per-layer elementwise kernels.
|
||
launch_reshape_and_cache_shuffle_5d(
|
||
cache_k,
|
||
cache_v,
|
||
self.k_buffer[layer_idx],
|
||
self.v_buffer[layer_idx],
|
||
loc,
|
||
)
|
||
return
|
||
|
||
_set_kv_buffer_impl(
|
||
cache_k,
|
||
cache_v,
|
||
self.k_buffer[layer_idx],
|
||
self.v_buffer[layer_idx],
|
||
loc,
|
||
row_dim=self.row_dim,
|
||
store_dtype=self.store_dtype,
|
||
device_module=self.device_module,
|
||
# size + page_size = real slots + the reserved padding slot (padded /
|
||
# dummy tokens write there); valid index range is [0, size + page_size).
|
||
size_limit=self.size + self.page_size,
|
||
alt_stream=self.alt_stream,
|
||
v_row_dim=self.v_row_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,
|
||
loc_2d: torch.Tensor,
|
||
commit_lens: torch.Tensor,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
k_scale: Optional[float] = None,
|
||
v_scale: Optional[float] = None,
|
||
layer_id_override: Optional[int] = None,
|
||
):
|
||
if layer_id_override is not None:
|
||
layer_id = layer_id_override
|
||
else:
|
||
layer_id = layer.layer_id
|
||
|
||
if loc_2d.ndim != 2:
|
||
raise ValueError(f"loc_2d must be rank-2, got shape={tuple(loc_2d.shape)}.")
|
||
if commit_lens.ndim != 1 or commit_lens.shape[0] != loc_2d.shape[0]:
|
||
raise ValueError(
|
||
"commit_lens must match loc_2d batch size: "
|
||
f"{tuple(commit_lens.shape)=} {tuple(loc_2d.shape)=}."
|
||
)
|
||
|
||
num_rows = int(loc_2d.numel())
|
||
if cache_k.shape[0] != num_rows or cache_v.shape[0] != num_rows:
|
||
raise ValueError(
|
||
"KV rows must match loc_2d size: "
|
||
f"{tuple(cache_k.shape)=} {tuple(cache_v.shape)=} {tuple(loc_2d.shape)=}."
|
||
)
|
||
|
||
if cache_k.dtype != self.dtype:
|
||
if k_scale is not None:
|
||
cache_k.div_(k_scale)
|
||
if v_scale is not None:
|
||
cache_v.div_(v_scale)
|
||
cache_k = cache_k.to(self.dtype)
|
||
cache_v = cache_v.to(self.dtype)
|
||
|
||
if self.store_dtype != self.dtype:
|
||
cache_k = cache_k.contiguous().view(self.store_dtype)
|
||
cache_v = cache_v.contiguous().view(self.store_dtype)
|
||
else:
|
||
cache_k = cache_k.contiguous()
|
||
cache_v = cache_v.contiguous()
|
||
|
||
if loc_2d.device != self.k_buffer[0].device:
|
||
loc_2d = loc_2d.to(device=self.k_buffer[0].device, non_blocking=True)
|
||
if commit_lens.device != self.k_buffer[0].device:
|
||
commit_lens = commit_lens.to(
|
||
device=self.k_buffer[0].device, non_blocking=True
|
||
)
|
||
if loc_2d.dtype != torch.int64:
|
||
loc_2d = loc_2d.to(torch.int64)
|
||
if commit_lens.dtype != torch.int32:
|
||
commit_lens = commit_lens.to(torch.int32)
|
||
|
||
if not (_is_cuda or _is_hip):
|
||
row_offsets = torch.arange(loc_2d.shape[1], device=loc_2d.device)
|
||
valid_mask = row_offsets[None, :] < commit_lens.to(torch.int64)[:, None]
|
||
valid_idx = torch.nonzero(valid_mask.reshape(-1), as_tuple=False).flatten()
|
||
if valid_idx.numel() == 0:
|
||
return
|
||
self.set_kv_buffer(
|
||
layer,
|
||
loc_2d.reshape(-1).index_select(0, valid_idx),
|
||
cache_k.index_select(0, valid_idx),
|
||
cache_v.index_select(0, valid_idx),
|
||
k_scale,
|
||
v_scale,
|
||
layer_id_override=layer_id,
|
||
)
|
||
return
|
||
|
||
# The tiled kernel takes one ROW_BYTES for both tensors, so an asymmetric V
|
||
# row would be written at K's width and bleed into the next slot. Only this
|
||
# path needs the gate; the non-CUDA branch above handles both widths.
|
||
if self.v_row_dim != self.row_dim:
|
||
raise NotImplementedError(
|
||
"prefix-valid commit requires equal-width K/V rows, got "
|
||
f"head_dim={self.head_dim} v_head_dim={self.v_head_dim}."
|
||
)
|
||
|
||
_set_kv_buffer_prefix_valid_impl(
|
||
cache_k,
|
||
cache_v,
|
||
self.k_buffer[layer_id - self.start_layer],
|
||
self.v_buffer[layer_id - self.start_layer],
|
||
loc_2d,
|
||
commit_lens,
|
||
row_dim=self.row_dim,
|
||
store_dtype=self.store_dtype,
|
||
)
|
||
|
||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
# Zero-layer pool (e.g. all-SWA model's full sub-pool) has no buffers.
|
||
if self.layer_num == 0:
|
||
return
|
||
|
||
# Catch stale indices here instead of as illegal-addr or silent KV corruption.
|
||
size_limit = self.size + self.page_size
|
||
maybe_detect_oob(tgt_loc, 0, size_limit, "move_kv_cache tgt_loc")
|
||
maybe_detect_oob(src_loc, 0, size_limit, "move_kv_cache src_loc")
|
||
|
||
if self.use_hnd:
|
||
pages_t, offs_t = tgt_loc // self.page_size, tgt_loc % self.page_size
|
||
pages_s, offs_s = src_loc // self.page_size, src_loc % self.page_size
|
||
for kb, vb in zip(self.k_buffer, self.v_buffer):
|
||
kb[pages_t, :, offs_t, :] = kb[pages_s, :, offs_s, :]
|
||
vb[pages_t, :, offs_t, :] = vb[pages_s, :, offs_s, :]
|
||
return
|
||
|
||
self._move_kv_cache_impl(tgt_loc, src_loc)
|
||
|
||
def _move_kv_cache_impl(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
# Physical move strategy. Override for layouts that change buffer
|
||
# identity (e.g. PageMajorMHATokenToKVPool always uses the native move).
|
||
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()
|
||
if N == 0:
|
||
return
|
||
|
||
assert self._kv_copy_config is not None, (
|
||
"KV copy not initialized. Set enable_kv_cache_copy=True in __init__"
|
||
)
|
||
|
||
cfg = self._kv_copy_config
|
||
cap = int(cfg.get("num_locs_upper", 256))
|
||
|
||
if N <= cap:
|
||
copy_all_layer_kv_cache_func(
|
||
self.data_ptrs,
|
||
self.data_strides,
|
||
tgt_loc,
|
||
src_loc,
|
||
N,
|
||
next_power_of_2(N),
|
||
cfg,
|
||
)
|
||
return
|
||
|
||
# Huge N: chunk, but each chunk's upper is still pow2(<= cap)
|
||
for start in range(0, N, cap):
|
||
end = min(start + cap, N)
|
||
chunk_len = end - start
|
||
copy_all_layer_kv_cache_func(
|
||
self.data_ptrs,
|
||
self.data_strides,
|
||
tgt_loc[start:end],
|
||
src_loc[start:end],
|
||
chunk_len,
|
||
next_power_of_2(chunk_len),
|
||
cfg,
|
||
)
|
||
|
||
|
||
class NoOpMHATokenToKVPool(MHATokenToKVPool):
|
||
"""KV cache pool that skips physical K/V buffer allocation.
|
||
|
||
Used in embedding-mode prefill-only workloads with the FA
|
||
fa_skip_kv_cache path, where no layer reads or writes KV cache because
|
||
attention uses raw K/V via flash_attn_varlen_func. Other prefill-only paths
|
||
such as scoring/MIS may benefit from the same idea later, but some still
|
||
stage K/V through paged cache today.
|
||
|
||
This class keeps the scheduler's view of pool capacity (self.size is
|
||
honored for admission) but allocates only (page_size, head_num, head_dim)
|
||
placeholder tensors per layer to satisfy any code paths that dereference
|
||
the buffers.
|
||
|
||
Callers MUST ensure no real set_kv_buffer/get_*_buffer calls happen against
|
||
this pool; those paths raise loudly so misuse is visible.
|
||
"""
|
||
|
||
def _create_buffers(self):
|
||
# No-op pool keeps tiny NHD placeholders regardless of SGLANG_USE_HND_KVCACHE
|
||
# (no real KV is stored), so force NHD here to keep the store/move fast paths.
|
||
self.use_hnd = False
|
||
self.kv_cache_layout = "nhd"
|
||
# Allocate minimal placeholder buffers. They exist purely so that code
|
||
# paths holding `k_buffer` / `v_buffer` references (pointer tables,
|
||
# layer-transfer counters, stride arithmetic) keep working without
|
||
# None-guards scattered across the codebase. Shape is
|
||
# [page_size, head_num, head_dim] per layer so that the unconditional
|
||
# `key_cache.view(-1, page_size, head_num, head_dim)` in the FA backend
|
||
# at the top of forward_extend succeeds regardless of --page-size.
|
||
# Total footprint is still on the order of KB vs GBs for a real pool.
|
||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||
self.k_buffer = [
|
||
torch.zeros(
|
||
(self.page_size, self.head_num, self.head_dim),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
self.v_buffer = [
|
||
torch.zeros(
|
||
(self.page_size, self.head_num, self.v_head_dim),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
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
|
||
for x in self.k_buffer + self.v_buffer
|
||
],
|
||
device=self.device,
|
||
)
|
||
|
||
def _finalize_allocation_log(self, num_tokens: int):
|
||
self.mem_usage = 0.0
|
||
placeholder_bytes = (
|
||
2
|
||
* self.layer_num
|
||
* self.page_size
|
||
* self.head_num
|
||
* max(self.head_dim, self.v_head_dim)
|
||
* self.store_dtype.itemsize
|
||
)
|
||
logger.info(
|
||
f"KV Cache skipped (no-op pool). Logical #tokens: {num_tokens}, "
|
||
f"physical K/V size: ~{placeholder_bytes / 1024:.1f} KB placeholder"
|
||
)
|
||
|
||
def get_kv_size_bytes(self):
|
||
# Report zero so downstream memory accounting matches reality.
|
||
return (0, 0)
|
||
|
||
def set_kv_buffer(self, *args, **kwargs):
|
||
raise RuntimeError(
|
||
"NoOpMHATokenToKVPool.set_kv_buffer was called. This pool is only "
|
||
"valid in prefill-only modes (e.g. --is-embedding, scoring) with "
|
||
"the FA backend's fa_skip_kv_cache path active; the attention "
|
||
"backend must never write to it. Check that the workload truly "
|
||
"performs no decode and that the FA backend's fa_skip_kv_cache "
|
||
"preconditions are met."
|
||
)
|
||
|
||
def get_key_buffer(self, layer_id: int):
|
||
# Return the placeholder. The FA backend reads this before taking the
|
||
# fa_skip_kv_cache branch (which does not use it); the placeholder shape
|
||
# is (page_size, head_num, head_dim) so downstream .view() calls succeed.
|
||
return self.k_buffer[layer_id - self.start_layer]
|
||
|
||
def get_value_buffer(self, layer_id: int):
|
||
return self.v_buffer[layer_id - self.start_layer]
|
||
|
||
def get_kv_buffer(self, layer_id: int):
|
||
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
|
||
|
||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
# no-op; embedding mode has no KV cache to move
|
||
return
|
||
|
||
|
||
class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
||
def _create_buffers(self):
|
||
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()
|
||
):
|
||
# [size, head_num, head_dim] for each layer
|
||
# The padded slot 0 is used for writing dummy outputs from padded tokens.
|
||
m = self.size + self.page_size
|
||
n = self.head_num
|
||
k = self.head_dim
|
||
|
||
scale_block_size = 16
|
||
self.store_dtype = torch.uint8
|
||
self.k_buffer = [
|
||
torch.zeros(
|
||
(m, n, k // 2),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
self.v_buffer = [
|
||
torch.zeros(
|
||
(m, n, k // 2),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
self.k_scale_buffer = [
|
||
torch.zeros(
|
||
(m, (n * k) // scale_block_size),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
self.v_scale_buffer = [
|
||
torch.zeros(
|
||
(m, (n * k) // scale_block_size),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
def _clear_buffers(self):
|
||
del self.k_buffer
|
||
del self.v_buffer
|
||
del self.k_scale_buffer
|
||
del self.v_scale_buffer
|
||
|
||
def _get_key_buffer(self, layer_id: int):
|
||
# for internal use of referencing
|
||
if self.store_dtype != self.dtype:
|
||
cache_k_nope_fp4 = self.k_buffer[layer_id - self.start_layer].view(
|
||
torch.uint8
|
||
)
|
||
cache_k_nope_fp4_sf = self.k_scale_buffer[layer_id - self.start_layer]
|
||
|
||
from sglang.srt.layers.quantization.kvfp4_tensor import (
|
||
FP4MXBlock16KVQuantizeUtil,
|
||
)
|
||
|
||
cache_k_nope_fp4_dequant = FP4MXBlock16KVQuantizeUtil.batched_dequantize(
|
||
cache_k_nope_fp4, cache_k_nope_fp4_sf
|
||
)
|
||
return cache_k_nope_fp4_dequant
|
||
return self.k_buffer[layer_id - self.start_layer]
|
||
|
||
def _get_value_buffer(self, layer_id: int):
|
||
# for internal use of referencing
|
||
if self.store_dtype != self.dtype:
|
||
cache_v_nope_fp4 = self.v_buffer[layer_id - self.start_layer].view(
|
||
torch.uint8
|
||
)
|
||
cache_v_nope_fp4_sf = self.v_scale_buffer[layer_id - self.start_layer]
|
||
|
||
from sglang.srt.layers.quantization.kvfp4_tensor import (
|
||
FP4MXBlock16KVQuantizeUtil,
|
||
)
|
||
|
||
cache_v_nope_fp4_dequant = FP4MXBlock16KVQuantizeUtil.batched_dequantize(
|
||
cache_v_nope_fp4, cache_v_nope_fp4_sf
|
||
)
|
||
return cache_v_nope_fp4_dequant
|
||
return self.v_buffer[layer_id - self.start_layer]
|
||
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc_info,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
k_scale: Optional[float] = None,
|
||
v_scale: Optional[float] = None,
|
||
layer_id_override: Optional[int] = None,
|
||
):
|
||
loc, _, _ = unwrap_write_loc(loc_info)
|
||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA-FP4)")
|
||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||
|
||
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)
|
||
if v_scale is not None:
|
||
cache_v.div_(v_scale)
|
||
|
||
from sglang.srt.layers.quantization.kvfp4_tensor import (
|
||
FP4MXBlock16KVQuantizeUtil,
|
||
)
|
||
|
||
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)
|
||
cache_v = cache_v.view(self.store_dtype)
|
||
|
||
cache_k_fp4_sf = cache_k_fp4_sf.view(self.store_dtype)
|
||
cache_v_fp4_sf = cache_v_fp4_sf.view(self.store_dtype)
|
||
|
||
if get_is_capture_mode() and self.alt_stream is not None:
|
||
# Overlap the copy of K and V cache for small batch size
|
||
current_stream = self.device_module.current_stream()
|
||
self.alt_stream.wait_stream(current_stream)
|
||
self.k_buffer[layer_id - self.start_layer][loc] = cache_k
|
||
|
||
self.k_scale_buffer[layer_id - self.start_layer][loc] = cache_k_fp4_sf
|
||
with self.device_module.stream(self.alt_stream):
|
||
self.v_buffer[layer_id - self.start_layer][loc] = cache_v
|
||
|
||
self.v_scale_buffer[layer_id - self.start_layer][loc] = cache_v_fp4_sf
|
||
current_stream.wait_stream(self.alt_stream)
|
||
else:
|
||
self.k_buffer[layer_id - self.start_layer][loc] = cache_k
|
||
self.v_buffer[layer_id - self.start_layer][loc] = cache_v
|
||
|
||
self.k_scale_buffer[layer_id - self.start_layer][loc] = cache_k_fp4_sf
|
||
self.v_scale_buffer[layer_id - self.start_layer][loc] = cache_v_fp4_sf
|
||
|
||
|
||
class PageMajorMHATokenToKVPool(MHATokenToKVPool):
|
||
"""MHA pool with the page-major page-granularity envelope layout.
|
||
|
||
NON-CONSTRUCTIBLE: the strided 4-D view builder and its write kernel are
|
||
gone, and ServerArgs rejects the static page-major arm at boot. The class
|
||
stays as the seat for the per-layer-view reimplementation.
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
*args,
|
||
kv_cache_layout: Optional[str] = None,
|
||
enable_kv_cache_copy: bool = False,
|
||
**kwargs,
|
||
):
|
||
assert kv_cache_layout in (
|
||
None,
|
||
"page_major_layer_major",
|
||
), f"PageMajorMHATokenToKVPool fixes its layout; got {kv_cache_layout!r}"
|
||
# The tiled copy kernel assumes stride == row bytes, which the strided 4-D
|
||
# views violate, so the copy path is never available here regardless of
|
||
# what the caller requested (the spec-decode call sites pass
|
||
# enable_kv_cache_copy=True). Always fall back to the native move.
|
||
super().__init__(
|
||
*args,
|
||
kv_cache_layout="page_major_layer_major",
|
||
enable_kv_cache_copy=False,
|
||
**kwargs,
|
||
)
|
||
|
||
def _create_buffers(self):
|
||
raise NotImplementedError(
|
||
"PageMajorMHATokenToKVPool: the strided 4-D envelope views were "
|
||
"removed; the static-pool page-major layout is temporarily "
|
||
"unsupported (ServerArgs rejects it at startup). "
|
||
"--enable-unified-memory provides the page-major layout with "
|
||
"per-layer views."
|
||
)
|
||
|
||
# The methods below assume the per-layer contiguous 3-D layout. The 4-D
|
||
# strided envelope views have no per-layer contiguous region (their bytes are
|
||
# interleaved layer-major within each page) and index page-major, not
|
||
# token-major. Inheriting them would silently mis-index; fail loudly instead.
|
||
|
||
def get_contiguous_buf_infos(self):
|
||
raise NotImplementedError(
|
||
"page-major layout has no per-layer contiguous regions; KV transfer / "
|
||
"disaggregation is unsupported (TODO: expose the single _raw buffer "
|
||
"with a page-aware transfer scheme)."
|
||
)
|
||
|
||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||
raise NotImplementedError(
|
||
"CPU offloading is unsupported under the page-major layout "
|
||
"(TODO: split token ids into page/slot for the 4-D index)."
|
||
)
|
||
|
||
def load_cpu_copy(
|
||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||
):
|
||
raise NotImplementedError(
|
||
"CPU offloading is unsupported under the page-major layout "
|
||
"(TODO: split token ids into page/slot for the 4-D index)."
|
||
)
|
||
|
||
def set_kv_buffer_prefix_valid(self, *args, **kwargs):
|
||
raise NotImplementedError(
|
||
"prefix-valid commit is unsupported under the page-major layout "
|
||
"(_set_kv_buffer_prefix_valid_impl assumes 3-D contiguous + row_dim)."
|
||
)
|
||
|
||
|
||
class MHATokenToKVPoolMXFP8(MHATokenToKVPool):
|
||
"""MHA KV cache pool for MXFP8 block-scaled FP8.
|
||
|
||
K/V data is stored as FP8 E4M3. Per-32-element UE8M0 scale factors are
|
||
stored beside it and passed to the FA4 MXFP8 kernel.
|
||
"""
|
||
|
||
MXFP8_SCALE_BLOCK_SIZE = 32
|
||
|
||
def _create_buffers(self):
|
||
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()
|
||
):
|
||
m = self.size + self.page_size
|
||
n = self.head_num
|
||
k = self.head_dim
|
||
v = self.v_head_dim
|
||
|
||
if k % self.MXFP8_SCALE_BLOCK_SIZE != 0:
|
||
raise ValueError(
|
||
f"MXFP8 KV cache requires head_dim divisible by "
|
||
f"{self.MXFP8_SCALE_BLOCK_SIZE}, got {k}."
|
||
)
|
||
if v % self.MXFP8_SCALE_BLOCK_SIZE != 0:
|
||
raise ValueError(
|
||
f"MXFP8 KV cache requires v_head_dim divisible by "
|
||
f"{self.MXFP8_SCALE_BLOCK_SIZE}, got {v}."
|
||
)
|
||
if not hasattr(torch, "float8_e8m0fnu"):
|
||
raise RuntimeError(
|
||
"MXFP8 KV cache requires torch.float8_e8m0fnu support."
|
||
)
|
||
if self.use_hnd:
|
||
# Buffers are NHD; the inherited HND move_kv_cache branch
|
||
# would silently relocate wrong bytes.
|
||
raise ValueError(
|
||
"MXFP8 KV cache does not support SGLANG_USE_HND_KVCACHE."
|
||
)
|
||
|
||
self.store_dtype = torch.float8_e4m3fn
|
||
self.k_buffer = [
|
||
torch.zeros((m, n, k), dtype=self.store_dtype, device=self.device)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
self.v_buffer = [
|
||
torch.zeros((m, n, v), dtype=self.store_dtype, device=self.device)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
# UE8M0 scales, one per 32-element block. For the production
|
||
# page_size==128 path they are stored interleaved in the FA4
|
||
# BlockScaledBasicChunk atom layout
|
||
# (num_pages, head, 32, page_size//32, sf_dim) and written by
|
||
# the store_sf_interleaved kernel; otherwise flat per slot. Must
|
||
# be zero-initialized (garbage 0xFF is e8m0 NaN).
|
||
k_sf_dim = k // self.MXFP8_SCALE_BLOCK_SIZE
|
||
v_sf_dim = v // self.MXFP8_SCALE_BLOCK_SIZE
|
||
self.mxfp8_sf_interleaved = self.page_size == 128
|
||
if self.mxfp8_sf_interleaved:
|
||
assert m % self.page_size == 0
|
||
num_pages = m // self.page_size
|
||
chunk = self.page_size // self.MXFP8_SCALE_BLOCK_SIZE
|
||
k_sf_shape = (
|
||
num_pages,
|
||
n,
|
||
self.MXFP8_SCALE_BLOCK_SIZE,
|
||
chunk,
|
||
k_sf_dim,
|
||
)
|
||
v_sf_shape = (
|
||
num_pages,
|
||
n,
|
||
self.MXFP8_SCALE_BLOCK_SIZE,
|
||
chunk,
|
||
v_sf_dim,
|
||
)
|
||
else:
|
||
k_sf_shape = (m, n, k_sf_dim)
|
||
v_sf_shape = (m, n, v_sf_dim)
|
||
self.k_scale_buffer = [
|
||
torch.zeros(
|
||
k_sf_shape, dtype=torch.float8_e8m0fnu, device=self.device
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
self.v_scale_buffer = [
|
||
torch.zeros(
|
||
v_sf_shape, dtype=torch.float8_e8m0fnu, device=self.device
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
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
|
||
for x in self.k_buffer + self.v_buffer
|
||
],
|
||
device=self.device,
|
||
)
|
||
# This override replaces the base allocation, so the PD-transfer
|
||
# descriptors for the packed data buffers are built here too.
|
||
self._kv_buffer_descs = self._build_kv_buffer_descs()
|
||
|
||
def _clear_buffers(self):
|
||
del self.k_buffer
|
||
del self.v_buffer
|
||
del self.k_scale_buffer
|
||
del self.v_scale_buffer
|
||
|
||
def _get_key_buffer(self, layer_id: int):
|
||
return self.k_buffer[layer_id - self.start_layer]
|
||
|
||
def _get_value_buffer(self, layer_id: int):
|
||
return self.v_buffer[layer_id - self.start_layer]
|
||
|
||
def get_kv_scale_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
idx = layer_id - self.start_layer
|
||
return self.k_scale_buffer[idx], self.v_scale_buffer[idx]
|
||
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc_info,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
k_scale: Optional[torch.Tensor] = None,
|
||
v_scale: Optional[torch.Tensor] = None,
|
||
layer_id_override: Optional[int] = None,
|
||
dcp_kv_mask: Optional[torch.Tensor] = None,
|
||
):
|
||
if dcp_kv_mask is not None:
|
||
raise NotImplementedError("MXFP8 KV cache does not support DCP KV masks.")
|
||
loc, _, _ = unwrap_write_loc(loc_info)
|
||
maybe_detect_oob(
|
||
loc, 0, self.size + self.page_size, "set_kv_buffer (MHA-MXFP8)"
|
||
)
|
||
layer_id = (
|
||
layer_id_override if layer_id_override is not None else layer.layer_id
|
||
)
|
||
idx = layer_id - self.start_layer
|
||
|
||
if k_scale is None or v_scale is None:
|
||
# Fused path (SGLANG_OPT_INKLING_MXFP8_FUSED_QUANT_STORE): the layer
|
||
# hands us bf16 K/V and one kernel quantizes + scatters the fp8
|
||
# payload and the interleaved UE8M0 scales.
|
||
if not self.mxfp8_sf_interleaved or cache_k.dtype == self.store_dtype:
|
||
raise ValueError("MXFP8 KV cache requires K and V scale tensors.")
|
||
from sglang.kernels.ops.quantization.mxfp8_quant import quant_store_kv_mxfp8
|
||
|
||
quant_store_kv_mxfp8(
|
||
cache_k,
|
||
cache_v,
|
||
loc,
|
||
self.k_buffer[idx],
|
||
self.v_buffer[idx],
|
||
self.k_scale_buffer[idx],
|
||
self.v_scale_buffer[idx],
|
||
page_size=self.page_size,
|
||
)
|
||
return
|
||
|
||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||
|
||
if get_is_capture_mode() and self.alt_stream is not None:
|
||
current_stream = self.device_module.current_stream()
|
||
self.alt_stream.wait_stream(current_stream)
|
||
self.k_buffer[idx][loc] = cache_k
|
||
self._write_scales(idx, loc, k_scale, v_scale)
|
||
with self.device_module.stream(self.alt_stream):
|
||
self.v_buffer[idx][loc] = cache_v
|
||
current_stream.wait_stream(self.alt_stream)
|
||
else:
|
||
self.k_buffer[idx][loc] = cache_k
|
||
self.v_buffer[idx][loc] = cache_v
|
||
self._write_scales(idx, loc, k_scale, v_scale)
|
||
|
||
def _write_scales(self, idx, loc, k_scale, v_scale):
|
||
"""Write per-token UE8M0 K/V scales — interleaved into the FA4
|
||
BlockScaledBasicChunk layout for page_size==128, flat otherwise."""
|
||
if self.mxfp8_sf_interleaved:
|
||
from sglang.kernels.ops.quantization.mxfp8_interleave_sf import (
|
||
store_sf_interleaved,
|
||
)
|
||
|
||
store_sf_interleaved(
|
||
k_scale, self.k_scale_buffer[idx], loc, page_size=self.page_size
|
||
)
|
||
store_sf_interleaved(
|
||
v_scale, self.v_scale_buffer[idx], loc, page_size=self.page_size
|
||
)
|
||
else:
|
||
self.k_scale_buffer[idx][loc] = k_scale
|
||
self.v_scale_buffer[idx][loc] = v_scale
|
||
|
||
def _read_sf_interleaved(self, sf_buf: torch.Tensor, loc: torch.Tensor):
|
||
"""Inverse of store_sf_interleaved: gather per-slot (T, head, sf_dim)
|
||
UE8M0 scales out of the interleaved BlockScaledBasicChunk buffer."""
|
||
num_pages, n = sf_buf.shape[0], sf_buf.shape[1]
|
||
sf_dim = sf_buf.shape[-1]
|
||
# (num_pages, n, page_size) as u32: 4 packed scales per u32.
|
||
buf_u32 = sf_buf.reshape(num_pages, n, -1).view(torch.int32)
|
||
off = loc % self.page_size
|
||
page = (loc // self.page_size).long()
|
||
chunk = self.page_size // self.MXFP8_SCALE_BLOCK_SIZE
|
||
ipos = (
|
||
(off % self.MXFP8_SCALE_BLOCK_SIZE) * chunk
|
||
+ (off // self.MXFP8_SCALE_BLOCK_SIZE)
|
||
).long()
|
||
heads = torch.arange(n, device=loc.device)
|
||
gathered = buf_u32[page[:, None], heads[None, :], ipos[:, None]] # (T, n) int32
|
||
return (
|
||
gathered.reshape(loc.shape[0], n, 1)
|
||
.view(torch.uint8)
|
||
.reshape(loc.shape[0], n, sf_dim)
|
||
.view(torch.float8_e8m0fnu)
|
||
)
|
||
|
||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
# The mamba extra_buffer allocator relocates KV rows during serving;
|
||
# scale rows must travel with their fp8 payload or dequant reads
|
||
# mismatched exponents.
|
||
if self.mxfp8_sf_interleaved:
|
||
from sglang.kernels.ops.quantization.mxfp8_interleave_sf import (
|
||
store_sf_interleaved,
|
||
)
|
||
|
||
for idx in range(self.layer_num):
|
||
self.k_buffer[idx][tgt_loc] = self.k_buffer[idx][src_loc]
|
||
self.v_buffer[idx][tgt_loc] = self.v_buffer[idx][src_loc]
|
||
k_sf = self._read_sf_interleaved(self.k_scale_buffer[idx], src_loc)
|
||
v_sf = self._read_sf_interleaved(self.v_scale_buffer[idx], src_loc)
|
||
store_sf_interleaved(
|
||
k_sf, self.k_scale_buffer[idx], tgt_loc, page_size=self.page_size
|
||
)
|
||
store_sf_interleaved(
|
||
v_sf, self.v_scale_buffer[idx], tgt_loc, page_size=self.page_size
|
||
)
|
||
else:
|
||
super().move_kv_cache(tgt_loc, src_loc)
|
||
for idx in range(self.layer_num):
|
||
self.k_scale_buffer[idx][tgt_loc] = self.k_scale_buffer[idx][src_loc]
|
||
self.v_scale_buffer[idx][tgt_loc] = self.v_scale_buffer[idx][src_loc]
|
||
|
||
def _read_scales(self, idx, loc):
|
||
"""Per-token UE8M0 K/V scales at ``loc``, inverse of ``_write_scales``."""
|
||
if self.mxfp8_sf_interleaved:
|
||
return (
|
||
self._read_sf_interleaved(self.k_scale_buffer[idx], loc),
|
||
self._read_sf_interleaved(self.v_scale_buffer[idx], loc),
|
||
)
|
||
return self.k_scale_buffer[idx][loc], self.v_scale_buffer[idx][loc]
|
||
|
||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||
# The scales travel with their fp8 payload; a restored slot dequantizes
|
||
# against mismatched exponents without them.
|
||
assert not self.use_hnd, (
|
||
"CPU KV offload indexes by slot (NHD); HND KV cache "
|
||
"(SGLANG_USE_HND_KVCACHE) is not supported with CPU offload yet."
|
||
)
|
||
current_platform.synchronize()
|
||
kv_cache_cpu = []
|
||
chunk_size = self.cpu_offloading_chunk_size
|
||
for layer_id in range(self.layer_num):
|
||
kv_cache_cpu.append([])
|
||
for i in range(0, len(indices), chunk_size):
|
||
chunk_indices = indices[i : i + chunk_size]
|
||
k_scale, v_scale = self._read_scales(layer_id, chunk_indices)
|
||
kv_cache_cpu[-1].append(
|
||
[
|
||
self.k_buffer[layer_id][chunk_indices].to(
|
||
"cpu", non_blocking=True
|
||
),
|
||
self.v_buffer[layer_id][chunk_indices].to(
|
||
"cpu", non_blocking=True
|
||
),
|
||
k_scale.to("cpu", non_blocking=True),
|
||
v_scale.to("cpu", non_blocking=True),
|
||
]
|
||
)
|
||
current_platform.synchronize()
|
||
return kv_cache_cpu
|
||
|
||
def load_cpu_copy(
|
||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||
):
|
||
assert not self.use_hnd, (
|
||
"CPU KV offload indexes by slot (NHD); HND KV cache "
|
||
"(SGLANG_USE_HND_KVCACHE) is not supported with CPU offload yet."
|
||
)
|
||
current_platform.synchronize()
|
||
device = self.k_buffer[0].device
|
||
chunk_size = self.cpu_offloading_chunk_size
|
||
for layer_id in range(self.layer_num):
|
||
for i in range(0, len(indices), chunk_size):
|
||
chunk_indices = indices[i : i + chunk_size]
|
||
k_cpu, v_cpu, k_scale_cpu, v_scale_cpu = kv_cache_cpu[layer_id][
|
||
i // chunk_size
|
||
]
|
||
assert k_cpu.shape[0] == v_cpu.shape[0] == len(chunk_indices)
|
||
self.k_buffer[layer_id][chunk_indices] = k_cpu.to(
|
||
device, non_blocking=True
|
||
)
|
||
self.v_buffer[layer_id][chunk_indices] = v_cpu.to(
|
||
device, non_blocking=True
|
||
)
|
||
self._write_scales(
|
||
layer_id,
|
||
chunk_indices,
|
||
k_scale_cpu.to(device, non_blocking=True),
|
||
v_scale_cpu.to(device, non_blocking=True),
|
||
)
|
||
current_platform.synchronize()
|
||
|
||
def get_kv_scale_buf_infos(self):
|
||
"""(ptrs, lens, item_lens) for the UE8M0 scale buffers, k then v.
|
||
|
||
The interleaved layout puts pages on the leading axis, so a page's
|
||
scales are one contiguous row; the flat layout is per slot.
|
||
"""
|
||
tensors = self.k_scale_buffer + self.v_scale_buffer
|
||
ptrs = [t.data_ptr() for t in tensors]
|
||
lens = [t.nbytes for t in tensors]
|
||
row_bytes = [t[0].nbytes for t in tensors]
|
||
if self.mxfp8_sf_interleaved:
|
||
item_lens = row_bytes
|
||
else:
|
||
item_lens = [rb * self.page_size for rb in row_bytes]
|
||
return ptrs, lens, item_lens
|
||
|
||
def set_kv_buffer_prefix_valid(self, *args, **kwargs):
|
||
raise NotImplementedError(
|
||
"prefix-valid commit is unsupported for MXFP8 KV cache "
|
||
"(it does not carry the scale buffers)."
|
||
)
|
||
|
||
def get_kv_size_bytes(self):
|
||
k_size_bytes = 0
|
||
v_size_bytes = 0
|
||
for k_cache in self.k_buffer:
|
||
k_size_bytes += get_tensor_size_bytes(k_cache)
|
||
for k_scale in self.k_scale_buffer:
|
||
k_size_bytes += get_tensor_size_bytes(k_scale)
|
||
for v_cache in self.v_buffer:
|
||
v_size_bytes += get_tensor_size_bytes(v_cache)
|
||
for v_scale in self.v_scale_buffer:
|
||
v_size_bytes += get_tensor_size_bytes(v_scale)
|
||
return k_size_bytes, v_size_bytes
|
||
|
||
|
||
class HybridLinearKVPool(KVCache):
|
||
"""KV cache with separate pools for full and linear attention layers."""
|
||
|
||
cpu_copy_carries_mamba = True
|
||
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
dtype: torch.dtype,
|
||
page_size: int,
|
||
head_num: int,
|
||
head_dim: int,
|
||
full_attention_layer_ids: List[int],
|
||
device: str,
|
||
mamba_pool: MambaPool,
|
||
enable_memory_saver: bool = False,
|
||
enable_kv_cache_copy: bool = False,
|
||
# TODO: refactor mla related args
|
||
use_mla: bool = False,
|
||
kv_lora_rank: int = None,
|
||
qk_rope_head_dim: int = None,
|
||
use_dsa: bool = False,
|
||
index_head_dim: Optional[int] = None,
|
||
kv_cache_dim: Optional[int] = None,
|
||
index_kpool: int = 1,
|
||
index_kpool_compress: bool = False,
|
||
tail_extra_slots: int = 0,
|
||
max_running_requests: Optional[int] = None,
|
||
skip_topk_layers: Optional[List[bool]] = 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,
|
||
post_capture_active: bool = False,
|
||
):
|
||
self.size = size
|
||
self.dtype = dtype
|
||
self.device = device
|
||
self.full_layer_nums = len(full_attention_layer_ids)
|
||
self.page_size = page_size
|
||
self.start_layer = start_layer if start_layer is not None else 0
|
||
self.layer_transfer_counter = None
|
||
self.head_num = head_num
|
||
self.head_dim = head_dim
|
||
self.mamba_pool = mamba_pool
|
||
# Identity even though the unified pool holds VIRTUAL mamba ids: its
|
||
# composite allocator implements neither `get_cpu_copy` nor
|
||
# `load_cpu_copy`, the only readers, so those ids never arrive here.
|
||
self._mamba_translate = lambda ids: ids
|
||
self.use_mla = use_mla
|
||
self.use_dsa = use_dsa
|
||
if full_kv_pool is not None:
|
||
# Shared-KV-pool path: the caller built a UnifiedMHATokenToKVPool
|
||
# aliasing the shared byte buffer.
|
||
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 {}
|
||
)
|
||
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=enable_memory_saver,
|
||
enable_kv_cache_copy=enable_kv_cache_copy,
|
||
**quant_method_kwarg,
|
||
**post_capture_kwargs,
|
||
)
|
||
elif use_dsa:
|
||
# DSA sparse full-attention layers share the MLA latent layout and
|
||
# additionally keep a paged index_k cache. Only full-attn layer count
|
||
# is allocated here; the wrapper translates global layer_id to dense.
|
||
assert index_head_dim is not None and kv_cache_dim is not None, (
|
||
"HybridLinearKVPool with use_dsa requires index_head_dim and kv_cache_dim"
|
||
)
|
||
self.full_kv_pool = DSATokenToKVPool(
|
||
size=size,
|
||
page_size=self.page_size,
|
||
kv_lora_rank=kv_lora_rank,
|
||
dtype=dtype,
|
||
qk_rope_head_dim=qk_rope_head_dim,
|
||
layer_num=self.full_layer_nums,
|
||
device=device,
|
||
index_head_dim=index_head_dim,
|
||
enable_memory_saver=enable_memory_saver,
|
||
kv_cache_dim=kv_cache_dim,
|
||
index_kpool=index_kpool,
|
||
index_kpool_compress=index_kpool_compress,
|
||
tail_extra_slots=tail_extra_slots,
|
||
max_running_requests=max_running_requests,
|
||
skip_topk_layers=skip_topk_layers,
|
||
)
|
||
else:
|
||
TokenToKVPoolClass = MLATokenToKVPool
|
||
|
||
if current_platform.is_out_of_tree():
|
||
TokenToKVPoolClass = current_platform.get_mla_kv_pool_cls()
|
||
elif _is_npu:
|
||
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
||
NPUMLATokenToKVPool,
|
||
)
|
||
|
||
TokenToKVPoolClass = NPUMLATokenToKVPool
|
||
|
||
self.full_kv_pool = TokenToKVPoolClass(
|
||
size=size,
|
||
page_size=self.page_size,
|
||
dtype=dtype,
|
||
layer_num=self.full_layer_nums,
|
||
device=device,
|
||
kv_lora_rank=kv_lora_rank,
|
||
qk_rope_head_dim=qk_rope_head_dim,
|
||
enable_memory_saver=enable_memory_saver,
|
||
)
|
||
self.full_attention_layer_id_mapping = {
|
||
id: i for i, id in enumerate(full_attention_layer_ids)
|
||
}
|
||
if use_mla:
|
||
self.mem_usage = self.get_kv_size_bytes() / GB
|
||
else:
|
||
k_size, v_size = self.get_kv_size_bytes()
|
||
self.mem_usage = (k_size + v_size) / GB
|
||
|
||
@property
|
||
def post_capture_active(self) -> bool:
|
||
return self.full_kv_pool.post_capture_active
|
||
|
||
@property
|
||
def post_capture_backed_bytes(self) -> int:
|
||
return self.full_kv_pool.post_capture_backed_bytes
|
||
|
||
def finalize_backing(self, config) -> None:
|
||
# Only the attention KV is resized; the mamba state cache is fixed pre-capture.
|
||
self.full_kv_pool._finalize_backing_tokens(config.max_total_num_tokens)
|
||
self.size = int(config.max_total_num_tokens)
|
||
|
||
@property
|
||
def dsa_kv_cache_store_fp8(self) -> bool:
|
||
return getattr(self.full_kv_pool, "dsa_kv_cache_store_fp8", False)
|
||
|
||
@property
|
||
def kv_cache_dim(self):
|
||
return getattr(self.full_kv_pool, "kv_cache_dim", None)
|
||
|
||
@property
|
||
def index_head_dim(self) -> Optional[int]:
|
||
return getattr(self.full_kv_pool, "index_head_dim", None)
|
||
|
||
@property
|
||
def quant_block_size(self) -> Optional[int]:
|
||
return getattr(self.full_kv_pool, "quant_block_size", None)
|
||
|
||
@property
|
||
def index_kpool(self) -> int:
|
||
return getattr(self.full_kv_pool, "index_kpool", 1)
|
||
|
||
@property
|
||
def index_kpool_compress(self) -> bool:
|
||
return bool(getattr(self.full_kv_pool, "index_kpool_compress", False))
|
||
|
||
@property
|
||
def tail_extra_slots(self) -> int:
|
||
return getattr(self.full_kv_pool, "tail_extra_slots", 0)
|
||
|
||
@property
|
||
def slots_per_page(self) -> int:
|
||
return getattr(self.full_kv_pool, "slots_per_page", self.page_size)
|
||
|
||
def get_kv_size_bytes(self):
|
||
return self.full_kv_pool.get_kv_size_bytes()
|
||
|
||
def get_kv_buffer_shape(self) -> Tuple[torch.Size, torch.Size]:
|
||
# Hybrid layer ids are global model-layer ids, while the backing pool
|
||
# is dense over only full-attention layers. Shape discovery does not
|
||
# need a global layer lookup, so delegate it to that backing pool.
|
||
return self.full_kv_pool.get_kv_buffer_shape()
|
||
|
||
def get_contiguous_buf_infos(self):
|
||
return self.full_kv_pool.get_contiguous_buf_infos()
|
||
|
||
def get_kv_layer_ids(self):
|
||
"""Global layer ids aligned with the full-attention KV buffers."""
|
||
layer_ids = list(self.full_attention_layer_id_mapping)
|
||
if self.use_mla and _is_npu and layer_ids:
|
||
data_ptrs, _, _ = self.get_contiguous_buf_infos()
|
||
return layer_ids * (len(data_ptrs) // len(layer_ids))
|
||
return layer_ids if self.use_mla else layer_ids * 2
|
||
|
||
def get_state_buf_infos(self):
|
||
mamba_data_ptrs, mamba_data_lens, mamba_item_lens = (
|
||
self.mamba_pool.get_contiguous_buf_infos()
|
||
)
|
||
return mamba_data_ptrs, mamba_data_lens, mamba_item_lens
|
||
|
||
def get_state_dim_per_tensor(self):
|
||
"""Get the sliceable dimension size for each mamba state tensor."""
|
||
return self.mamba_pool.get_state_dim_per_tensor()
|
||
|
||
def get_state_layer_ids(self):
|
||
"""Global layer id per mamba state entry, aligned with get_state_buf_infos()."""
|
||
return self.mamba_pool.get_state_layer_ids()
|
||
|
||
def get_state_slice_outer_counts(self):
|
||
"""Get the row count preceding each mamba state slice axis."""
|
||
return self.mamba_pool.get_state_slice_outer_counts()
|
||
|
||
def get_state_conv_shard_groups(self):
|
||
"""Per-tensor conv sub-block dims (GDN) aligned with the state list."""
|
||
return self.mamba_pool.get_state_conv_shard_groups()
|
||
|
||
def maybe_get_custom_mem_pool(self):
|
||
return self.full_kv_pool.maybe_get_custom_mem_pool()
|
||
|
||
def _transfer_full_attention_id(self, layer_id: int):
|
||
if layer_id not in self.full_attention_layer_id_mapping:
|
||
raise ValueError(
|
||
f"{layer_id=} not in full attention layers: {self.full_attention_layer_id_mapping.keys()}"
|
||
)
|
||
return self.full_attention_layer_id_mapping[layer_id]
|
||
|
||
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
|
||
self.layer_transfer_counter = layer_transfer_counter
|
||
# The layer-wise wait logic is executed at the Hybrid LinearPool level;
|
||
# no additional wait is needed in the full_kv_pool
|
||
self.full_kv_pool.register_layer_transfer_counter(None)
|
||
|
||
def _wait_for_layer(self, layer_id: int):
|
||
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, 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, 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):
|
||
self._wait_for_layer(layer_id)
|
||
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
|
||
)
|
||
|
||
def get_kv_scale_buffer(self, layer_id: int):
|
||
# MXFP8 full_kv_pool exposes per-32 UE8M0 K/V scale buffers.
|
||
self._wait_for_layer(layer_id)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
return self.full_kv_pool.get_kv_scale_buffer(layer_id)
|
||
|
||
@contextmanager
|
||
def _transfer_id_context(self, layer: RadixAttention):
|
||
@contextmanager
|
||
def _patch_layer_id(layer):
|
||
original_layer_id = layer.layer_id
|
||
layer.layer_id = self._transfer_full_attention_id(layer.layer_id)
|
||
try:
|
||
yield
|
||
finally:
|
||
layer.layer_id = original_layer_id
|
||
|
||
with _patch_layer_id(layer):
|
||
yield
|
||
|
||
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,
|
||
dcp_kv_mask: Optional[torch.Tensor] = None,
|
||
):
|
||
# Write-location info lives in the metadata (`KVWriteLoc`). `full_loc` is the
|
||
# unified pool's pre-translated PHYSICAL loc (None for a static pool, where
|
||
# `loc` is already physical) — either way the pool writes a PHYSICAL loc.
|
||
loc, _, full_loc = unwrap_write_loc(loc)
|
||
layer_id = self._transfer_full_attention_id(layer.layer_id)
|
||
if not self.use_mla:
|
||
write_loc = full_loc if full_loc is not None else loc
|
||
self.full_kv_pool.set_kv_buffer(
|
||
layer,
|
||
write_loc,
|
||
cache_k,
|
||
cache_v,
|
||
k_scale,
|
||
v_scale,
|
||
layer_id_override=layer_id,
|
||
dcp_kv_mask=dcp_kv_mask,
|
||
)
|
||
else:
|
||
# Mirror the MHA branch: `full_loc` is the unified pool's
|
||
# pre-translated (kernel-facing) loc; None for a static pool.
|
||
write_loc = full_loc if full_loc is not None else loc
|
||
with self._transfer_id_context(layer):
|
||
self.full_kv_pool.set_kv_buffer(
|
||
layer,
|
||
write_loc,
|
||
cache_k,
|
||
cache_v,
|
||
)
|
||
|
||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
self.full_kv_pool.move_kv_cache(tgt_loc, src_loc)
|
||
|
||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||
kv_cpu = self.full_kv_pool.get_cpu_copy(indices, req_pool_index=req_pool_index)
|
||
# mamba_pool stores PHYSICAL ids; translate the (unified-pool virtual) ids first.
|
||
mamba_cpu = (
|
||
self.mamba_pool.get_cpu_copy(self._mamba_translate(mamba_indices))
|
||
if mamba_indices is not None
|
||
else None
|
||
)
|
||
return kv_cpu, mamba_cpu
|
||
|
||
def load_cpu_copy(
|
||
self, cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||
):
|
||
kv_cpu, mamba_cpu = cache_cpu
|
||
self.full_kv_pool.load_cpu_copy(kv_cpu, indices, req_pool_index=req_pool_index)
|
||
if mamba_cpu is not None and mamba_indices is not None:
|
||
self.mamba_pool.load_cpu_copy(
|
||
mamba_cpu, self._mamba_translate(mamba_indices)
|
||
)
|
||
|
||
def get_v_head_dim(self):
|
||
# Use start_layer to handle pipeline parallelism where layer 0
|
||
# may not be present in this stage's buffer.
|
||
return self.full_kv_pool.get_value_buffer(self.full_kv_pool.start_layer).shape[
|
||
-1
|
||
]
|
||
|
||
def set_mla_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_k_nope: torch.Tensor,
|
||
cache_k_rope: torch.Tensor,
|
||
):
|
||
assert self.use_mla, "set_mla_kv_buffer called when use_mla is False"
|
||
with self._transfer_id_context(layer):
|
||
self.full_kv_pool.set_mla_kv_buffer(layer, loc, cache_k_nope, cache_k_rope)
|
||
|
||
def get_mla_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
dst_dtype: Optional[torch.dtype] = None,
|
||
):
|
||
assert self.use_mla, "get_mla_kv_buffer called when use_mla is False"
|
||
# Read door -- same kernel-facing contract as the write door: `loc` is
|
||
# a read-index tensor already translated at its production site
|
||
# (fetch_mha_one_shot_kv_indices / prepare_chunked_kv_indices); the
|
||
# pool never translates.
|
||
with self._transfer_id_context(layer):
|
||
return self.full_kv_pool.get_mla_kv_buffer(layer, loc, dst_dtype)
|
||
|
||
def set_index_k_scale_buffer(
|
||
self,
|
||
layer_id: int,
|
||
loc: torch.Tensor,
|
||
index_k: torch.Tensor,
|
||
index_k_scale: torch.Tensor,
|
||
) -> None:
|
||
assert self.use_dsa, "set_index_k_scale_buffer called when use_dsa is False"
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
self.full_kv_pool.set_index_k_scale_buffer(
|
||
layer_id, loc, index_k, index_k_scale
|
||
)
|
||
|
||
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||
assert self.use_dsa, (
|
||
"get_index_k_with_scale_buffer called when use_dsa is False"
|
||
)
|
||
self._wait_for_layer(layer_id)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
return self.full_kv_pool.get_index_k_with_scale_buffer(layer_id)
|
||
|
||
def get_broadcastable_index_k_with_scale_buffer(
|
||
self, layer_id: int
|
||
) -> torch.Tensor:
|
||
assert self.use_dsa, (
|
||
"get_broadcastable_index_k_with_scale_buffer called when use_dsa is False"
|
||
)
|
||
self._wait_for_layer(layer_id)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
if hasattr(self.full_kv_pool, "_get_broadcastable_index_buffer"):
|
||
return self.full_kv_pool._get_broadcastable_index_buffer(layer_id)
|
||
return self.full_kv_pool.get_index_k_with_scale_buffer(layer_id)
|
||
|
||
def invalidate_index_buffer_for_layer(self, layer_id: int) -> None:
|
||
if not self.use_dsa or not hasattr(
|
||
self.full_kv_pool, "invalidate_index_buffer_for_layer"
|
||
):
|
||
return
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
self.full_kv_pool.invalidate_index_buffer_for_layer(layer_id)
|
||
|
||
def get_index_k_continuous(
|
||
self,
|
||
layer_id: int,
|
||
seq_len: int,
|
||
page_indices: torch.Tensor,
|
||
):
|
||
assert self.use_dsa, "get_index_k_continuous called when use_dsa is False"
|
||
self._wait_for_layer(layer_id)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
return self.full_kv_pool.get_index_k_continuous(layer_id, seq_len, page_indices)
|
||
|
||
def get_index_k_scale_continuous(
|
||
self,
|
||
layer_id: int,
|
||
seq_len: int,
|
||
page_indices: torch.Tensor,
|
||
):
|
||
assert self.use_dsa, "get_index_k_scale_continuous called when use_dsa is False"
|
||
self._wait_for_layer(layer_id)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
return self.full_kv_pool.get_index_k_scale_continuous(
|
||
layer_id, seq_len, page_indices
|
||
)
|
||
|
||
def get_index_k_scale_buffer(
|
||
self,
|
||
layer_id: int,
|
||
seq_len_tensor: torch.Tensor,
|
||
page_indices: torch.Tensor,
|
||
seq_len_sum: int,
|
||
max_seq_len: int,
|
||
):
|
||
assert self.use_dsa, "get_index_k_scale_buffer called when use_dsa is False"
|
||
self._wait_for_layer(layer_id)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
return self.full_kv_pool.get_index_k_scale_buffer(
|
||
layer_id, seq_len_tensor, page_indices, seq_len_sum, max_seq_len
|
||
)
|
||
|
||
def get_compress_tail_buffers(
|
||
self, layer_id: int
|
||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
assert self.use_dsa, "get_compress_tail_buffers called when use_dsa is False"
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
return self.full_kv_pool.get_compress_tail_buffers(layer_id)
|
||
|
||
def kpool_decode_update_index_cache(
|
||
self,
|
||
layer_id: int,
|
||
key: torch.Tensor,
|
||
slot_score: torch.Tensor,
|
||
ape: torch.Tensor,
|
||
block_tables: torch.Tensor,
|
||
req_pool_indices: torch.Tensor,
|
||
positions: torch.Tensor,
|
||
seq_lens: torch.Tensor,
|
||
out_cache_loc: torch.Tensor,
|
||
round_scale: bool = False,
|
||
) -> None:
|
||
assert self.use_dsa, (
|
||
"kpool_decode_update_index_cache called when use_dsa is False"
|
||
)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
self.full_kv_pool.kpool_decode_update_index_cache(
|
||
layer_id=layer_id,
|
||
key=key,
|
||
slot_score=slot_score,
|
||
ape=ape,
|
||
block_tables=block_tables,
|
||
req_pool_indices=req_pool_indices,
|
||
positions=positions,
|
||
seq_lens=seq_lens,
|
||
out_cache_loc=out_cache_loc,
|
||
round_scale=round_scale,
|
||
)
|
||
|
||
def set_compress_tail_for_request(
|
||
self,
|
||
layer_id: int,
|
||
req_pool_idx: torch.Tensor,
|
||
key_tail: torch.Tensor,
|
||
score_tail: torch.Tensor,
|
||
n_remain: int,
|
||
dst_logical_start: int,
|
||
) -> None:
|
||
assert self.use_dsa, (
|
||
"set_compress_tail_for_request called when use_dsa is False"
|
||
)
|
||
layer_id = self._transfer_full_attention_id(layer_id)
|
||
self.full_kv_pool.set_compress_tail_for_request(
|
||
layer_id=layer_id,
|
||
req_pool_idx=req_pool_idx,
|
||
key_tail=key_tail,
|
||
score_tail=score_tail,
|
||
n_remain=n_remain,
|
||
dst_logical_start=dst_logical_start,
|
||
)
|
||
|
||
|
||
class MLATokenToKVPool(KVCache):
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
page_size: int,
|
||
dtype: torch.dtype,
|
||
kv_lora_rank: int,
|
||
qk_rope_head_dim: int,
|
||
layer_num: int,
|
||
device: str,
|
||
enable_memory_saver: bool,
|
||
start_layer: Optional[int] = None,
|
||
end_layer: Optional[int] = None,
|
||
use_dsa: bool = False,
|
||
override_kv_cache_dim: Optional[int] = None,
|
||
):
|
||
super().__init__(
|
||
size,
|
||
page_size,
|
||
dtype,
|
||
layer_num,
|
||
device,
|
||
enable_memory_saver,
|
||
start_layer,
|
||
end_layer,
|
||
)
|
||
|
||
self.kv_lora_rank = kv_lora_rank
|
||
self.qk_rope_head_dim = qk_rope_head_dim
|
||
self.use_dsa = use_dsa
|
||
self.dsa_kv_cache_store_fp8 = (
|
||
use_dsa
|
||
and dtype == torch.float8_e4m3fn
|
||
and override_kv_cache_dim is not None
|
||
)
|
||
# When override_kv_cache_dim is provided with dsa model, we assume the
|
||
# override kv cache dim is correct and use it directly.
|
||
self.kv_cache_dim = (
|
||
override_kv_cache_dim
|
||
if self.dsa_kv_cache_store_fp8
|
||
else (kv_lora_rank + qk_rope_head_dim)
|
||
)
|
||
|
||
self._create_buffers()
|
||
|
||
self.data_ptrs = torch.tensor(
|
||
[x.data_ptr() for x in self.kv_buffer],
|
||
dtype=torch.uint64,
|
||
device=self.device,
|
||
)
|
||
if not use_dsa:
|
||
# DSA will allocate indexer KV cache later and then log the total size
|
||
self._finalize_allocation_log(size)
|
||
|
||
def _create_buffers(self):
|
||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||
with (
|
||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||
if self.custom_mem_pool
|
||
else nullcontext()
|
||
):
|
||
# The padded slot 0 is used for writing dummy outputs from padded tokens.
|
||
self.kv_buffer = [
|
||
torch.zeros(
|
||
(self.size + self.page_size, 1, self.kv_cache_dim),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
def _clear_buffers(self):
|
||
del self.kv_buffer
|
||
|
||
def get_kv_size_bytes(self):
|
||
assert hasattr(self, "kv_buffer")
|
||
kv_size_bytes = 0
|
||
for kv_cache in self.kv_buffer:
|
||
kv_size_bytes += get_tensor_size_bytes(kv_cache)
|
||
return kv_size_bytes
|
||
|
||
# for disagg
|
||
def get_contiguous_buf_infos(self):
|
||
# MLA has only one kv_buffer, so only the information of this buffer needs to be returned.
|
||
kv_data_ptrs = [self.kv_buffer[i].data_ptr() for i in range(self.layer_num)]
|
||
kv_data_lens = [self.kv_buffer[i].nbytes for i in range(self.layer_num)]
|
||
kv_item_lens = [
|
||
self.kv_buffer[i][0].nbytes * self.page_size for i in range(self.layer_num)
|
||
]
|
||
return kv_data_ptrs, kv_data_lens, kv_item_lens
|
||
|
||
def get_key_buffer(self, layer_id: int):
|
||
if self.layer_transfer_counter is not None:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
|
||
if self.store_dtype != self.dtype:
|
||
return self.kv_buffer[layer_id - self.start_layer].view(self.dtype)
|
||
|
||
return self.kv_buffer[layer_id - self.start_layer]
|
||
|
||
def get_value_buffer(self, layer_id: int):
|
||
if self.layer_transfer_counter is not None:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
|
||
if self.store_dtype != self.dtype:
|
||
return self.kv_buffer[layer_id - self.start_layer][
|
||
..., : self.kv_lora_rank
|
||
].view(self.dtype)
|
||
return self.kv_buffer[layer_id - self.start_layer][..., : self.kv_lora_rank]
|
||
|
||
def get_kv_buffer(self, layer_id: int):
|
||
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
|
||
|
||
# Has the WRITE loc arriving here already had the DCP owner rule resolved?
|
||
# False: this pool takes a WIDENED loc. The unified pool resolves it in
|
||
# `KVIndexTranslator.rebind_write_loc` and flips this. Not derivable from
|
||
# `kernel_page_blocks`: that is `layer_num`, so a rank owning one
|
||
# full-attention layer is translated with blocks_per_page 1.
|
||
write_loc_is_dcp_resolved = False
|
||
|
||
@property
|
||
def _write_loc_dcp_span(self) -> int:
|
||
"""How many logical ids one stored row spans in the write-loc space."""
|
||
return 1 if self.write_loc_is_dcp_resolved else get_parallel().attn_dcp_size
|
||
|
||
def _scatter_mla_rows(
|
||
self,
|
||
dst_buffer: torch.Tensor,
|
||
loc: torch.Tensor,
|
||
cache_k_nope: torch.Tensor,
|
||
cache_k_rope: torch.Tensor,
|
||
) -> None:
|
||
if self.write_loc_is_dcp_resolved:
|
||
set_mla_kv_buffer_triton(dst_buffer, loc, cache_k_nope, cache_k_rope)
|
||
else:
|
||
set_mla_kv_buffer_dcp_sharded_triton(
|
||
dst_buffer, loc, cache_k_nope, cache_k_rope
|
||
)
|
||
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc_info,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
layer_id_override: Optional[int] = None,
|
||
):
|
||
loc, _, _ = unwrap_write_loc(loc_info)
|
||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
|
||
maybe_detect_kernel_facing_loc(
|
||
loc, self.page_size, self.kernel_page_blocks, "set_kv_buffer (MLA)"
|
||
)
|
||
layer_id = (
|
||
layer_id_override if layer_id_override is not None else layer.layer_id
|
||
)
|
||
assert not self.dsa_kv_cache_store_fp8
|
||
# No DCP-aware variant is possible: the two backends reaching this door
|
||
# disagree on the loc space (flashinfer-MLA widened, Triton collapsed).
|
||
assert self.write_loc_is_dcp_resolved or not get_parallel().dcp_enabled, (
|
||
"MLATokenToKVPool.set_kv_buffer has no DCP-aware write path. Under "
|
||
"--dcp-size > 1 the MLA write must go through set_mla_kv_buffer, "
|
||
"whose kernel resolves the owner rule; reaching the combined-row "
|
||
"door means an attention backend took a write path that never "
|
||
"declared which loc space it emits."
|
||
)
|
||
if cache_k.dtype != self.dtype:
|
||
cache_k = cache_k.to(self.dtype)
|
||
|
||
if self.store_dtype != self.dtype:
|
||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k.view(
|
||
self.store_dtype
|
||
)
|
||
else:
|
||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
|
||
|
||
def _write_mla_kv_buffer(
|
||
self,
|
||
dst_buffer: torch.Tensor,
|
||
loc: torch.Tensor,
|
||
cache_k_nope: torch.Tensor,
|
||
cache_k_rope: torch.Tensor,
|
||
) -> None:
|
||
assert not (
|
||
self.write_loc_is_dcp_resolved
|
||
and (self.use_dsa or self.dsa_kv_cache_store_fp8)
|
||
), "the DSA write paths have no resolved-loc variant"
|
||
if _is_hip and self.use_dsa and self.dtype == fp8_dtype:
|
||
# HIP FP8 path uses raw MLA KV layout (nope + rope) without per-block scales.
|
||
# Fuse BF16/FP16 -> FP8 cast with paged KV write.
|
||
set_mla_kv_buffer_triton_fp8_quant(
|
||
dst_buffer,
|
||
loc,
|
||
cache_k_nope,
|
||
cache_k_rope,
|
||
fp8_dtype,
|
||
)
|
||
elif self.dsa_kv_cache_store_fp8:
|
||
# OPTIMIZATION: Quantize k_nope and k_rope separately to avoid concat overhead
|
||
# This also enables reuse of set_mla_kv_buffer_triton two-tensor write path
|
||
# quantize_k_cache_separate returns (nope_part, rope_part) as uint8 bytes
|
||
cache_k_nope_fp8, cache_k_rope_fp8 = quantize_k_cache_separate(
|
||
cache_k_nope, cache_k_rope
|
||
)
|
||
|
||
# Reuse existing two-tensor write kernel (works with FP8 byte layout)
|
||
# cache_k_nope_fp8: (num_tokens, 1, 528) uint8 [nope_fp8(512) | scales(16)]
|
||
# cache_k_rope_fp8: (num_tokens, 1, 128) uint8 [rope_bf16_bytes(128)]
|
||
self._scatter_mla_rows(dst_buffer, loc, cache_k_nope_fp8, cache_k_rope_fp8)
|
||
else:
|
||
if cache_k_nope.dtype != self.dtype:
|
||
cache_k_nope = cache_k_nope.to(self.dtype)
|
||
if cache_k_rope is not None and cache_k_rope.numel() > 0:
|
||
cache_k_rope = cache_k_rope.to(self.dtype)
|
||
if self.store_dtype != self.dtype:
|
||
cache_k_nope = cache_k_nope.view(self.store_dtype)
|
||
if cache_k_rope is not None and cache_k_rope.numel() > 0:
|
||
cache_k_rope = cache_k_rope.view(self.store_dtype)
|
||
|
||
self._scatter_mla_rows(dst_buffer, loc, cache_k_nope, cache_k_rope)
|
||
|
||
def set_mla_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_k_nope: torch.Tensor,
|
||
cache_k_rope: torch.Tensor,
|
||
layer_id_override: Optional[int] = None,
|
||
):
|
||
# loc is widened under DCP unless the pool declares it resolved.
|
||
maybe_detect_oob(
|
||
loc,
|
||
0,
|
||
(self.size + self.page_size) * self._write_loc_dcp_span,
|
||
"set_mla_kv_buffer (MLA)",
|
||
)
|
||
maybe_detect_kernel_facing_loc(
|
||
loc, self.page_size, self.kernel_page_blocks, "set_mla_kv_buffer (MLA)"
|
||
)
|
||
layer_id = (
|
||
layer_id_override if layer_id_override is not None else layer.layer_id
|
||
)
|
||
self._write_mla_kv_buffer(
|
||
self.kv_buffer[layer_id - self.start_layer],
|
||
loc,
|
||
cache_k_nope,
|
||
cache_k_rope,
|
||
)
|
||
|
||
def get_mla_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
dst_dtype: Optional[torch.dtype] = None,
|
||
):
|
||
# get k nope and k rope from the kv buffer, and optionally cast them to dst_dtype.
|
||
layer_id = layer.layer_id
|
||
kv_buffer = self.get_key_buffer(layer_id)
|
||
dst_dtype = dst_dtype or self.dtype
|
||
cache_k_nope = torch.empty(
|
||
(loc.shape[0], 1, self.kv_lora_rank),
|
||
dtype=dst_dtype,
|
||
device=kv_buffer.device,
|
||
)
|
||
if self.qk_rope_head_dim == 0:
|
||
cache_k_rope = None
|
||
else:
|
||
cache_k_rope = torch.empty(
|
||
(loc.shape[0], 1, self.qk_rope_head_dim),
|
||
dtype=dst_dtype,
|
||
device=kv_buffer.device,
|
||
)
|
||
get_mla_kv_buffer_triton(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
||
return cache_k_nope, cache_k_rope
|
||
|
||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
"""Relocate accepted-token combined MLA KV (latent + rope) per layer."""
|
||
size_limit = self.size + self.page_size
|
||
maybe_detect_oob(tgt_loc, 0, size_limit, "move_kv_cache tgt_loc")
|
||
maybe_detect_oob(src_loc, 0, size_limit, "move_kv_cache src_loc")
|
||
|
||
if tgt_loc.numel() == 0:
|
||
return
|
||
|
||
tgt_loc_flat = tgt_loc.view(-1).long()
|
||
src_loc_flat = src_loc.view(-1).long()
|
||
for kv_cache in self.kv_buffer:
|
||
kv_cache[tgt_loc_flat] = kv_cache[src_loc_flat]
|
||
|
||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||
indices = maybe_dcp_kernel_indices(
|
||
indices, self._write_loc_dcp_span, get_parallel().attn_dcp_rank
|
||
)
|
||
current_platform.synchronize()
|
||
kv_cache_cpu = []
|
||
chunk_size = self.cpu_offloading_chunk_size
|
||
for layer_id in range(self.layer_num):
|
||
kv_cache_cpu.append([])
|
||
if self.kv_buffer[layer_id].shape[0] == 0:
|
||
continue
|
||
for i in range(0, len(indices), chunk_size):
|
||
chunk_indices = indices[i : i + chunk_size]
|
||
kv_cpu = self.kv_buffer[layer_id][chunk_indices].to(
|
||
"cpu", non_blocking=True
|
||
)
|
||
kv_cache_cpu[-1].append(kv_cpu)
|
||
current_platform.synchronize()
|
||
return kv_cache_cpu
|
||
|
||
def load_cpu_copy(
|
||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||
):
|
||
indices = maybe_dcp_kernel_indices(
|
||
indices, self._write_loc_dcp_span, get_parallel().attn_dcp_rank
|
||
)
|
||
current_platform.synchronize()
|
||
chunk_size = self.cpu_offloading_chunk_size
|
||
for layer_id in range(self.layer_num):
|
||
if self.kv_buffer[layer_id].shape[0] == 0:
|
||
continue
|
||
for i in range(0, len(indices), chunk_size):
|
||
chunk_indices = indices[i : i + chunk_size]
|
||
kv_cpu = kv_cache_cpu[layer_id][i // chunk_size]
|
||
assert kv_cpu.shape[0] == len(chunk_indices)
|
||
kv_chunk = kv_cpu.to(self.kv_buffer[0].device, non_blocking=True)
|
||
self.kv_buffer[layer_id][chunk_indices] = kv_chunk
|
||
current_platform.synchronize()
|
||
|
||
|
||
class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
||
def _create_buffers(self):
|
||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||
with (
|
||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||
if self.custom_mem_pool
|
||
else nullcontext()
|
||
):
|
||
# The padded slot 0 is used for writing dummy outputs from padded tokens.
|
||
m = self.size + self.page_size
|
||
n = 1 # head_num
|
||
k = self.kv_cache_dim # head_dim
|
||
|
||
scale_block_size = 16
|
||
self.store_dtype = torch.uint8
|
||
|
||
self.kv_buffer = [
|
||
torch.zeros(
|
||
(m, n, k // 2),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
self.kv_scale_buffer = [
|
||
torch.zeros(
|
||
(m, k // scale_block_size),
|
||
dtype=self.store_dtype,
|
||
device=self.device,
|
||
)
|
||
for _ in range(self.layer_num)
|
||
]
|
||
|
||
def _clear_buffers(self):
|
||
del self.kv_buffer
|
||
del self.kv_scale_buffer
|
||
|
||
def get_key_buffer(self, layer_id: int):
|
||
if self.layer_transfer_counter is not None:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
|
||
if self.store_dtype != self.dtype:
|
||
cache_k_nope_fp4 = self.kv_buffer[layer_id - self.start_layer].view(
|
||
torch.uint8
|
||
)
|
||
cache_k_nope_fp4_sf = self.kv_scale_buffer[layer_id - self.start_layer]
|
||
|
||
from sglang.srt.layers.quantization.kvfp4_tensor import (
|
||
FP4MXBlock16KVQuantizeUtil,
|
||
)
|
||
|
||
cache_k_nope_fp4_dequant = FP4MXBlock16KVQuantizeUtil.batched_dequantize(
|
||
cache_k_nope_fp4, cache_k_nope_fp4_sf
|
||
)
|
||
return cache_k_nope_fp4_dequant
|
||
|
||
return self.kv_buffer[layer_id - self.start_layer]
|
||
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc_info,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
):
|
||
# loc_info may be a KVWriteLoc; MLA pools have no SWA target.
|
||
loc, _, _ = unwrap_write_loc(loc_info)
|
||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA-FP4)")
|
||
layer_id = layer.layer_id
|
||
assert not self.dsa_kv_cache_store_fp8
|
||
if cache_k.dtype != self.dtype:
|
||
from sglang.srt.layers.quantization.kvfp4_tensor import (
|
||
FP4MXBlock16KVQuantizeUtil,
|
||
)
|
||
|
||
cache_k_fp4, cache_k_fp4_sf = FP4MXBlock16KVQuantizeUtil.batched_quantize(
|
||
cache_k
|
||
)
|
||
|
||
if self.store_dtype != self.dtype:
|
||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k_fp4.view(
|
||
self.store_dtype
|
||
)
|
||
self.kv_scale_buffer[layer_id - self.start_layer][loc] = (
|
||
cache_k_fp4_sf.view(self.store_dtype)
|
||
)
|
||
else:
|
||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
|
||
|
||
def set_mla_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_k_nope: torch.Tensor,
|
||
cache_k_rope: torch.Tensor,
|
||
):
|
||
maybe_detect_oob(
|
||
loc, 0, self.size + self.page_size, "set_mla_kv_buffer (MLA-FP4)"
|
||
)
|
||
layer_id = layer.layer_id
|
||
|
||
if self.dsa_kv_cache_store_fp8:
|
||
# original cache_k: (num_tokens, num_heads 1, hidden 576); we unsqueeze the page_size=1 dim here
|
||
# TODO no need to cat
|
||
cache_k = torch.cat([cache_k_nope, cache_k_rope], dim=-1)
|
||
cache_k = quantize_k_cache(cache_k.unsqueeze(1)).squeeze(1)
|
||
cache_k = cache_k.view(self.store_dtype)
|
||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
|
||
else:
|
||
if cache_k_nope.dtype != self.dtype:
|
||
from sglang.srt.layers.quantization.kvfp4_tensor import (
|
||
FP4MXBlock16KVQuantizeUtil,
|
||
)
|
||
|
||
cache_k_nope_fp4, cache_k_nope_fp4_sf = (
|
||
FP4MXBlock16KVQuantizeUtil.batched_quantize(cache_k_nope)
|
||
)
|
||
if cache_k_rope is not None and cache_k_rope.numel() > 0:
|
||
cache_k_rope_fp4, cache_k_rope_fp4_sf = (
|
||
FP4MXBlock16KVQuantizeUtil.batched_quantize(cache_k_rope)
|
||
)
|
||
else:
|
||
cache_k_rope_fp4 = None
|
||
cache_k_rope_fp4_sf = None
|
||
|
||
if self.store_dtype != self.dtype:
|
||
cache_k_nope = cache_k_nope.view(self.store_dtype)
|
||
if cache_k_rope is not None and cache_k_rope.numel() > 0:
|
||
cache_k_rope = cache_k_rope.view(self.store_dtype)
|
||
|
||
self._scatter_mla_rows(
|
||
self.kv_buffer[layer_id - self.start_layer],
|
||
loc,
|
||
cache_k_nope_fp4,
|
||
cache_k_rope_fp4,
|
||
)
|
||
set_mla_kv_scale_buffer_triton(
|
||
self.kv_scale_buffer[layer_id - self.start_layer],
|
||
loc,
|
||
cache_k_nope_fp4_sf,
|
||
cache_k_rope_fp4_sf,
|
||
)
|
||
|
||
|
||
class DSATokenToKVPool(MLATokenToKVPool):
|
||
quant_block_size = 128
|
||
index_k_with_scale_buffer_dtype = torch.uint8
|
||
rope_storage_dtype = torch.bfloat16 # rope is always stored in bf16
|
||
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
page_size: int,
|
||
kv_lora_rank: int,
|
||
dtype: torch.dtype,
|
||
qk_rope_head_dim: int,
|
||
layer_num: int,
|
||
device: str,
|
||
index_head_dim: int,
|
||
enable_memory_saver: bool,
|
||
kv_cache_dim: int,
|
||
start_layer: Optional[int] = None,
|
||
end_layer: Optional[int] = None,
|
||
index_buf_size: Optional[int] = None,
|
||
index_kpool: int = 1,
|
||
index_kpool_compress: bool = False,
|
||
tail_extra_slots: int = 0,
|
||
max_running_requests: Optional[int] = None,
|
||
skip_topk_layers: Optional[List[bool]] = None,
|
||
):
|
||
override_dim = (
|
||
kv_cache_dim if kv_cache_dim != kv_lora_rank + qk_rope_head_dim else None
|
||
)
|
||
|
||
super().__init__(
|
||
size,
|
||
page_size,
|
||
dtype,
|
||
kv_lora_rank,
|
||
qk_rope_head_dim,
|
||
layer_num,
|
||
device,
|
||
enable_memory_saver,
|
||
start_layer,
|
||
end_layer,
|
||
use_dsa=True,
|
||
override_kv_cache_dim=override_dim,
|
||
)
|
||
# self.index_k_dtype = torch.float8_e4m3fn
|
||
# self.index_k_scale_dtype = torch.float32
|
||
self.index_head_dim = index_head_dim
|
||
self.index_kpool = index_kpool
|
||
self.index_kpool_compress = index_kpool_compress
|
||
self.tail_extra_slots = tail_extra_slots
|
||
self.slots_per_page = self.page_size
|
||
if index_buf_size is None:
|
||
index_buf_size = size
|
||
self.index_buf_size = index_buf_size
|
||
# num head == 1 and head dim == 128 for index_k in DSA
|
||
assert index_head_dim == 128
|
||
|
||
self.skip_topk_layers = (
|
||
list(skip_topk_layers)
|
||
if skip_topk_layers is not None
|
||
else [False] * layer_num
|
||
)
|
||
assert len(self.skip_topk_layers) == layer_num
|
||
|
||
if _is_hip:
|
||
if aiter_can_use_preshuffle_paged_mqa():
|
||
assert self.page_size % 16 == 0, (
|
||
f"HIP preshuffle requires page_size to be a multiple of 16, got {self.page_size}"
|
||
)
|
||
else:
|
||
assert self.page_size == 1, (
|
||
f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
|
||
)
|
||
elif is_xpu():
|
||
assert self.page_size in (
|
||
64,
|
||
128,
|
||
), f"XPU DSA requires page_size 64 or 128, got {self.page_size}"
|
||
else:
|
||
assert self.page_size == 64
|
||
self.index_key_cache = self._create_index_key_cache()
|
||
self._init_kpool_compress_tail_buffers(
|
||
index_kpool=index_kpool,
|
||
index_kpool_compress=index_kpool_compress,
|
||
tail_extra_slots=tail_extra_slots,
|
||
index_head_dim=index_head_dim,
|
||
layer_num=layer_num,
|
||
device=device,
|
||
max_running_requests=max_running_requests,
|
||
)
|
||
self._finalize_allocation_log(size)
|
||
|
||
def _create_index_key_cache(self) -> IndexKeyCache:
|
||
return IndexKeyCache(self, self.index_buf_size)
|
||
|
||
def _should_allocate_index_layer(self, local_layer_idx: int) -> bool:
|
||
return not self.skip_topk_layers[local_layer_idx]
|
||
|
||
@property
|
||
def index_k_with_scale_buffer(self):
|
||
# Preserve direct HiCache access while storage lives behind the facade.
|
||
return self.index_key_cache.buffer
|
||
|
||
def _init_kpool_compress_tail_buffers(
|
||
self,
|
||
index_kpool: int,
|
||
index_kpool_compress: bool,
|
||
tail_extra_slots: int,
|
||
index_head_dim: int,
|
||
layer_num: int,
|
||
device: str,
|
||
max_running_requests: Optional[int],
|
||
) -> None:
|
||
"""Keep request tails on the pool so they follow the index-cache lifecycle."""
|
||
self.kpool_use_compress = index_kpool > 1 and index_kpool_compress
|
||
|
||
if not self.kpool_use_compress:
|
||
self._compress_tail_k = None
|
||
self._compress_tail_score = None
|
||
return
|
||
|
||
assert max_running_requests is not None, (
|
||
"DSATokenToKVPool with kpool compress requires max_running_requests"
|
||
)
|
||
# +1 mirrors req_to_token_pool.size + 1 used by the indexer to
|
||
# provide an extra slot for invalid / sentinel req indices.
|
||
req_pool_size = max_running_requests + 1
|
||
tail_dtype = torch.bfloat16
|
||
tail_width = index_kpool + tail_extra_slots
|
||
with (
|
||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||
if self.custom_mem_pool
|
||
else nullcontext()
|
||
):
|
||
self._compress_tail_k: Optional[List[torch.Tensor]] = [
|
||
torch.zeros(
|
||
req_pool_size if self._should_allocate_index_layer(i) else 0,
|
||
tail_width,
|
||
index_head_dim,
|
||
dtype=tail_dtype,
|
||
device=device,
|
||
)
|
||
for i in range(layer_num)
|
||
]
|
||
self._compress_tail_score: Optional[List[torch.Tensor]] = [
|
||
torch.zeros(
|
||
req_pool_size if self._should_allocate_index_layer(i) else 0,
|
||
tail_width,
|
||
index_head_dim,
|
||
dtype=tail_dtype,
|
||
device=device,
|
||
)
|
||
for i in range(layer_num)
|
||
]
|
||
|
||
def get_compress_tail_buffers(
|
||
self, layer_id: int
|
||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
assert self.kpool_use_compress, (
|
||
"get_compress_tail_buffers called when kpool compress is disabled"
|
||
)
|
||
idx = layer_id - self.start_layer
|
||
return (
|
||
self._compress_tail_k[idx],
|
||
self._compress_tail_score[idx],
|
||
)
|
||
|
||
def get_compress_tail_buf_infos(self):
|
||
if not self.kpool_use_compress:
|
||
return [], [], []
|
||
transfer_layer_ids = list(range(self.layer_num))
|
||
# Keep zero-row indexShare entries in the pointer list so layer offsets
|
||
# stay aligned across PD peers; item_len=0 makes transfer backends skip them.
|
||
tail_buffers = [self._compress_tail_k[i] for i in transfer_layer_ids] + [
|
||
self._compress_tail_score[i] for i in transfer_layer_ids
|
||
]
|
||
data_ptrs = [buf.data_ptr() for buf in tail_buffers]
|
||
data_lens = [buf.nbytes for buf in tail_buffers]
|
||
item_lens = [buf[0].nbytes if buf.shape[0] > 0 else 0 for buf in tail_buffers]
|
||
return data_ptrs, data_lens, item_lens
|
||
|
||
def kpool_decode_update_index_cache(
|
||
self,
|
||
layer_id: int,
|
||
key: torch.Tensor,
|
||
slot_score: torch.Tensor,
|
||
ape: torch.Tensor,
|
||
block_tables: torch.Tensor,
|
||
req_pool_indices: torch.Tensor,
|
||
positions: torch.Tensor,
|
||
seq_lens: torch.Tensor,
|
||
out_cache_loc: torch.Tensor,
|
||
round_scale: bool = False,
|
||
) -> None:
|
||
from sglang.srt.layers.attention.dsa.kpool_fp8_index import (
|
||
kpool_decode_update_and_maybe_write_cache,
|
||
)
|
||
|
||
assert self.kpool_use_compress, (
|
||
"kpool_decode_update_index_cache called when kpool compress is disabled"
|
||
)
|
||
idx = layer_id - self.start_layer
|
||
buf = self.get_index_k_with_scale_buffer(layer_id)
|
||
kpool_decode_update_and_maybe_write_cache(
|
||
pool=self,
|
||
buf=buf,
|
||
tail_k=self._compress_tail_k[idx],
|
||
tail_score=self._compress_tail_score[idx],
|
||
key=key,
|
||
slot_score=slot_score,
|
||
ape=ape,
|
||
block_tables=block_tables,
|
||
req_pool_indices=req_pool_indices,
|
||
positions=positions,
|
||
seq_lens=seq_lens,
|
||
out_cache_loc=out_cache_loc,
|
||
round_scale=round_scale,
|
||
)
|
||
|
||
def set_compress_tail_for_request(
|
||
self,
|
||
layer_id: int,
|
||
req_pool_idx: torch.Tensor,
|
||
key_tail: torch.Tensor,
|
||
score_tail: torch.Tensor,
|
||
n_remain: int,
|
||
dst_logical_start: int,
|
||
) -> None:
|
||
"""Leave the ring untouched at a pool boundary; no tail carries over."""
|
||
assert self.kpool_use_compress, (
|
||
"set_compress_tail_for_request called when kpool compress is disabled"
|
||
)
|
||
idx = layer_id - self.start_layer
|
||
if n_remain > 0:
|
||
slots = (
|
||
torch.arange(n_remain, device=key_tail.device, dtype=torch.long)
|
||
+ int(dst_logical_start)
|
||
) % self._compress_tail_k[idx].shape[1]
|
||
self._compress_tail_k[idx][req_pool_idx, slots] = key_tail
|
||
self._compress_tail_score[idx][req_pool_idx, slots] = score_tail
|
||
|
||
def _clear_buffers(self):
|
||
super()._clear_buffers()
|
||
self.index_key_cache.clear()
|
||
if hasattr(self, "_compress_tail_k") and self._compress_tail_k is not None:
|
||
del self._compress_tail_k
|
||
del self._compress_tail_score
|
||
|
||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
"""Move latent KV and the DSA indexer cache (key + scale) in lockstep."""
|
||
super().move_kv_cache(tgt_loc, src_loc)
|
||
self.index_key_cache.move(tgt_loc, src_loc)
|
||
|
||
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||
return self.index_key_cache.get_local_buffer(layer_id)
|
||
|
||
def get_index_k_continuous(
|
||
self,
|
||
layer_id: int,
|
||
seq_len: int,
|
||
page_indices: torch.Tensor,
|
||
):
|
||
return self.index_key_cache.get_k_continuous(layer_id, seq_len, page_indices)
|
||
|
||
def get_index_k_scale_continuous(
|
||
self,
|
||
layer_id: int,
|
||
seq_len: int,
|
||
page_indices: torch.Tensor,
|
||
):
|
||
return self.index_key_cache.get_k_scale_continuous(
|
||
layer_id, seq_len, page_indices
|
||
)
|
||
|
||
def get_index_k_scale_buffer(
|
||
self,
|
||
layer_id: int,
|
||
seq_len_tensor: torch.Tensor,
|
||
page_indices: torch.Tensor,
|
||
seq_len_sum: int,
|
||
max_seq_len: int,
|
||
):
|
||
return self.index_key_cache.get_k_and_scale(
|
||
layer_id, seq_len_tensor, page_indices, seq_len_sum, max_seq_len
|
||
)
|
||
|
||
def set_index_k_scale_buffer(
|
||
self,
|
||
layer_id: int,
|
||
loc: torch.Tensor,
|
||
index_k: torch.Tensor,
|
||
index_k_scale: torch.Tensor,
|
||
) -> None:
|
||
self.index_key_cache.store_quantized(layer_id, loc, index_k, index_k_scale)
|
||
|
||
def _get_compress_tail_cpu_copy(self, req_pool_index):
|
||
if not self.kpool_use_compress or req_pool_index is None:
|
||
return None
|
||
|
||
tail_k_cpu = []
|
||
tail_score_cpu = []
|
||
for tail_k, tail_score in zip(self._compress_tail_k, self._compress_tail_score):
|
||
if tail_k.shape[0] == 0:
|
||
tail_k_cpu.append(None)
|
||
tail_score_cpu.append(None)
|
||
continue
|
||
tail_k_cpu.append(tail_k[req_pool_index].to("cpu", non_blocking=True))
|
||
tail_score_cpu.append(
|
||
tail_score[req_pool_index].to("cpu", non_blocking=True)
|
||
)
|
||
return tail_k_cpu, tail_score_cpu
|
||
|
||
def _load_compress_tail_cpu_copy(self, tail_k_cpu, tail_score_cpu, req_pool_index):
|
||
if (
|
||
not self.kpool_use_compress
|
||
or req_pool_index is None
|
||
or tail_k_cpu is None
|
||
or tail_score_cpu is None
|
||
):
|
||
return
|
||
|
||
for tail_k, tail_score, saved_k, saved_score in zip(
|
||
self._compress_tail_k,
|
||
self._compress_tail_score,
|
||
tail_k_cpu,
|
||
tail_score_cpu,
|
||
):
|
||
if tail_k.shape[0] == 0 or saved_k is None or saved_score is None:
|
||
continue
|
||
tail_k[req_pool_index] = saved_k.to(tail_k.device, non_blocking=True)
|
||
tail_score[req_pool_index] = saved_score.to(
|
||
tail_score.device, non_blocking=True
|
||
)
|
||
|
||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||
# Retraction reuses index-cache pages; offload index/scale with KV so resume cannot read another request's entries.
|
||
kv_cache_cpu = super().get_cpu_copy(indices, mamba_indices=mamba_indices)
|
||
cpu_copy = {
|
||
"kv": kv_cache_cpu,
|
||
"index_k": self.index_key_cache.cpu_copy(indices),
|
||
}
|
||
compress_tail = self._get_compress_tail_cpu_copy(req_pool_index)
|
||
if compress_tail is not None:
|
||
cpu_copy["tail_k"], cpu_copy["tail_score"] = compress_tail
|
||
torch.cuda.synchronize()
|
||
return cpu_copy
|
||
|
||
def load_cpu_copy(
|
||
self,
|
||
kv_cache_cpu_dict,
|
||
indices,
|
||
mamba_indices=None,
|
||
req_pool_index=None,
|
||
):
|
||
super().load_cpu_copy(
|
||
kv_cache_cpu_dict["kv"],
|
||
indices,
|
||
mamba_indices=mamba_indices,
|
||
req_pool_index=req_pool_index,
|
||
)
|
||
self.index_key_cache.load_cpu_copy(kv_cache_cpu_dict["index_k"], indices)
|
||
self._load_compress_tail_cpu_copy(
|
||
kv_cache_cpu_dict.get("tail_k"),
|
||
kv_cache_cpu_dict.get("tail_score"),
|
||
req_pool_index,
|
||
)
|
||
torch.cuda.synchronize()
|
||
|
||
def get_state_buf_infos(self):
|
||
return self.index_key_cache.state_buf_infos()
|
||
|
||
def get_kv_size_bytes(self):
|
||
kv_size_bytes = super().get_kv_size_bytes()
|
||
for index_k_cache in self.index_k_with_scale_buffer:
|
||
kv_size_bytes += get_tensor_size_bytes(index_k_cache)
|
||
return kv_size_bytes
|
||
|
||
|
||
def move_kv_cache_native(
|
||
k_buffer: List[torch.Tensor],
|
||
v_buffer: List[torch.Tensor],
|
||
tgt_loc: torch.Tensor,
|
||
src_loc: torch.Tensor,
|
||
):
|
||
"""Move token-granular K/V rows from ``src_loc`` to ``tgt_loc``.
|
||
|
||
Buffers are the per-layer 3-D ``[max_slots, head_num, head_dim]`` pools;
|
||
direct advanced indexing on dim 0.
|
||
"""
|
||
if tgt_loc.numel() == 0:
|
||
return
|
||
|
||
tgt_loc_flat = tgt_loc.view(-1).long()
|
||
src_loc_flat = src_loc.view(-1).long()
|
||
for k_cache, v_cache in zip(k_buffer, v_buffer):
|
||
k_cache[tgt_loc_flat] = k_cache[src_loc_flat]
|
||
v_cache[tgt_loc_flat] = v_cache[src_loc_flat]
|
||
|
||
|
||
@triton.jit
|
||
def masked_set_kv_buffer_kernel(
|
||
k_ptr,
|
||
v_ptr,
|
||
k_buffer_ptr,
|
||
v_buffer_ptr,
|
||
loc_ptr,
|
||
mask_ptr,
|
||
N: tl.constexpr,
|
||
H: tl.constexpr,
|
||
D: tl.constexpr,
|
||
CHUNK: tl.constexpr,
|
||
k_stride_B: tl.constexpr,
|
||
k_stride_H: tl.constexpr,
|
||
v_stride_B: tl.constexpr,
|
||
v_stride_H: tl.constexpr,
|
||
):
|
||
pid = tl.program_id(0)
|
||
if pid >= N:
|
||
return
|
||
|
||
do_write = tl.load(mask_ptr + pid) != 0
|
||
if not do_write:
|
||
return
|
||
|
||
loc = tl.load(loc_ptr + pid)
|
||
total = H * D
|
||
num_chunks = tl.cdiv(total, CHUNK)
|
||
|
||
for c in range(num_chunks):
|
||
offs = tl.arange(0, CHUNK)
|
||
idx = c * CHUNK + offs
|
||
mask = idx < total
|
||
row = idx // D
|
||
col = idx % D
|
||
|
||
key = tl.load(k_ptr + pid * k_stride_B + row * k_stride_H + col, mask=mask)
|
||
tl.store(k_buffer_ptr + loc * H * D + idx, key, mask=mask)
|
||
|
||
value = tl.load(v_ptr + pid * v_stride_B + row * v_stride_H + col, mask=mask)
|
||
tl.store(v_buffer_ptr + loc * H * D + idx, value, mask=mask)
|
||
|
||
|
||
class MHATokenToKOnlyPool(KVCache):
|
||
"""K-only pool for MiniMax sparse layers whose index branch never reads V
|
||
(``sparse_disable_index_value``); allocating V would waste memory."""
|
||
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
page_size: int,
|
||
dtype: torch.dtype,
|
||
head_num: int,
|
||
head_dim: int,
|
||
layer_num: int,
|
||
device: str,
|
||
enable_memory_saver: bool,
|
||
start_layer: Optional[int] = None,
|
||
end_layer: Optional[int] = None,
|
||
):
|
||
super().__init__(
|
||
size,
|
||
page_size,
|
||
dtype,
|
||
layer_num,
|
||
device,
|
||
enable_memory_saver,
|
||
start_layer,
|
||
end_layer,
|
||
)
|
||
self.head_num = head_num
|
||
self.head_dim = head_dim
|
||
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()
|
||
):
|
||
self.k_buffer = [
|
||
torch.zeros(
|
||
(size + page_size, head_num, head_dim),
|
||
dtype=self.store_dtype,
|
||
device=device,
|
||
)
|
||
for _ in range(layer_num)
|
||
]
|
||
self._finalize_allocation_log(size)
|
||
|
||
def _get_key_buffer(self, layer_id: int):
|
||
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]
|
||
|
||
def register_layer_transfer_counter(
|
||
self, layer_transfer_counter: LayerDoneCounter
|
||
) -> None:
|
||
self.layer_transfer_counter = layer_transfer_counter
|
||
|
||
def get_key_buffer(self, layer_id: int):
|
||
if self.layer_transfer_counter is not None:
|
||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||
return self._get_key_buffer(layer_id)
|
||
|
||
def set_k_buffer(
|
||
self,
|
||
layer_id: int,
|
||
loc: torch.Tensor,
|
||
cache_k: torch.Tensor,
|
||
) -> None:
|
||
if cache_k.dtype != self.dtype:
|
||
cache_k = cache_k.to(self.dtype)
|
||
if self.store_dtype != self.dtype:
|
||
cache_k = cache_k.view(self.store_dtype)
|
||
self.k_buffer[layer_id][loc] = cache_k
|
||
|
||
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
|
||
raise NotImplementedError("MHATokenToKOnlyPool does not allocate V")
|
||
|
||
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
raise NotImplementedError("MHATokenToKOnlyPool does not allocate V")
|
||
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
k_scale: Optional[float] = None,
|
||
v_scale: Optional[float] = None,
|
||
layer_id_override: Optional[int] = None,
|
||
) -> None:
|
||
# Routed through MiniMaxSparseKVPool.set_index_k_buffer instead.
|
||
raise NotImplementedError(
|
||
"MHATokenToKOnlyPool: use set_index_k_buffer on the parent "
|
||
"MiniMaxSparseKVPool — this pool does not store V"
|
||
)
|
||
|
||
def get_kv_size_bytes(self):
|
||
k_size_bytes = sum(get_tensor_size_bytes(k) for k in self.k_buffer)
|
||
return k_size_bytes, 0
|
||
|
||
|
||
class MiniMaxSparseKVPool(KVCache):
|
||
def __init__(
|
||
self,
|
||
size: int,
|
||
page_size: int,
|
||
dtype: torch.dtype,
|
||
head_num: int,
|
||
head_dim: int,
|
||
idx_head_dim: int,
|
||
dense_layer_ids: List[int],
|
||
sparse_layer_ids: List[int],
|
||
device: str,
|
||
disable_value_sparse_layer_ids: Optional[List[int]] = None,
|
||
enable_memory_saver: bool = False,
|
||
index_dtype: Optional[torch.dtype] = None,
|
||
start_layer: Optional[int] = None,
|
||
end_layer: Optional[int] = None,
|
||
main_pool_cls=MHATokenToKVPool,
|
||
index_kv_pool_cls=MHATokenToKVPool,
|
||
index_k_pool_cls=MHATokenToKOnlyPool,
|
||
):
|
||
# Do not call super().__init__() — delegate to sub-pools instead.
|
||
self.size = size
|
||
self.page_size = page_size
|
||
self.dtype = dtype
|
||
self.device = device
|
||
self.use_minimax_fused_kv_index_store = (
|
||
envs.SGLANG_OPT_USE_MINIMAX_FUSED_KV_INDEX_STORE.get()
|
||
)
|
||
|
||
local_dense_layer_ids = [
|
||
lid for lid in dense_layer_ids if start_layer <= lid < end_layer
|
||
]
|
||
local_sparse_layer_ids = [
|
||
lid for lid in sparse_layer_ids if start_layer <= lid < end_layer
|
||
]
|
||
|
||
index_dtype = index_dtype if index_dtype is not None else dtype
|
||
|
||
# Split sparse layers by V policy: kv_sparse (index_kv_pool holds K+V) vs
|
||
# k_only_sparse (index_k_pool holds only K; V is never read).
|
||
disable_set = set(disable_value_sparse_layer_ids or [])
|
||
local_kv_sparse_layer_ids = [
|
||
g for g in local_sparse_layer_ids if g not in disable_set
|
||
]
|
||
local_k_only_sparse_layer_ids = [
|
||
g for g in local_sparse_layer_ids if g in disable_set
|
||
]
|
||
|
||
# Membership check across all sparse layers, regardless of split.
|
||
self.sparse_layer_id_mapping: dict[int, int] = {
|
||
gid: i for i, gid in enumerate(local_sparse_layer_ids)
|
||
}
|
||
# Per-sub-pool local indices.
|
||
self.index_kv_layer_id_mapping: dict[int, int] = {
|
||
gid: i for i, gid in enumerate(local_kv_sparse_layer_ids)
|
||
}
|
||
self.index_k_layer_id_mapping: dict[int, int] = {
|
||
gid: i for i, gid in enumerate(local_k_only_sparse_layer_ids)
|
||
}
|
||
|
||
self.main_pool = main_pool_cls(
|
||
size=size,
|
||
page_size=page_size,
|
||
dtype=dtype,
|
||
head_num=head_num,
|
||
head_dim=head_dim,
|
||
layer_num=len(local_dense_layer_ids) + len(local_sparse_layer_ids),
|
||
device=device,
|
||
enable_memory_saver=enable_memory_saver,
|
||
start_layer=start_layer,
|
||
end_layer=end_layer,
|
||
)
|
||
|
||
self.index_kv_pool: Optional[MHATokenToKVPool] = (
|
||
index_kv_pool_cls(
|
||
size=size,
|
||
page_size=page_size,
|
||
dtype=index_dtype,
|
||
head_num=1,
|
||
head_dim=idx_head_dim,
|
||
layer_num=len(local_kv_sparse_layer_ids),
|
||
device=device,
|
||
enable_memory_saver=enable_memory_saver,
|
||
)
|
||
if local_kv_sparse_layer_ids
|
||
else None
|
||
)
|
||
|
||
self.index_k_pool: Optional[MHATokenToKOnlyPool] = (
|
||
index_k_pool_cls(
|
||
size=size,
|
||
page_size=page_size,
|
||
dtype=index_dtype,
|
||
head_num=1,
|
||
head_dim=idx_head_dim,
|
||
layer_num=len(local_k_only_sparse_layer_ids),
|
||
device=device,
|
||
enable_memory_saver=enable_memory_saver,
|
||
)
|
||
if local_k_only_sparse_layer_ids
|
||
else None
|
||
)
|
||
|
||
self.mem_usage = self.main_pool.mem_usage
|
||
if self.index_kv_pool is not None:
|
||
self.mem_usage += self.index_kv_pool.mem_usage
|
||
if self.index_k_pool is not None:
|
||
self.mem_usage += self.index_k_pool.mem_usage
|
||
|
||
# HiCacheController reads these from the top-level KV pool wrapper.
|
||
self.layer_num = self.main_pool.layer_num
|
||
self.start_layer = self.main_pool.start_layer
|
||
self.end_layer = self.main_pool.end_layer
|
||
# PD disaggregation reads these directly (no fallback) off the wrapper.
|
||
self.head_num = self.main_pool.head_num
|
||
self.head_dim = self.main_pool.head_dim
|
||
self.layer_transfer_counter = None
|
||
|
||
def register_layer_transfer_counter(
|
||
self, layer_transfer_counter: LayerDoneCounter
|
||
) -> None:
|
||
self.layer_transfer_counter = layer_transfer_counter
|
||
|
||
def get_kv_cache_quant_method(self) -> Any:
|
||
# The base unwrap chain only knows full_kv_pool/swa_kv_pool; the dense
|
||
# KV (what attention backends quantize against) lives in main_pool here.
|
||
return self.main_pool.get_kv_cache_quant_method()
|
||
|
||
def _wait_for_layer(self, layer_id: int) -> None:
|
||
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) -> torch.Tensor:
|
||
self._wait_for_layer(layer_id)
|
||
return self.main_pool.get_key_buffer(layer_id)
|
||
|
||
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
|
||
self._wait_for_layer(layer_id)
|
||
return self.main_pool.get_value_buffer(layer_id)
|
||
|
||
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
self._wait_for_layer(layer_id)
|
||
return self.main_pool.get_kv_buffer(layer_id)
|
||
|
||
def get_index_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||
self._wait_for_layer(layer_id)
|
||
mapped_id = self.index_kv_layer_id_mapping.get(layer_id)
|
||
if mapped_id is None:
|
||
raise ValueError(
|
||
f"layer_id={layer_id} does not have an index V cache "
|
||
f"(either dense, or in the K-only group). "
|
||
f"index_kv layers: {list(self.index_kv_layer_id_mapping.keys())}"
|
||
)
|
||
return self.index_kv_pool.get_kv_buffer(mapped_id)
|
||
|
||
def get_index_k_buffer(self, layer_id: int) -> torch.Tensor:
|
||
self._wait_for_layer(layer_id)
|
||
# First try the K-only pool; fall back to the index_kv pool's K side
|
||
# so callers that just need K work for both sparse subgroups.
|
||
mapped_id = self.index_k_layer_id_mapping.get(layer_id)
|
||
if mapped_id is not None:
|
||
return self.index_k_pool.get_key_buffer(mapped_id)
|
||
mapped_id = self.index_kv_layer_id_mapping.get(layer_id)
|
||
if mapped_id is not None:
|
||
return self.index_kv_pool.get_key_buffer(mapped_id)
|
||
raise ValueError(
|
||
f"layer_id={layer_id} is not a sparse attention layer; "
|
||
f"sparse layers: {list(self.sparse_layer_id_mapping.keys())}"
|
||
)
|
||
|
||
def set_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
k_scale: Optional[float] = None,
|
||
v_scale: Optional[float] = None,
|
||
) -> None:
|
||
"""Write main K/V at `loc`. Works for any layer (dense or sparse).
|
||
|
||
Scale semantics follow MHATokenToKVPool: None means unit scale;
|
||
a non-None scale is applied with an in-place div_ before the fp8 cast.
|
||
"""
|
||
self.main_pool.set_kv_buffer(
|
||
layer,
|
||
loc,
|
||
cache_k,
|
||
cache_v,
|
||
k_scale,
|
||
v_scale,
|
||
)
|
||
|
||
def set_index_kv_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_idx_k: torch.Tensor,
|
||
cache_idx_v: torch.Tensor,
|
||
k_scale: Optional[float] = None,
|
||
v_scale: Optional[float] = None,
|
||
) -> None:
|
||
mapped_id = self.index_kv_layer_id_mapping.get(layer.layer_id)
|
||
if mapped_id is None:
|
||
raise ValueError(
|
||
f"layer.layer_id={layer.layer_id} does not have an index V "
|
||
f"cache (either dense, or in the K-only group). "
|
||
f"index_kv layers: {list(self.index_kv_layer_id_mapping.keys())}"
|
||
)
|
||
self.index_kv_pool.set_kv_buffer(
|
||
layer,
|
||
loc,
|
||
cache_idx_k,
|
||
cache_idx_v,
|
||
k_scale,
|
||
v_scale,
|
||
layer_id_override=mapped_id,
|
||
)
|
||
|
||
def set_index_k_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_idx_k: torch.Tensor,
|
||
k_scale: Optional[float] = None,
|
||
) -> None:
|
||
mapped_id = self.index_k_layer_id_mapping.get(layer.layer_id)
|
||
if mapped_id is None:
|
||
raise ValueError(
|
||
f"layer.layer_id={layer.layer_id} is not in the K-only "
|
||
f"sparse group. K-only layers: "
|
||
f"{list(self.index_k_layer_id_mapping.keys())}"
|
||
)
|
||
sub_pool = self.index_k_pool
|
||
if cache_idx_k.dtype != sub_pool.dtype:
|
||
if k_scale is not None:
|
||
cache_idx_k = cache_idx_k / k_scale
|
||
sub_pool.set_k_buffer(mapped_id, loc, cache_idx_k)
|
||
|
||
def _can_fuse_kv_index_store(
|
||
self,
|
||
index_pool: MHATokenToKVPool,
|
||
cache_k: torch.Tensor,
|
||
cache_idx_k: torch.Tensor,
|
||
) -> bool:
|
||
"""Fast-path precondition: CUDA, no per-store quantization, and a uniform
|
||
head byte size shared by main and index caches."""
|
||
main = self.main_pool
|
||
return (
|
||
self.use_minimax_fused_kv_index_store
|
||
and _is_cuda
|
||
# No dtype conversion / fp8 scaling on either side (the fused kernel
|
||
# is a raw byte copy, it does not quantize).
|
||
and main.store_dtype == main.dtype
|
||
and index_pool.store_dtype == index_pool.dtype
|
||
and cache_k.dtype == main.dtype
|
||
and cache_idx_k.dtype == index_pool.dtype
|
||
# Uniform head byte size collapses head_dim + dtype into one constant.
|
||
and main.dtype == index_pool.dtype
|
||
and main.head_dim == index_pool.head_dim
|
||
# 128-bit vector copy requires a 16-byte-aligned head size.
|
||
and (main.head_dim * main.dtype.itemsize) % 16 == 0
|
||
)
|
||
|
||
def set_fused_kv_index_buffer(
|
||
self,
|
||
layer: RadixAttention,
|
||
loc: torch.Tensor,
|
||
cache_k: torch.Tensor,
|
||
cache_v: torch.Tensor,
|
||
cache_idx_k: torch.Tensor,
|
||
cache_idx_v: Optional[torch.Tensor],
|
||
k_scale: Optional[float] = None,
|
||
v_scale: Optional[float] = None,
|
||
idx_k_scale: Optional[float] = None,
|
||
idx_v_scale: Optional[float] = None,
|
||
) -> None:
|
||
"""Store main K/V + index K (+ optional index V) for a sparse layer in
|
||
one fused JIT launch, falling back to separate stores when not applicable."""
|
||
disable_value = cache_idx_v is None
|
||
index_pool = self.index_k_pool if disable_value else self.index_kv_pool
|
||
|
||
if index_pool is not None and self._can_fuse_kv_index_store(
|
||
index_pool, cache_k, cache_idx_k
|
||
):
|
||
from sglang.kernels.ops.kvcache.minimax_store_kv_index import store_kv_index
|
||
|
||
main = self.main_pool
|
||
head_bytes = main.head_dim * main.dtype.itemsize
|
||
if disable_value:
|
||
idx_k_cache = self.get_index_k_buffer(layer.layer_id).flatten(1)
|
||
idx_v_cache = None
|
||
else:
|
||
ik, iv = self.get_index_kv_buffer(layer.layer_id)
|
||
idx_k_cache, idx_v_cache = ik.flatten(1), iv.flatten(1)
|
||
store_kv_index(
|
||
cache_k.flatten(1),
|
||
cache_v.flatten(1),
|
||
main.get_key_buffer(layer.layer_id).flatten(1),
|
||
main.get_value_buffer(layer.layer_id).flatten(1),
|
||
cache_idx_k.flatten(1),
|
||
idx_k_cache,
|
||
None if disable_value else cache_idx_v.flatten(1),
|
||
idx_v_cache,
|
||
loc,
|
||
num_kv_heads=main.head_num,
|
||
head_bytes=head_bytes,
|
||
)
|
||
return
|
||
|
||
# Fallback: separate stores (identical semantics; quantizes for fp8
|
||
# pools — the fused raw-byte path is disqualified there by
|
||
# _can_fuse_kv_index_store's dtype-equality checks). Scales use the
|
||
# None-means-unit convention throughout: MHATokenToKVPool.set_kv_buffer
|
||
# applies any non-None scale with an IN-PLACE div_ (extra kernel +
|
||
# caller-tensor mutation), which must not fire for unit scale.
|
||
self.set_kv_buffer(layer, loc, cache_k, cache_v, k_scale, v_scale)
|
||
if disable_value:
|
||
self.set_index_k_buffer(layer, loc, cache_idx_k, idx_k_scale)
|
||
else:
|
||
self.set_index_kv_buffer(
|
||
layer,
|
||
loc,
|
||
cache_idx_k,
|
||
cache_idx_v,
|
||
idx_k_scale,
|
||
idx_v_scale,
|
||
)
|
||
|
||
def get_kv_size_bytes(self):
|
||
sub_pools = [self.main_pool, self.index_kv_pool, self.index_k_pool]
|
||
sizes = [p.get_kv_size_bytes() for p in sub_pools if p is not None]
|
||
return sum(k for k, _ in sizes), sum(v for _, v in sizes)
|
||
|
||
def get_contiguous_buf_infos(self):
|
||
# Main K/V only; index buffers ride the state-buffer channel.
|
||
return self.main_pool.get_contiguous_buf_infos()
|
||
|
||
def get_index_k_state_buf_infos(self):
|
||
# Per-page item_len (MHATokenToKVPool convention); index rows share the
|
||
# main-KV `loc`, so the transfer reuses the same page-ids.
|
||
pool = self.index_k_pool
|
||
n = pool.layer_num
|
||
data_ptrs = [pool.k_buffer[i].data_ptr() for i in range(n)]
|
||
data_lens = [pool.k_buffer[i].nbytes for i in range(n)]
|
||
item_lens = [pool.k_buffer[i][0].nbytes * pool.page_size for i in range(n)]
|
||
return data_ptrs, data_lens, item_lens
|
||
|
||
def maybe_get_custom_mem_pool(self):
|
||
return self.main_pool.maybe_get_custom_mem_pool()
|
||
|
||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||
# TODO: spec-decode needs sub-pools built with enable_kv_cache_copy=True,
|
||
# then delegate to main_pool/index_pool.move_kv_cache.
|
||
raise NotImplementedError(
|
||
"move_kv_cache is not yet supported for MiniMaxSparseKVPool: "
|
||
"sub-pools must be built with enable_kv_cache_copy=True first."
|
||
)
|
||
|
||
def get_v_head_dim(self):
|
||
# Use start_layer to handle pipeline parallelism where layer 0
|
||
# may not be present in this stage's buffer.
|
||
return self.main_pool.get_value_buffer(self.main_pool.start_layer).shape[-1]
|