[unified-memory] Support MLA-hybrid-Mamba (Kimi-Linear) on the Triton backend (#32971)

This commit is contained in:
Cheng Wan
2026-07-30 22:10:34 -07:00
committed by GitHub
parent 06ccaef24a
commit e23ccb15f0
10 changed files with 1163 additions and 48 deletions
@@ -60,6 +60,7 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
h,
initial_state,
initial_state_indices,
stride_init_state,
cu_seqlens,
chunk_offsets,
T,
@@ -113,9 +114,13 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
stride_k = Hg * K
stride_w = H * K
index = tl.load(initial_state_indices + i_n).to(tl.int32)
h0 = initial_state + index * stride_h
ht = initial_state + index * stride_h
# Slot stride comes from the caller (initial_state.stride(0)): the state pool
# may be an envelope-strided view (page-major / unified memory), where the
# per-slot pitch spans ALL layers' state, not H*V*K. int64: envelope pitches
# overflow an int32 index product.
index = tl.load(initial_state_indices + i_n).to(tl.int64)
h0 = initial_state + index * stride_init_state
ht = initial_state + index * stride_init_state
if USE_INITIAL_STATE:
h0 = h0 + i_h * V * K
if INPLACE_UPDATE:
@@ -355,6 +360,9 @@ def chunk_gated_delta_rule_fwd_h(
h=h,
initial_state=initial_state,
initial_state_indices=initial_state_indices,
# Envelope-strided state pools (page-major / unified memory) have a
# per-slot pitch != H*V*K; contiguous pools pass exactly H*V*K.
stride_init_state=(initial_state.stride(0) if initial_state is not None else 0),
cu_seqlens=cu_seqlens,
chunk_offsets=chunk_offsets,
T=T,
@@ -160,10 +160,13 @@ class TritonAttnBackend(AttentionBackend):
# byte-identical to the slot-based envelope.
self.page_size = getattr(model_runner, "page_size", 1) or 1
# Unified pool v2p hook (None = no-op): req_to_token holds VIRTUAL ids but
# kernels need PHYSICAL. Applied eagerly so the captured graph has no translate.
# kernels need the kernel-facing id space — PHYSICAL for MHA, DENSE for the
# dense-view MLA pool (translate_kv_loc_dense falls back to the physical
# translate when kernel_page_multiplier == 1, so preferring it is exact for
# both). Applied eagerly so the captured graph has no translate.
self._translate_kv_loc = getattr(
self.token_to_kv_pool_allocator, "translate_kv_loc", None
)
self.token_to_kv_pool_allocator, "translate_kv_loc_dense", None
) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None)
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.topk = model_runner.server_args.speculative_eagle_topk or 0
@@ -1243,6 +1246,9 @@ class TritonAttnBackend(AttentionBackend):
cache_loc = forward_batch.out_cache_loc
if isinstance(pool, SWAKVPool) and pool.layers_mapping[layer.layer_id][1]:
cache_loc = pool.translate_loc_from_full_to_swa(cache_loc)
elif self._translate_kv_loc is not None:
# Unified pool: buffers are indexed in the kernel-facing id space.
cache_loc = self._translate_kv_loc(cache_loc)
k_buffer, v_buffer = pool.get_kv_buffer(layer.layer_id)
k = k_buffer[cache_loc]
v = v_buffer[cache_loc]
@@ -1710,7 +1716,15 @@ class TritonAttnBackend(AttentionBackend):
k.div_(layer.k_scale)
self.token_to_kv_pool.set_kv_buffer(
layer,
forward_batch.out_cache_loc,
# `full_loc` carries the pre-translated loc under the unified
# pool, refreshed into a capture-stable buffer before replay —
# translating inside set_kv_buffer would be captured and replay
# a stale v2p. None (-> raw loc) for static pools.
KVWriteLoc(
forward_batch.out_cache_loc,
self.forward_metadata.swa_out_cache_loc,
full_loc=self.forward_metadata.out_cache_loc_full_physical,
),
k,
v,
)
@@ -430,9 +430,6 @@ class KVCacheConfigurator:
config = self.mambaish_config
assert config is not None
assert (
not self.use_mla_backend
), "unified memory pool does not support MLA-hybrid-Mamba yet"
# The full sub-pool is page-aware (via `MultiEndedAllocator(page_size=...)`);
# the mamba sub-pool stays page=1.
assert self.page_size >= 1, f"page_size must be >= 1, got {self.page_size}"
@@ -462,6 +459,12 @@ class KVCacheConfigurator:
end_layer=self.layer_info.end_layer,
is_draft_worker=self.is_draft_worker,
use_mla_backend=self.use_mla_backend,
kv_lora_rank=(
self.model_config.kv_lora_rank if self.use_mla_backend else None
),
qk_rope_head_dim=(
self.model_config.qk_rope_head_dim if self.use_mla_backend else None
),
mamba_layer_ids=mamba_layer_ids,
full_attention_layer_ids=full_attention_layer_ids,
mamba2_cache_params=config.mamba2_cache_params,
@@ -108,6 +108,74 @@ def build_page_major_mha_views(
return k_buffer, v_buffer
def mla_entry_bytes(*, layer_num: int, kv_cache_dim: int, itemsize: int) -> int:
"""Bytes occupied by one MLA slot across all layers (single latent row, no V)."""
return layer_num * kv_cache_dim * itemsize
def build_dense_mla_views(
raw: torch.Tensor,
*,
layer_num: int,
kv_cache_dim: int,
store_dtype: torch.dtype,
page_size: int,
num_pages: int,
anchor_bytes: int = 0,
) -> List[torch.Tensor]:
"""Per-layer DENSE views over ``raw`` for MLA in the page-major layout.
The page envelope is ``[L0_latent * ps | L1_latent * ps | ...]``. Because all
MLA layers share one uniform row size (``kv_cache_dim``), the envelope is
itself a valid dense paged pool under a re-numbered index space: folding the
layer offset ``l * ps * kv_cache_dim`` into each view's storage_offset makes
every per-layer view a plain CONTIGUOUS ``(num_pages * layer_num * ps, 1,
kv_cache_dim)`` tensor, addressed by the layer-independent dense id
dense(t) = (t // ps) * (ps * layer_num) + t % ps (t = physical token)
so one shared block table (entry = page * layer_num) serves every layer, and
kernels that require ``.view(-1, page_size, kv_cache_dim)`` (trtllm/cutlass/
flashmla) work on the views natively.
The views overlap each other (view ``l+1`` is view ``l`` shifted by ``ps``
rows); that is safe because layer ``l`` is only ever indexed at dense ids,
which always resolve to layer-``l`` bytes relative to view ``l``'s origin.
Layer ``layer_num-1``'s view extends ``(layer_num-1) * ps`` rows past the
last page envelope, so ``raw`` must carry at least one extra page envelope
of tail padding (``UnifiedKVPool``'s ``view_tail_pad_bytes``).
"""
itemsize = store_dtype.itemsize
row_bytes = kv_cache_dim * itemsize
page_bytes = page_size * layer_num * row_bytes
n_dense = num_pages * layer_num * page_size
assert anchor_bytes % itemsize == 0
last_view_end = (
anchor_bytes + (layer_num - 1) * page_size * row_bytes + (n_dense * row_bytes)
)
assert last_view_end <= raw.numel() * raw.itemsize, (
f"build_dense_mla_views: layer {layer_num - 1}'s view ends at byte "
f"{last_view_end} but the raw buffer holds only "
f"{raw.numel() * raw.itemsize} bytes; allocate the tail pad "
f"(one page envelope = {page_bytes} B) via view_tail_pad_bytes"
)
as_dtype_view = raw.view(store_dtype)
views: List[torch.Tensor] = []
for layer in range(layer_num):
base_bytes = anchor_bytes + layer * page_size * row_bytes
assert base_bytes % itemsize == 0
views.append(
torch.as_strided(
as_dtype_view,
size=(n_dense, 1, kv_cache_dim),
stride=(kv_cache_dim, kv_cache_dim, 1),
storage_offset=base_bytes // itemsize,
)
)
return views
def mamba_entry_bytes(
*,
layer_num: int,
+13 -1
View File
@@ -3551,6 +3551,10 @@ class HybridLinearKVPool(KVCache):
# virtual->physical mamba-slot translate for the HiCache offload path;
# identity for a static pool, the allocator's `translate` for the unified pool.
self._mamba_translate = lambda ids: ids
# virtual->dense full-KV translate for the model-level MLA entry points
# (`set_mla_kv_buffer` / `get_mla_kv_buffer` receive VIRTUAL locs);
# identity for a static pool, `translate_kv_loc_dense` for the unified pool.
self._full_translate = lambda ids: ids
self.use_mla = use_mla
if full_kv_pool is not None:
# Shared-KV-pool path: the caller built a UnifiedMHATokenToKVPool
@@ -3791,10 +3795,13 @@ class HybridLinearKVPool(KVCache):
dcp_kv_mask=dcp_kv_mask,
)
else:
# Mirror the MHA branch: `full_loc` is the unified pool's
# pre-translated (dense) 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,
loc,
write_loc,
cache_k,
cache_v,
)
@@ -3831,6 +3838,10 @@ class HybridLinearKVPool(KVCache):
cache_k_rope: torch.Tensor,
):
assert self.use_mla, "set_mla_kv_buffer called when use_mla is False"
# Model-level MLA entry point: `loc` is a VIRTUAL loc under the unified
# pool (eager prefill only; the decode write goes through set_kv_buffer's
# pre-translated `full_loc`), so translate to the dense id space here.
loc = self._full_translate(loc)
with self._transfer_id_context(layer):
self.full_kv_pool.set_mla_kv_buffer(layer, loc, cache_k_nope, cache_k_rope)
@@ -3841,6 +3852,7 @@ class HybridLinearKVPool(KVCache):
dst_dtype: Optional[torch.dtype] = None,
):
assert self.use_mla, "get_mla_kv_buffer called when use_mla is False"
loc = self._full_translate(loc)
with self._transfer_id_context(layer):
return self.full_kv_pool.get_mla_kv_buffer(layer, loc, dst_dtype)
@@ -111,6 +111,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
need_sort: bool = False,
forward_stream: Optional[torch.cuda.Stream] = None,
lazy_compaction: bool = False,
kernel_page_multiplier: int = 1,
):
spec = unified_buffer.spec(sub_pool_name)
max_slots = unified_buffer.max_slots(sub_pool_name)
@@ -130,6 +131,11 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
self.entry_bytes = spec.entry_bytes()
self.min_slot_index = unified_buffer.min_slot_index(sub_pool_name)
self.is_id_owner = is_id_owner
# Dense (kernel-facing) index space scale: the page-major envelope of a
# multi-layer uniform-entry sub-pool (MLA) is a valid dense paged pool
# once page ids are scaled by layer_num — `translate_kv_loc_dense` emits
# that space. 1 for sub-pools whose kernels take real physical ids.
self.kernel_page_multiplier = kernel_page_multiplier
# Overlap mode: `free` drops a wait_stream(forward_stream) barrier so its
# v2p writes + move kernel serialize after the in-flight forward.
self.forward_stream = forward_stream
@@ -681,6 +687,59 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
result = phys_pages * self.page_size + offsets
return torch.clamp_min(result, 0)
def translate_kv_loc_dense(
self,
virt_tokens: torch.Tensor,
*,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Translate virtual token ids to DENSE (kernel-facing) ids.
dense(t) = (t // ps) * (ps * kernel_page_multiplier) + t % ps for the
physical token t — i.e. `translate_kv_loc` with the page stride scaled by
`kernel_page_multiplier` (= layer_num for a dense-view MLA sub-pool; see
`build_dense_mla_views`). Internal machinery (compaction, in-flight write
sets) MUST keep using `translate_kv_loc`: dense ids are for kernels only.
The tombstone clamp routes -1 entries to dense id 0 — inside the page-0
reserved sink for every layer view. Supports ``out=`` like
`translate_kv_loc` for cuda-graph buffer stability.
"""
if self.kernel_page_multiplier == 1:
return self.translate_kv_loc(virt_tokens, out=out)
if out is not None:
assert out.dtype == torch.int64, (
f"translate_kv_loc_dense: out= dtype must be int64 (matches v2p), "
f"got {out.dtype}"
)
assert out.shape == virt_tokens.shape, (
f"translate_kv_loc_dense: out= shape {tuple(out.shape)} must "
f"match virt_tokens shape {tuple(virt_tokens.shape)}"
)
with record_function("MultiEndedAlloc.translate_kv_loc_dense"):
dense_page_stride = self.page_size * self.kernel_page_multiplier
if self.page_size == 1:
# dense = phys * multiplier; tombstone -1 scales negative → clamp 0.
if out is not None:
tmp = torch.index_select(self.virtual_to_physical, 0, virt_tokens)
tmp = torch.clamp_min(tmp * dense_page_stride, 0)
out.copy_(tmp)
return out
result = torch.index_select(self.virtual_to_physical, 0, virt_tokens)
return torch.clamp_min(result * dense_page_stride, 0)
virt_pages = virt_tokens // self.page_size
offsets = virt_tokens % self.page_size
if out is not None:
torch.index_select(self.virtual_to_physical, 0, virt_pages, out=out)
out.mul_(dense_page_stride)
out.add_(offsets)
# tombstoned page: -1*dense_page_stride + offset < 0
out.clamp_(min=0)
return out
phys_pages = self.virtual_to_physical[virt_pages]
result = phys_pages * dense_page_stride + offsets
return torch.clamp_min(result, 0)
# -- alloc --
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
@@ -1654,12 +1713,13 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
need_sort: bool = False,
forward_stream: Optional[torch.cuda.Stream] = None,
lazy_compaction: bool = False,
full_kernel_page_multiplier: int = 1,
):
full_max = unified_buffer.max_slots("full")
super().__init__(
size=full_max - 1,
page_size=page_size,
dtype=unified_buffer.mha_spec("full").store_dtype,
dtype=unified_buffer.spec("full").get_dtype(),
device=device,
kvcache=kvcache,
need_sort=need_sort,
@@ -1681,6 +1741,7 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
need_sort=need_sort,
forward_stream=forward_stream,
lazy_compaction=lazy_compaction,
kernel_page_multiplier=full_kernel_page_multiplier,
)
self.mamba_allocator = MultiEndedAllocator(
kvcache=kvcache.mamba_pool,
@@ -1833,6 +1894,20 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
result = self.full_attn_allocator.translate_kv_loc(loc, out=out)
return result
@property
def kernel_page_multiplier(self) -> int:
return self.full_attn_allocator.kernel_page_multiplier
def translate_kv_loc_dense(
self,
loc: torch.Tensor,
*,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Full-pool virtual TOKEN ids -> DENSE (kernel-facing) ids. Falls back
to the physical translate when `kernel_page_multiplier == 1` (MHA)."""
return self.full_attn_allocator.translate_kv_loc_dense(loc, out=out)
def is_slot_allocated(self, slot: int) -> bool:
return self.full_attn_allocator.is_slot_allocated(slot)
@@ -35,6 +35,7 @@ from torch.profiler import record_function
from sglang.kernels.ops.kvcache.cache_move import store_cache_4d_kernel
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.mem_cache.layout.page_major import (
build_dense_mla_views,
build_page_major_mamba_views,
build_page_major_mha_views,
)
@@ -42,6 +43,7 @@ from sglang.srt.mem_cache.memory_pool import (
HybridReqToTokenPool,
MambaPool,
MHATokenToKVPool,
MLATokenToKVPool,
move_kv_cache_native,
unwrap_write_loc,
)
@@ -139,6 +141,40 @@ class MHASubPoolSpec(SubPoolSpec):
return self.store_dtype
@dataclass(frozen=True, kw_only=True)
class MLASubPoolSpec(SubPoolSpec):
"""Per-slot layout of one MLA-shaped sub-pool.
One latent row (``kv_lora_rank + qk_rope_head_dim``) per token per layer; V
is a prefix slice of the same row, so there is no separate V region. Not a
subclass of ``MHASubPoolSpec`` the K+V byte math and the ``v_head_dim > 0``
invariant there do not apply.
"""
kv_lora_rank: int
qk_rope_head_dim: int
store_dtype: torch.dtype
def __post_init__(self):
super().__post_init__()
assert (
self.kv_lora_rank > 0
), f"kv_lora_rank must be positive; got {self.kv_lora_rank}"
assert (
self.qk_rope_head_dim > 0
), f"qk_rope_head_dim must be positive; got {self.qk_rope_head_dim}"
@property
def kv_cache_dim(self) -> int:
return self.kv_lora_rank + self.qk_rope_head_dim
def entry_bytes(self) -> int:
return self.layer_num * self.kv_cache_dim * self.store_dtype.itemsize
def get_dtype(self) -> torch.dtype:
return self.store_dtype
@dataclass(frozen=True, kw_only=True)
class MambaSubPoolSpec(SubPoolSpec):
"""Per-slot layout of one Mamba-shaped sub-pool."""
@@ -188,6 +224,7 @@ class UnifiedKVPool:
device: str,
enable_memory_saver: bool,
page_size: int = 1,
view_tail_pad_bytes: int = 0,
):
assert page_size >= 1, f"page_size must be >= 1; got {page_size}"
assert len(sub_pool_specs) == 2, (
@@ -213,25 +250,44 @@ class UnifiedKVPool:
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
)
# `view_tail_pad_bytes` extends the ALLOCATION only (dense MLA views are
# per-layer shifted, so the last layer's view reaches past the final page
# envelope); all slot/watermark math stays on the unpadded `total_bytes`.
self.view_tail_pad_bytes = view_tail_pad_bytes
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
self._raw = torch.empty(total_bytes, dtype=torch.uint8, device=device)
self._raw = torch.empty(
total_bytes + view_tail_pad_bytes, dtype=torch.uint8, device=device
)
self._raw.zero_() # unset slots must read as zeros (matches non-shared)
self._max_slots: Dict[str, int] = {}
self._anchor_bytes: Dict[str, int] = {}
self._min_slot_index: Dict[str, int] = {}
# MHA: (k_buffer, v_buffer); Mamba: (conv_state_list, temporal_state)
# MHA: (k_buffer, v_buffer); MLA: [per-layer dense views];
# Mamba: (conv_state_list, temporal_state)
self._mha_views: Dict[str, Tuple[List[torch.Tensor], List[torch.Tensor]]] = {}
self._mla_views: Dict[str, List[torch.Tensor]] = {}
self._mamba_views: Dict[str, Tuple[List[torch.Tensor], torch.Tensor]] = {}
# Slot-0 dummy writes for both pools land in [0, entry_max); each pool's
# first allocatable slot is chosen so real data starts at >= entry_max.
# Slot-0 dummy writes for both pools land in the reserved low-byte sink;
# each pool's first allocatable slot is chosen so real data starts past it.
# For a page-aware sub-pool the slot-0 write touches layer blocks spread
# across the WHOLE page-0 envelope (up to page_size * entry_bytes), not
# just one slot envelope — reserve the max of both.
entry_max = max(s.entry_bytes() for s in sub_pool_specs)
reserved_floor = max(
[entry_max]
+ [
page_size * s.entry_bytes()
for s in sub_pool_specs
if not isinstance(s, MambaSubPoolSpec) # mamba is page_size=1
]
)
for spec in sub_pool_specs:
entry_bytes = spec.entry_bytes()
max_slots = total_bytes // entry_bytes
min_slot_index = (entry_max + entry_bytes - 1) // entry_bytes # ceil
min_slot_index = (reserved_floor + entry_bytes - 1) // entry_bytes # ceil
if max_slots <= min_slot_index:
raise RuntimeError(
f"UnifiedKVPool: sub-pool {spec.name!r} fits only {max_slots} "
@@ -249,6 +305,13 @@ class UnifiedKVPool:
max_slots,
page_size=page_size,
)
elif isinstance(spec, MLASubPoolSpec):
self._mla_views[spec.name] = self._build_mla_views(
spec,
anchor,
max_slots,
page_size=page_size,
)
elif isinstance(spec, MambaSubPoolSpec):
self._mamba_views[spec.name] = self._build_mamba_views(
spec, anchor, max_slots
@@ -289,6 +352,13 @@ class UnifiedKVPool:
), f"sub-pool {name!r} is {type(s).__name__}, expected MHASubPoolSpec"
return s
def mla_spec(self, name: str) -> MLASubPoolSpec:
s = self._specs_by_name[name]
assert isinstance(
s, MLASubPoolSpec
), f"sub-pool {name!r} is {type(s).__name__}, expected MLASubPoolSpec"
return s
def mamba_spec(self, name: str) -> MambaSubPoolSpec:
s = self._specs_by_name[name]
assert isinstance(
@@ -310,6 +380,9 @@ class UnifiedKVPool:
def mha_views_for(self, name: str) -> Tuple[List[torch.Tensor], List[torch.Tensor]]:
return self._mha_views[name]
def mla_views_for(self, name: str) -> List[torch.Tensor]:
return self._mla_views[name]
def mamba_views_for(self, name: str) -> Tuple[List[torch.Tensor], torch.Tensor]:
return self._mamba_views[name]
@@ -332,6 +405,23 @@ class UnifiedKVPool:
anchor_bytes=anchor_bytes,
)
def _build_mla_views(
self,
spec: MLASubPoolSpec,
anchor_bytes: int,
max_slots: int,
page_size: int,
) -> List[torch.Tensor]:
return build_dense_mla_views(
self._raw,
layer_num=spec.layer_num,
kv_cache_dim=spec.kv_cache_dim,
store_dtype=spec.store_dtype,
page_size=page_size,
num_pages=max_slots // page_size,
anchor_bytes=anchor_bytes,
)
def _build_mamba_views(
self, spec: MambaSubPoolSpec, anchor_bytes: int, max_slots: int
) -> Tuple[List[torch.Tensor], torch.Tensor]:
@@ -497,6 +587,91 @@ class UnifiedMHATokenToKVPool(MHATokenToKVPool):
)
class UnifiedMLATokenToKVPool(MLATokenToKVPool):
"""MLA KV pool whose per-layer `kv_buffer` entries are DENSE views into a
`UnifiedKVPool` (see `build_dense_mla_views`).
Loc-space contract: every loc this pool receives through the KVCache API
(`set_kv_buffer` / `set_mla_kv_buffer` / `get_mla_kv_buffer`, and the
kv_indices consumed by attention kernels reading `get_key_buffer` /
`get_value_buffer`) is a DENSE id the `translate_kv_loc_dense` output
dense(t) = (t // ps) * (ps * layer_num) + t % ps
which is layer-independent (the layer offset is folded into each view's
storage_offset), so the stock `MLATokenToKVPool` read/write methods work on
the views unmodified. The ONE exception is `move_kv_cache`: the allocator's
compaction calls it with REAL physical token ids, and it is overridden to
relocate whole page envelopes on the raw buffer.
"""
def __init__(
self,
*,
unified_buffer: UnifiedKVPool,
sub_pool_name: str,
kv_cache_dtype: torch.dtype,
page_size: int = 1,
):
spec = unified_buffer.mla_spec(sub_pool_name)
store_dtype = _store_dtype_for(kv_cache_dtype)
assert spec.store_dtype == store_dtype, (
f"sub-pool {sub_pool_name!r} store dtype {spec.store_dtype} does not "
f"match kv cache dtype {kv_cache_dtype} (store {store_dtype})"
)
self._unified_buffer = unified_buffer
self._sub_pool_name = sub_pool_name
self._kv_views = unified_buffer.mla_views_for(sub_pool_name)
max_slots = unified_buffer.max_slots(sub_pool_name)
self._num_pages = max_slots // page_size
self._page_bytes = page_size * spec.entry_bytes()
# Dense row count per view; also the OOB bound for dense locs.
self._dense_size = self._num_pages * spec.layer_num * page_size
super().__init__(
# OOB checks bound locs by `size + page_size`; dense ids run to
# `_dense_size` (page 0 is the reserved padding sink).
size=self._dense_size - page_size,
page_size=page_size,
dtype=kv_cache_dtype,
kv_lora_rank=spec.kv_lora_rank,
qk_rope_head_dim=spec.qk_rope_head_dim,
layer_num=spec.layer_num,
device=unified_buffer.device,
enable_memory_saver=False, # buffer owned by UnifiedKVPool
)
def _create_buffers(self):
self.kv_buffer = self._kv_views
def _clear_buffers(self):
# Lifetime owned by UnifiedKVPool; do not delete the views.
pass
def get_kv_size_bytes(self):
return 0 # UnifiedKVPool logs the total; per-sub-pool would double-count
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
"""Relocate whole page envelopes.
`tgt_loc`/`src_loc` are REAL physical token ids (NOT dense ids): both
compaction paths expand page ids into page-major-ordered token runs
(`pages[:, None] * ps + offsets`), relied on here to recover the page
lists. One contiguous envelope copy replaces the per-layer strided moves.
"""
if tgt_loc.numel() == 0:
return
ps = self.page_size
tgt_pages = tgt_loc.view(-1, ps)[:, 0] // ps
src_pages = src_loc.view(-1, ps)[:, 0] // ps
with record_function("UnifiedMLA.move_kv_cache"):
env = self._unified_buffer._raw[: self._num_pages * self._page_bytes].view(
self._num_pages, self._page_bytes
)
env[tgt_pages] = env[src_pages]
class UnifiedMambaPool(MambaPool):
"""Mamba state pool whose conv/temporal state are strided views into a `UnifiedKVPool`.
@@ -766,10 +941,12 @@ class UnifiedHybridReqToTokenPool(HybridReqToTokenPool):
mamba_envelope_layout: bool = False,
enable_linear_replayssm: bool = False,
linear_replayssm_cache_len: int = 16,
enable_gdn_replayssm_spec: bool = False,
):
# mamba_envelope_layout / speculative_eagle_topk / enable_linear_replayssm /
# linear_replayssm_cache_len: accepted to match the parent signature but NOT
# forwarded — the shared pool's conv/temporal state are fixed-shape views.
# linear_replayssm_cache_len / enable_gdn_replayssm_spec: accepted to match
# the parent signature but NOT forwarded — the shared pool's conv/temporal
# state are fixed-shape views (replayssm/spec are gated off under unified).
assert mamba_size == self._shared_mamba_size, (
f"UnifiedHybridReqToTokenPool._init_mamba_pool: mamba_size={mamba_size} "
f"!= unified_buffer.max_slots({self._mamba_sub_pool_name!r}) - 1 "
@@ -837,6 +1014,8 @@ def init_unified_mamba_pools(
end_layer: int,
is_draft_worker: bool,
use_mla_backend: bool,
kv_lora_rank: Optional[int] = None,
qk_rope_head_dim: Optional[int] = None,
mamba_layer_ids: List[int],
full_attention_layer_ids: List[int],
mamba2_cache_params,
@@ -860,22 +1039,37 @@ def init_unified_mamba_pools(
UnifiedMambaTokenToKVPoolAllocator,
)
assert (
not use_mla_backend
), "unified memory pool does not support MLA-hybrid-Mamba yet"
# Full sub-pool is page-aware; mamba stays page=1 (state is per-request).
assert page_size >= 1, f"page_size must be >= 1, got {page_size}"
store_dtype = _store_dtype_for(kv_cache_dtype)
# full-attn at the high-byte end (grow-down), mamba at the low-byte end (grow-up).
full_spec = MHASubPoolSpec(
name="full",
layer_num=len(full_attention_layer_ids),
head_num=head_num,
head_dim=head_dim,
store_dtype=store_dtype,
grow_direction="down",
)
if use_mla_backend:
assert kv_lora_rank and qk_rope_head_dim, (
"init_unified_mamba_pools: MLA-hybrid-Mamba needs kv_lora_rank and "
f"qk_rope_head_dim; got {kv_lora_rank} / {qk_rope_head_dim}"
)
assert not is_draft_worker, (
"init_unified_mamba_pools: draft workers (speculative decoding) are "
"not supported with the MLA unified pool"
)
full_spec = MLASubPoolSpec(
name="full",
layer_num=len(full_attention_layer_ids),
kv_lora_rank=kv_lora_rank,
qk_rope_head_dim=qk_rope_head_dim,
store_dtype=store_dtype,
grow_direction="down",
)
else:
full_spec = MHASubPoolSpec(
name="full",
layer_num=len(full_attention_layer_ids),
head_num=head_num,
head_dim=head_dim,
store_dtype=store_dtype,
grow_direction="down",
)
cp = mamba2_cache_params
mamba_spec = MambaSubPoolSpec(
name="mamba",
@@ -891,12 +1085,16 @@ def init_unified_mamba_pools(
max_total_num_tokens * full_spec.entry_bytes()
+ max_mamba_cache_size * mamba_spec.entry_bytes()
)
# Dense MLA views are per-layer shifted, so the last layer's view reaches one
# page envelope past the final page — allocation-only tail pad (~page bytes).
view_tail_pad_bytes = page_size * full_spec.entry_bytes() if use_mla_backend else 0
shared_pool = UnifiedKVPool(
total_bytes=total_bytes,
sub_pool_specs=[full_spec, mamba_spec],
device=device,
enable_memory_saver=enable_memory_saver,
page_size=page_size,
view_tail_pad_bytes=view_tail_pad_bytes,
)
req_to_token_pool = UnifiedHybridReqToTokenPool(
unified_buffer=shared_pool,
@@ -913,13 +1111,23 @@ def init_unified_mamba_pools(
enable_overlap_schedule=not disable_overlap_schedule,
start_layer=start_layer,
)
unified_full_kv_pool = UnifiedMHATokenToKVPool(
unified_buffer=shared_pool,
sub_pool_name="full",
page_size=page_size,
start_layer=start_layer,
end_layer=end_layer,
)
if use_mla_backend:
# start_layer stays 0: HybridLinearKVPool patches layer ids to the dense
# 0..N-1 index via _transfer_id_context before every MLA pool call.
unified_full_kv_pool = UnifiedMLATokenToKVPool(
unified_buffer=shared_pool,
sub_pool_name="full",
kv_cache_dtype=kv_cache_dtype,
page_size=page_size,
)
else:
unified_full_kv_pool = UnifiedMHATokenToKVPool(
unified_buffer=shared_pool,
sub_pool_name="full",
page_size=page_size,
start_layer=start_layer,
end_layer=end_layer,
)
full_attn_layer_ids_for_pool = (
[0] if is_draft_worker else list(full_attention_layer_ids)
)
@@ -945,6 +1153,9 @@ def init_unified_mamba_pools(
need_sort=need_sort,
forward_stream=forward_stream,
lazy_compaction=lazy_compaction,
full_kernel_page_multiplier=(
len(full_attention_layer_ids) if use_mla_backend else 1
),
)
# Wrap the composite's mamba MultiEndedAllocator in a slot allocator (PHYSICAL view).
@@ -956,23 +1167,43 @@ def init_unified_mamba_pools(
# `_mamba_translate` feeds the HiCache offload path, GATED OFF here — wired but inert.
req_to_token_pool.mamba_allocator = mamba_slot_allocator
token_to_kv_pool._mamba_translate = mamba_slot_allocator.translate
if use_mla_backend:
# Model-level MLA entry points (`set_mla_kv_buffer` / `get_mla_kv_buffer`)
# receive VIRTUAL locs and translate to the dense space internally
# (eager-prefill-only paths; never captured in a cuda graph).
token_to_kv_pool._full_translate = allocator.translate_kv_loc_dense
logger.info(
"[unified-memory-pool] ============================================================"
)
logger.info(
"[unified-memory-pool] UNIFIED MEMORY POOL ENABLED -- path=Mamba hybrid"
)
logger.info(
"[unified-memory-pool] full_layers=%d, mamba_layers=%d, head_num=%d, head_dim=%d, "
"page_size=%d, is_draft_worker=%s",
len(full_attention_layer_ids),
len(mamba_layer_ids),
head_num,
head_dim,
page_size,
is_draft_worker,
"[unified-memory-pool] UNIFIED MEMORY POOL ENABLED -- path=Mamba hybrid (%s full side)",
"MLA" if use_mla_backend else "MHA",
)
if use_mla_backend:
logger.info(
"[unified-memory-pool] full_layers=%d, mamba_layers=%d, kv_lora_rank=%d, "
"qk_rope_head_dim=%d, page_size=%d (dense views, kernel_page_multiplier=%d, "
"view_tail_pad=%d B)",
len(full_attention_layer_ids),
len(mamba_layer_ids),
kv_lora_rank,
qk_rope_head_dim,
page_size,
len(full_attention_layer_ids),
view_tail_pad_bytes,
)
else:
logger.info(
"[unified-memory-pool] full_layers=%d, mamba_layers=%d, head_num=%d, head_dim=%d, "
"page_size=%d, is_draft_worker=%s",
len(full_attention_layer_ids),
len(mamba_layer_ids),
head_num,
head_dim,
page_size,
is_draft_worker,
)
logger.info(
"[unified-memory-pool] total_bytes=%d, max_total_num_tokens=%d, max_mamba_cache_size=%d, "
"max_num_reqs=%d, speculative_num_draft_tokens=%s",