diff --git a/python/sglang/kernels/ops/attention/decode_attention.py b/python/sglang/kernels/ops/attention/decode_attention.py index 3884e1ce4..6c1ec8046 100644 --- a/python/sglang/kernels/ops/attention/decode_attention.py +++ b/python/sglang/kernels/ops/attention/decode_attention.py @@ -195,53 +195,18 @@ def _mla_launch_plan( def _extract_kv_strides(buf, page_size: int): - """Extract (slot_stride, head_stride, page_stride, tok_stride) for a - KV buffer that may be: - - 3-D ``[max_slots, head_num, head_dim]`` (legacy / non-shared) — the - contiguous layout most callers use. page/tok strides are synthesized - so the kernel's PAGE_SIZE>1 math collapses to ``kv_loc * stride(0)``. - - 4-D ``[num_pages, page_size, head_num, head_dim]`` (shared - pool). page/tok strides come from stride(0)/stride(1) directly; - legacy ``stride_bs`` is set to 0 (unused at PAGE_SIZE>1). + """Extract (slot_stride, head_stride, page_stride, tok_stride) for a 3-D + ``[max_slots, head_num, head_dim]`` KV buffer. Returns a 4-tuple of ints suitable for passing as ``stride_buf_*bs``, ``stride_buf_*h``, ``stride_buf_*page``, ``stride_buf_*tok``. """ - if buf.ndim == 4: - # 4-D view ``[num_pages, page_size, head_num, head_dim]``. - # stride(0) = per-PAGE stride (page_bytes/itemsize) - # stride(1) = within-page per-TOKEN stride (k_row/v_row bytes/itemsize) - # The PAGE_SIZE>1 kernel branch uses page_stride/tok_stride and does - # NOT read slot_stride. slot_stride is consumed ONLY by the - # PAGE_SIZE==1 branch (``offs = kv_loc * stride_buf_*bs``), where one - # page holds exactly one slot, so the per-slot stride is the per-page - # stride — NOT the within-page token stride. Concretely the per-slot - # stride is ``page_stride // page_size`` (= entry_bytes/itemsize), - # which at ps=1 equals page_stride. Using ``tok_stride`` here (one - # layer's k_row) would make the ps=1 read address ``kv_loc * k_row`` - # instead of ``kv_loc * entry_bytes`` and read the wrong slot. - page_stride = buf.stride(0) - tok_stride = buf.stride(1) - head_stride = buf.stride(2) - slot_stride = ( - page_stride // page_size - ) # per-slot stride; == page_stride at ps=1 - assert buf.shape[1] == page_size, ( - f"4-D KV buffer's dim-1 must equal page_size; got " - f"shape[1]={buf.shape[1]}, page_size={page_size}" - ) - elif buf.ndim == 3: - # Legacy 3-D ``[N, head, dim]``. Synthesize page/tok strides such - # that ``(kv_loc // ps) * page_stride + (kv_loc % ps) * tok_stride - # == kv_loc * slot_stride`` for the page-aware branch — this lets - # the same kernel handle non-shared paged-allocator buffers without - # any caller adjustment. - slot_stride = buf.stride(0) - head_stride = buf.stride(1) - page_stride = slot_stride * page_size - tok_stride = slot_stride - else: # pragma: no cover + if buf.ndim != 3: raise ValueError(f"unexpected KV buffer ndim={buf.ndim}, shape={buf.shape}") + slot_stride = buf.stride(0) + head_stride = buf.stride(1) + page_stride = slot_stride * page_size + tok_stride = slot_stride return slot_stride, head_stride, page_stride, tok_stride @@ -464,9 +429,6 @@ def _decode_att_m_fwd( Lk = k_buffer.shape[-1] Lv = v_buffer.shape[-1] - # head_num lives in the dim immediately before the head_dim. For 3-D - # ``[N, head_num, head_dim]`` that's dim 1; for 4-D - # ``[num_pages, page_size, head_num, head_dim]`` that's dim 2. kv_head_num = k_buffer.shape[-2] batch, head_num = q.shape[0], q.shape[1] @@ -1204,8 +1166,7 @@ def decode_attention_fwd( # schedule from kv_indptr on-device, so this path involves no host sync and is safe to # capture in a CUDA graph. Whether Lean pays off for a given shape is decided cheaply by # the backend's host-side seqlen gate (lean_decode_seqlen_gate) before we get here. - # Lean supports both the contiguous 3-D [N, head, dim] and paged 4-D - # [num_pages, page_size, head, dim] KV layouts (page-aware address math in the kernel). + # Lean handles both page sizes: the kernel does page-aware address math. # ROCm/AMD only: Lean is validated on MI300X/MI355X; CUDA/NVIDIA uses the standard kernel. if ( _is_hip @@ -1587,8 +1548,8 @@ def _lean_attention_decode_kernel( # Load K transposed: [BLOCK_DMODEL, BLOCK_N] so qk = q @ k directly. # Page-aware KV address math (mirrors the standard grouped kernel): at - # PAGE_SIZE==1 the slot index addresses directly; otherwise it splits into - # (page_id, tok_in_p) for a [num_pages, page_size, head, dim] paged buffer. + # PAGE_SIZE==1 the slot index addresses directly; otherwise it splits + # into (page_id, tok_in_p). if PAGE_SIZE == 1: offs_buf_k = ( kv_loc[None, :] * stride_buf_kbs @@ -1968,13 +1929,11 @@ def _decode_lean_attention_fwd( ``total_programs`` is the fixed persistent-grid size (2× device CU count). The kernel derives its own tile schedule from ``kv_indptr`` on-device, so no host sync is needed and the launch is CUDA-graph capturable. ``Mp``, ``Lp``, ``Op``, ``locks`` are pre-allocated - persistent-grid partial-result buffers reused across decode steps. ``page_size`` selects - the KV address math: 1 for a contiguous ``[N, head, dim]`` buffer, >1 for a paged - ``[num_pages, page_size, head, dim]`` buffer (strides via ``_extract_kv_strides``). + persistent-grid partial-result buffers reused across decode steps. ``page_size`` + selects the KV address math over the ``[N, head, dim]`` buffer (strides via + ``_extract_kv_strides``). """ batch, head_num = q.shape[0], q.shape[1] - # head_num lives at dim -2 for both the 3-D [N, head, dim] and 4-D paged - # [num_pages, page_size, head, dim] layouts. num_kv_heads = k_buffer.shape[-2] Lk = k_buffer.shape[-1] Lv = v_buffer.shape[-1] diff --git a/python/sglang/kernels/ops/attention/metadata.py b/python/sglang/kernels/ops/attention/metadata.py index 545ae24b3..54762daff 100644 --- a/python/sglang/kernels/ops/attention/metadata.py +++ b/python/sglang/kernels/ops/attention/metadata.py @@ -193,7 +193,7 @@ def _fused_metadata_kernel_general( use_swa: tl.constexpr, SHIFT: tl.constexpr, BLOCK_COLS: tl.constexpr, - # Unified-memory dense-view path (page-major envelope shared with the mamba + # Unified-memory per-layer-view path (page-major envelope shared with the mamba # sub-pool). Both default to the identity for the statically-partitioned # pool, where req_to_token already holds physical ids. v2p_ptr=None, @@ -316,7 +316,7 @@ def _fused_metadata_kernel_ps1_no_swa( max_seq_pages, seq_len_delta: tl.constexpr, BLOCK_COLS: tl.constexpr, - # Unified-memory dense-view path; identity defaults for the static pool. + # Unified-memory per-layer-view path; identity defaults for the static pool. v2p_ptr=None, PAGE_MULT: tl.constexpr = 1, ): diff --git a/python/sglang/kernels/ops/kvcache/__init__.py b/python/sglang/kernels/ops/kvcache/__init__.py index 40cbb6f98..e2ee65933 100644 --- a/python/sglang/kernels/ops/kvcache/__init__.py +++ b/python/sglang/kernels/ops/kvcache/__init__.py @@ -72,7 +72,6 @@ _TRITON_KERNELS = [ ("trtllm_mha_graph_metadata", "update_trtllm_mha_graph_metadata"), ("aiter_unified_attention", "scatter_ragged_to_page_table_kernel"), ("aiter_unified_attention", "scatter_req_to_token_to_page_table_kernel"), - ("cache_move", "store_cache_4d"), ("cache_move", "set_kv_buffer_prefix_valid_tiled"), ("cache_move", "copy_all_layer_kv_cache_tiled"), ("mla_buffer", "set_mla_kv_buffer_triton"), diff --git a/python/sglang/kernels/ops/kvcache/cache_move.py b/python/sglang/kernels/ops/kvcache/cache_move.py index e4e408f9b..cf0e0e4cc 100644 --- a/python/sglang/kernels/ops/kvcache/cache_move.py +++ b/python/sglang/kernels/ops/kvcache/cache_move.py @@ -122,185 +122,3 @@ def copy_all_layer_kv_cache_func( num_warps=kv_copy_config["num_warps"], num_stages=2, ) - - -# --------------------------------------------------------------------------- -# store_cache_4d — single-launch Triton write into the 4-D page-major envelope -# K/V views. At `PAGE_SIZE = 1` the kernel constexpr-folds to byte-identical -# addresses as the slot-major envelope view; at `PAGE_SIZE > 1` it uses the -# same `(page_id, tok_in_p)` split the attention read kernels use. -# --------------------------------------------------------------------------- - - -@triton.jit -def store_cache_4d_kernel( - k_view_ptr, - v_view_ptr, - cache_k_ptr, - cache_v_ptr, - loc_ptr, - # Strides in ELEMENTS (not bytes); wrapper passes view.stride(D) - # directly. K and V may have different head_dim → different per-token - # strides, so we carry both. - stride_k_page, - stride_k_tok, - stride_v_page, - stride_v_tok, - stride_src_k_row, - stride_src_v_row, - K_ROW_DIM: tl.constexpr, # head_num * head_dim - V_ROW_DIM: tl.constexpr, # head_num * v_head_dim - PAGE_SIZE: tl.constexpr, - BLOCK: tl.constexpr, -): - """Token-parallel Triton write into a 4-D envelope-strided K/V view. - - Grid: ``(N, ceil(max(K_ROW_DIM, V_ROW_DIM) / BLOCK), 2)`` where: - - axis 0 → one program per token (loc[i]) - - axis 1 → blocks within one slot's K (or V) row - - axis 2 → 0 = K, 1 = V (two-tensor write fused into one launch) - - For each token i, the kernel writes: - page_id = loc[i] // PAGE_SIZE - tok_in_p = loc[i] % PAGE_SIZE - k_view[page_id, tok_in_p, :, :] = cache_k[i, :, :] - v_view[page_id, tok_in_p, :, :] = cache_v[i, :, :] - - Cuda-graph safe: no Python branching on tensor values, no `.item()`, - all shapes/strides known at launch time. - """ - pid_n = tl.program_id(0) - pid_b = tl.program_id(1) - pid_kv = tl.program_id(2) - - # 1. Resolve destination slot in the 4-D view. - loc = tl.load(loc_ptr + pid_n).to(tl.int64) - if PAGE_SIZE == 1: - page_id = loc - tok_in_p = tl.zeros([], dtype=tl.int64) - else: - page_id = loc // PAGE_SIZE - tok_in_p = loc % PAGE_SIZE - - # 2. Compute per-tensor source/dest pointers. - base_off = pid_b * BLOCK + tl.arange(0, BLOCK) - - if pid_kv == 0: - mask = base_off < K_ROW_DIM - # The trailing (head_num, head_dim) axes of `k_view` are - # contiguous: stride[-1]==1, stride[-2]==head_dim. So we can - # treat them as a flat K_ROW_DIM dimension addressed by `base_off`. - # The wrapper asserts this invariant. - src_ptr = cache_k_ptr + pid_n * stride_src_k_row + base_off - dst_ptr = ( - k_view_ptr + page_id * stride_k_page + tok_in_p * stride_k_tok + base_off - ) - else: - mask = base_off < V_ROW_DIM - src_ptr = cache_v_ptr + pid_n * stride_src_v_row + base_off - dst_ptr = ( - v_view_ptr + page_id * stride_v_page + tok_in_p * stride_v_tok + base_off - ) - - src = tl.load(src_ptr, mask=mask) - tl.store(dst_ptr, src, mask=mask) - - -def store_cache_4d( - k_view: torch.Tensor, - v_view: torch.Tensor, - cache_k: torch.Tensor, - cache_v: torch.Tensor, - loc: torch.Tensor, - page_size: int, -) -> None: - """One-launch Triton write into the 4-D page-major envelope K/V views. - - Writes ``cache_k[i]`` and ``cache_v[i]`` to - ``k_view[loc[i]//ps, loc[i]%ps, :, :]`` (and analogously for V) for - ``i in [0, N)``. - - Contract: - - ``k_view``, ``v_view``: 4-D ``(num_pages, page_size, head_num, - head_dim*)``, contiguous in the trailing ``(head_num, head_dim)`` - dims (i.e., ``stride[-1] == 1`` and ``stride[-2] == head_dim``). - - ``cache_k``, ``cache_v``: 3-D ``(N, head_num, head_dim*)``, - contiguous in the trailing ``(head_num, head_dim)`` dims. - - ``loc``: 1-D int64 or int32, N elements, values in - ``[0, num_pages * page_size)``. The caller is responsible for - clamping any negative entries to ≥ 0. - - At ``page_size == 1`` the kernel produces byte-identical output - to the legacy advanced-indexing path. - - Returns nothing; writes in place. - """ - if loc.numel() == 0: - return - assert k_view.is_cuda and v_view.is_cuda, "store_cache_4d: CUDA only" - assert k_view.ndim == 4 and v_view.ndim == 4, ( - f"store_cache_4d: k_view/v_view must be 4-D, " - f"got {k_view.ndim}/{v_view.ndim}" - ) - assert cache_k.ndim == 3 and cache_v.ndim == 3, ( - f"store_cache_4d: cache_k/cache_v must be 3-D, " - f"got {cache_k.ndim}/{cache_v.ndim}" - ) - assert cache_k.shape[0] == cache_v.shape[0] == loc.numel(), ( - "store_cache_4d: cache_k/cache_v/loc batch dim mismatch: " - f"{cache_k.shape[0]}, {cache_v.shape[0]}, {loc.numel()}" - ) - assert k_view.dtype == v_view.dtype == cache_k.dtype == cache_v.dtype, ( - "store_cache_4d: dtype mismatch: " - f"k_view={k_view.dtype}, v_view={v_view.dtype}, " - f"cache_k={cache_k.dtype}, cache_v={cache_v.dtype}" - ) - # Stride invariants — the kernel addresses (head_num, head_dim) as one - # flat ROW_DIM dimension; this requires the trailing two dims to be - # contiguous. This holds for the page-major envelope views - # (k_stride = (page_bytes/itemsize, k_row_bytes/itemsize, head_dim, 1)) and - # for cache_k/cache_v produced by the model forward. - assert k_view.stride(-1) == 1 and k_view.stride(-2) == k_view.shape[-1], ( - f"store_cache_4d: k_view trailing dims must be contiguous; " - f"got stride={k_view.stride()}, shape={tuple(k_view.shape)}" - ) - assert v_view.stride(-1) == 1 and v_view.stride(-2) == v_view.shape[-1], ( - f"store_cache_4d: v_view trailing dims must be contiguous; " - f"got stride={v_view.stride()}, shape={tuple(v_view.shape)}" - ) - assert cache_k.stride(-1) == 1 and cache_k.stride(-2) == cache_k.shape[-1], ( - f"store_cache_4d: cache_k trailing dims must be contiguous; " - f"got stride={cache_k.stride()}, shape={tuple(cache_k.shape)}" - ) - assert cache_v.stride(-1) == 1 and cache_v.stride(-2) == cache_v.shape[-1], ( - f"store_cache_4d: cache_v trailing dims must be contiguous; " - f"got stride={cache_v.stride()}, shape={tuple(cache_v.shape)}" - ) - - head_num = k_view.shape[2] - head_dim = k_view.shape[3] - v_head_dim = v_view.shape[3] - K_ROW_DIM = head_num * head_dim - V_ROW_DIM = head_num * v_head_dim - BLOCK = 128 - N = loc.numel() - row_dim_max = max(K_ROW_DIM, V_ROW_DIM) - grid = (N, triton.cdiv(row_dim_max, BLOCK), 2) - - store_cache_4d_kernel[grid]( - k_view, - v_view, - cache_k, - cache_v, - loc, - k_view.stride(0), - k_view.stride(1), - v_view.stride(0), - v_view.stride(1), - cache_k.stride(0), - cache_v.stride(0), - K_ROW_DIM=K_ROW_DIM, - V_ROW_DIM=V_ROW_DIM, - PAGE_SIZE=page_size, - BLOCK=BLOCK, - num_warps=4, - ) diff --git a/python/sglang/kernels/ops/kvcache/kv_indices.py b/python/sglang/kernels/ops/kvcache/kv_indices.py index 1ee8804e6..4f5f75f08 100644 --- a/python/sglang/kernels/ops/kvcache/kv_indices.py +++ b/python/sglang/kernels/ops/kvcache/kv_indices.py @@ -105,11 +105,11 @@ def create_flashmla_kv_indices_triton( req_to_token_ptr_stride: tl.constexpr, kv_indices_ptr_stride: tl.constexpr, PAGED_SIZE: tl.constexpr = 64, - # Unified-memory dense-view path (page-major envelope shared with the mamba + # Unified-memory per-layer-view path (page-major envelope shared with the mamba # sub-pool). req_to_token holds VIRTUAL token ids; the block table the MLA - # kernel consumes must hold DENSE page ids. When v2p_ptr is given, map each + # kernel consumes must hold kernel-facing page ids. When v2p_ptr is given, map each # virtual page through it to the physical page, then scale by PAGE_MULT - # (= num MLA layers) so the entry addresses the layer's dense per-page block + # (= num MLA layers) so the entry addresses the layer's per-page block # in the (num_pages*L, page_size, kv_cache_dim) reshaped view. Both default # to the identity (v2p_ptr None, PAGE_MULT 1) for the static pool. v2p_ptr=None, diff --git a/python/sglang/srt/arg_groups/kv_cache_hook.py b/python/sglang/srt/arg_groups/kv_cache_hook.py index def9f3dd9..5a58b474e 100644 --- a/python/sglang/srt/arg_groups/kv_cache_hook.py +++ b/python/sglang/srt/arg_groups/kv_cache_hook.py @@ -9,6 +9,7 @@ from typing import Any from sglang.srt.arg_groups.overrides import ( attention_backends_of, declare_resolution, + model_config_of, resolved_view, resolving_view, use_mla_backend, @@ -223,7 +224,7 @@ def handle_unified_memory_pool(server_args: Any) -> None: assert cfg.speculative_algorithm in (None, "DSPARK"), ( "--enable-unified-memory only supports --speculative-algorithm " "DSPARK (chain draft); other speculative algorithms are not yet " - "audited for the unified pool's virtual/dense loc translation. Got " + "audited for the unified pool's virtual/kernel-facing loc translation. Got " f"--speculative-algorithm={cfg.speculative_algorithm!r}." ) if cfg.speculative_algorithm == "DSPARK": @@ -243,7 +244,7 @@ def handle_unified_memory_pool(server_args: Any) -> None: f"attention backends {sorted(spec_allowed)} for both prefill " f"and decode; got {sorted(spec_backends)}. flashinfer / fa3 do " "not translate speculative verify indices to the unified " - "pool's dense space yet." + "pool's kernel-facing space yet." ) assert not (cfg.enable_hierarchical_cache or cfg.enable_lmcache), ( "--enable-unified-memory is not yet compatible with hierarchical / " @@ -283,15 +284,39 @@ def handle_page_major_kv_layout(server_args: Any): ) if not cfg.enable_page_major_kv_layout: return + assert cfg.enable_unified_memory, ( + "--enable-page-major-kv-layout without --enable-unified-memory is " + "temporarily unsupported: the strided MHA K/V views were removed " + "and the static-pool page-major layout awaits its per-layer-view " + "reimplementation. Run with --enable-unified-memory, or drop " + "--enable-page-major-kv-layout." + ) + from sglang.srt.mem_cache.unified_memory_pool import ( + unified_memory_supported_for_model, + ) + + model_config = model_config_of(server_args) + assert unified_memory_supported_for_model( + model_config, use_mla_backend=use_mla_backend(server_args) + ), ( + "--enable-unified-memory requires uniform K/V rows " + "(head_dim == v_head_dim); this model has " + f"head_dim={model_config.head_dim}, " + f"v_head_dim={model_config.v_head_dim}, " + f"swa_head_dim={model_config.swa_head_dim}, " + f"swa_v_head_dim={model_config.swa_v_head_dim}. The unified " + "pool's per-layer views require a uniform row width; run " + "this model without --enable-unified-memory." + ) # Only the Triton attention kernels read the strided 4-D envelope K/V # views; FA3 / FlashInfer do not. EXCEPTION: the unified-memory MLA pool - # exposes each layer as a DENSE contiguous per-layer view - # (build_dense_mla_views), which the paged MLA kernels consume directly, - # with their kv_indices / block tables remapped to dense ids. Names below + # exposes each layer as a contiguous per-layer view + # (build_mla_views), which the paged MLA kernels consume directly, + # with their kv_indices / block tables remapped to kernel-facing ids. Names below # are the RESOLVED ids from attention_backends_of: "flashinfer" is # FlashInferMLAAttnBackend for an MLA model, "trtllm_mla" the trtllm # decode kernel; "cutedsl_mla" and "tokenspeed_mla" subclass - # TRTLLMMLABackend and inherit its dense read/write path; "fa3" remaps its + # TRTLLMMLABackend and inherit its read/write path; "fa3" remaps its # page_table (in-kernel for captured decode, one funnel for eager). # flashmla / cutlass_mla share the create_flashmla block-table path and # can be added the same way once exercised. diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 04280dff4..f3cd3b29f 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -193,7 +193,7 @@ class FlashAttentionBackend(AttentionBackend): self.needs_cpu_seq_lens = False self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA # Unified pool: req_to_token holds VIRTUAL ids but the MLA per-layer views - # are DENSE, so every page_table needs remapping. MLA-only -- the MHA/SWA + # are kernel-facing, so every page_table needs remapping. MLA-only -- the MHA/SWA # sub-pools keep the strided envelope layout FA3 cannot read at all. self._unified_hooks = unified_mla_hooks(model_runner.token_to_kv_pool_allocator) self._unified_dense = self._unified_hooks.enabled and self.use_mla @@ -1101,16 +1101,16 @@ class FlashAttentionBackend(AttentionBackend): # the remap into normal_decode_set_metadata, which must write in place. # # Placed BEFORE the `// page_size` reduction, in token space: since - # dense(t) = phys_page * (ps * L) + t % ps, dense(page_start) // ps is + # kernel_id(t) = phys_page * (ps * L) + t % ps, dense(page_start) // ps is # phys_page * L, the dense page id the kernel wants. One site then serves - # both page sizes, and it inherits translate_kv_loc_dense's tombstone + # both page sizes, and it inherits translate_kv_loc_for_kernel's tombstone # clamp so an unwritten req_to_token slot lands in the page-0 sink. if self._unified_dense and metadata.page_table is not None: # Flattened: the page_size == 1 translate path uses index_select, # which rejects a 2-D index. pt = metadata.page_table metadata.page_table = ( - self._unified_hooks.translate_kv_loc_dense(pt.reshape(-1)) + self._unified_hooks.translate_kv_loc_for_kernel(pt.reshape(-1)) .to(torch.int32) .view(pt.shape) ) diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index eed97d760..295cd60a0 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -850,7 +850,7 @@ class FlashInferMLAIndicesUpdaterDecode: # Unified dense MLA pool: VIRTUAL -> DENSE kv_indices (see prefill updater). self._translate_kv_loc_dense = unified_mla_hooks( model_runner.token_to_kv_pool_allocator - ).translate_kv_loc_dense + ).translate_kv_loc_for_kernel def update( self, @@ -918,7 +918,7 @@ class FlashInferMLAIndicesUpdaterDecode: # [:paged_kernel_lens_sum] prefix the index kernel just filled is # translated; the stale tail is left alone so it can never index the # v2p table out of bounds. The int64 translate result narrows back to - # the buffer's int32 on copy_ (flashinfer requires int32; dense ids + # the buffer's int32 on copy_ (flashinfer requires int32; kernel-facing ids # fit comfortably). if self._translate_kv_loc_dense is not None: valid = kv_indices[:paged_kernel_lens_sum] @@ -989,11 +989,11 @@ class FlashInferMLAIndicesUpdaterPrefill: self.req_to_token = model_runner.req_to_token_pool.req_to_token self.prefill_wrapper_ragged = attn_backend.prefill_wrapper_ragged # Unified dense MLA pool: kv_indices built from req_to_token are VIRTUAL; - # the paged wrapper reads the dense per-layer view, so remap them to DENSE + # the paged wrapper reads the per-layer view, so remap them to kernel-facing # token ids. None (identity) unless the unified MLA pool is active. self._translate_kv_loc_dense = unified_mla_hooks( model_runner.token_to_kv_pool_allocator - ).translate_kv_loc_dense + ).translate_kv_loc_for_kernel def update( self, diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 018157b50..64dbe6405 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -210,11 +210,11 @@ class TritonAttnBackend(AttentionBackend): 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 the kernel-facing id space — PHYSICAL for MHA, DENSE for the - # dense-view MLA pool (translate_kv_loc_dense falls back to the physical + # per-layer-view MLA pool (translate_kv_loc_for_kernel 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_dense", None + self.token_to_kv_pool_allocator, "translate_kv_loc_for_kernel", None ) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None) self.num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_steps = get_spec().speculative_num_steps diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index a64a60e4a..a43b28c78 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -285,22 +285,22 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): # Tree-mask scratch is fetched from the target backend only. self.is_draft_runner = model_runner.is_draft_worker - # Unified-memory dense-view hooks (None on the static pool). req_to_token - # holds VIRTUAL token ids; the block table needs DENSE page ids, so the + # Unified-memory per-layer-view hooks (None on the static pool). req_to_token + # holds VIRTUAL token ids; the block table needs kernel-facing page ids, so the # kv-index kernels gather virtual->physical page through `_v2p_page_table` # then scale by `_kernel_page_multiplier` (= num MLA layers). See - # build_dense_mla_views / create_flashmla_kv_indices_triton. + # build_mla_views / create_flashmla_kv_indices_triton. _hooks = unified_mla_hooks(model_runner.token_to_kv_pool_allocator) self._v2p_page_table = _hooks.v2p_page_table self._kernel_page_multiplier = _hooks.kernel_page_multiplier self._unified_mla = _hooks.enabled # virtual token id -> DENSE kernel-facing id, for the KV write loc. - self._translate_kv_loc_dense = _hooks.translate_kv_loc_dense - # Per-forward dense write loc ([:n] view of a capture-stable buffer), + self._translate_kv_loc_dense = _hooks.translate_kv_loc_for_kernel + # Per-forward kernel-facing write loc ([:n] view of a capture-stable buffer), # set by the cuda-graph out-graph hook; None on the eager path (where the # write translates through the pool's _full_translate hook instead). - self._decode_dense_loc: Optional[torch.Tensor] = None - self.cuda_graph_out_cache_loc_dense: Optional[torch.Tensor] = None + self._decode_kernel_loc: Optional[torch.Tensor] = None + self.cuda_graph_out_cache_loc_kernel: Optional[torch.Tensor] = None # Fused KV-scatter + q-concat on the decode dense-loc path (one launch # instead of set_mla_kv_buffer + concat_mla_absorb_q). Disabled under # async asserts: the fused path writes the pool directly and would @@ -408,7 +408,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): # out-of-graph in init_forward_metadata_out_graph so the in-graph # set_mla_kv_buffer captures no translate. if self._unified_mla: - self.cuda_graph_out_cache_loc_dense = torch.zeros( + self.cuda_graph_out_cache_loc_kernel = torch.zeros( max_num_tokens, dtype=torch.int64, device=self.device ) num_tokens_per_req = max_num_tokens // max_bs @@ -635,7 +635,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): ): out_cache_loc = forward_batch.out_cache_loc n = out_cache_loc.shape[0] - dst = self.cuda_graph_out_cache_loc_dense[:n] + dst = self.cuda_graph_out_cache_loc_kernel[:n] self._translate_kv_loc_dense(out_cache_loc, out=dst) # Replay-prep receives the RAW (unpadded) out_cache_loc # (build_replay_fb_view), but the captured write kernel consumes the @@ -644,16 +644,16 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): # earlier larger replays — a stale tail scatters pad-row garbage into # live KV pages. Mirrors the runner's PaddingPolicy.ZERO on its own # out_cache_loc slot. - self.cuda_graph_out_cache_loc_dense[n:].zero_() - self._decode_dense_loc = dst + self.cuda_graph_out_cache_loc_kernel[n:].zero_() + self._decode_kernel_loc = dst else: - self._decode_dense_loc = None + self._decode_kernel_loc = None def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize the metadata for a forward pass.""" - # Eager path: no capture-stable dense write loc; the pool's _full_translate + # Eager path: no capture-stable kernel-facing write loc; the pool's _full_translate # hook translates the write loc (safe out of a cuda graph). - self._decode_dense_loc = None + self._decode_kernel_loc = None # Delegate to parent for non-decode modes. if ( forward_batch.forward_mode.is_extend() @@ -1058,8 +1058,8 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): if cos_sin_cache is None: if save_kv_cache and self._fused_set_kv_concat_q_fp8: loc = ( - self._decode_dense_loc - if self._decode_dense_loc is not None + self._decode_kernel_loc + if self._decode_kernel_loc is not None else ( None if self._unified_mla else forward_batch.out_cache_loc ) @@ -1099,15 +1099,15 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): assert ( k is not None and k_rope is not None ), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None." - if self._decode_dense_loc is not None: - # cuda-graph path: dense write loc precomputed out-of-graph, so + if self._decode_kernel_loc is not None: + # cuda-graph path: kernel-facing write loc precomputed out-of-graph, so # the in-graph write captures no translate allocation. if merge_query and self._fused_set_kv_concat_q: # Fused: KV scatter + [q_nope | q_rope] concat in one # launch; None when the inputs are not covered. query = self._set_kv_and_concat_q_fused( layer=layer, - loc=self._decode_dense_loc, + loc=self._decode_kernel_loc, k=k, k_rope=k_rope, q=q, @@ -1115,7 +1115,11 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): ) if query is None: self.token_to_kv_pool.set_mla_kv_buffer( - layer, self._decode_dense_loc, k, k_rope, loc_is_dense=True + layer, + self._decode_kernel_loc, + k, + k_rope, + loc_is_kernel_facing=True, ) else: # eager (or static pool): the pool's _full_translate handles it. @@ -1259,9 +1263,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): assert ( k is not None and k_rope is not None ), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None." - if self._decode_dense_loc is not None: + if self._decode_kernel_loc is not None: self.token_to_kv_pool.set_mla_kv_buffer( - layer, self._decode_dense_loc, k, k_rope, loc_is_dense=True + layer, self._decode_kernel_loc, k, k_rope, loc_is_kernel_facing=True ) else: self.token_to_kv_pool.set_mla_kv_buffer( diff --git a/python/sglang/srt/layers/attention/unified_mem_hooks.py b/python/sglang/srt/layers/attention/unified_mem_hooks.py index c15f39283..6839926f6 100644 --- a/python/sglang/srt/layers/attention/unified_mem_hooks.py +++ b/python/sglang/srt/layers/attention/unified_mem_hooks.py @@ -37,7 +37,7 @@ class UnifiedMLAHooks(msgspec.Struct, frozen=True): # Page-level virtual->physical table, gathered through by block-table kernels. v2p_page_table: Optional[torch.Tensor] # Virtual token id -> DENSE kernel-facing id (tombstones clamped to the sink). - translate_kv_loc_dense: Optional[Callable[..., torch.Tensor]] + translate_kv_loc_for_kernel: Optional[Callable[..., torch.Tensor]] # Dense page stride scale (= number of full-attention MLA layers). kernel_page_multiplier: int enabled: bool @@ -45,18 +45,18 @@ class UnifiedMLAHooks(msgspec.Struct, frozen=True): _STATIC_POOL = UnifiedMLAHooks( v2p_page_table=None, - translate_kv_loc_dense=None, + translate_kv_loc_for_kernel=None, kernel_page_multiplier=1, enabled=False, ) def unified_mla_hooks(allocator) -> UnifiedMLAHooks: - """Probe ``allocator`` for the unified-pool dense-view hooks. + """Probe ``allocator`` for the unified-pool per-layer-view hooks. Detection keys on the v2p table, NOT on ``kernel_page_multiplier > 1``: a rank owning exactly ONE full-attention layer has multiplier 1 while its - ``req_to_token`` is still virtual. There the dense id collapses onto the + ``req_to_token`` is still virtual. There the kernel-facing id collapses onto the physical id, so the v2p gather alone is the whole translation. """ v2p = getattr(allocator, "full_v2p_page_table", None) @@ -64,7 +64,9 @@ def unified_mla_hooks(allocator) -> UnifiedMLAHooks: return _STATIC_POOL return UnifiedMLAHooks( v2p_page_table=v2p, - translate_kv_loc_dense=getattr(allocator, "translate_kv_loc_dense", None), + translate_kv_loc_for_kernel=getattr( + allocator, "translate_kv_loc_for_kernel", None + ), kernel_page_multiplier=getattr(allocator, "kernel_page_multiplier", 1), enabled=True, ) diff --git a/python/sglang/srt/mem_cache/layout/page_major.py b/python/sglang/srt/mem_cache/layout/page_major.py index 7f91215a2..ff8158318 100644 --- a/python/sglang/srt/mem_cache/layout/page_major.py +++ b/python/sglang/srt/mem_cache/layout/page_major.py @@ -10,8 +10,8 @@ Across pages the layout is envelope-major (one ``page_bytes`` block per page). At ``page_size == 1`` a page is a single slot, so the within-page block is the per-slot ``[L0_K | L0_V | L1_K | L1_V | ...]`` envelope (token-granularity). -These builders produce per-layer strided views into a raw ``uint8`` buffer; they -hold no allocator/ownership state. ``anchor_bytes`` is the byte offset of the +These builders produce per-layer views into a raw ``uint8`` buffer; they hold +no allocator/ownership state. ``anchor_bytes`` is the byte offset of the pool's region inside the raw buffer (0 for a standalone pool). """ @@ -36,7 +36,7 @@ def mha_entry_bytes( return layer_num * (k_row_bytes + v_row_bytes) -def build_page_major_mha_views( +def build_mha_views( raw: torch.Tensor, *, layer_num: int, @@ -48,63 +48,59 @@ def build_page_major_mha_views( num_pages: int, anchor_bytes: int = 0, ) -> Tuple[List[torch.Tensor], List[torch.Tensor]]: - """Per-layer K/V views over ``raw`` in the page-major layer-major layout. + """Per-layer K/V views over ``raw`` for uniform-row MHA. - Each returned view is 4-D ``(num_pages, page_size, head_num, head_dim*)`` - with constant strides: + The page envelope ``[L0_K*ps | L0_V*ps | L1_K*ps | ...]`` is a uniform + array of ``2*layer_num`` row-blocks when K and V rows are equally wide, so + it is a valid paged pool under - stride[0] = page_bytes / itemsize # next page - stride[1] = k_row_bytes / itemsize # next slot within layer L's K block - stride[2] = head_dim # next head - stride[3] = 1 # next element + kernel_id(t) = (t // ps) * (ps * 2 * layer_num) + t % ps - V is analogous with ``v_row_bytes`` / ``v_head_dim``. A token id ``t`` reads - page ``t // page_size``, slot ``t % page_size``. + with layer ``l``'s K at block ``2l`` and its V at block ``2l+1``. Each view + is a contiguous ``(num_pages * 2 * layer_num * ps, head_num, head_dim)``. + + Views overlap by ``ps`` rows per block, safe because an id always resolves + inside its own block; the last view runs ``(2*layer_num - 1) * ps`` rows + past the envelope, so ``raw`` needs ``UnifiedKVPool.view_tail_pad_bytes``. """ + assert head_dim == v_head_dim, ( + f"build_mha_views requires uniform rows (head_dim == v_head_dim); " + f"got head_dim={head_dim}, v_head_dim={v_head_dim}. Asymmetric-KV " + "models cannot use the unified pool (screened out at startup)." + ) itemsize = store_dtype.itemsize - k_row_bytes = head_num * head_dim * itemsize - v_row_bytes = head_num * v_head_dim * itemsize - entry_bytes = layer_num * (k_row_bytes + v_row_bytes) - page_bytes = page_size * entry_bytes + row_elems = head_num * head_dim + row_bytes = row_elems * itemsize + blocks = 2 * layer_num + page_bytes = page_size * blocks * row_bytes + n_rows = num_pages * blocks * page_size assert anchor_bytes % itemsize == 0 - assert k_row_bytes % itemsize == 0 - assert v_row_bytes % itemsize == 0 - assert page_bytes % itemsize == 0 + last_view_end = ( + anchor_bytes + (blocks - 1) * page_size * row_bytes + n_rows * row_bytes + ) + assert last_view_end <= raw.numel() * raw.itemsize, ( + f"build_mha_views: block {blocks - 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) - stride_page = page_bytes // itemsize - stride_tok_k = k_row_bytes // itemsize - stride_tok_v = v_row_bytes // itemsize - - k_shape = (num_pages, page_size, head_num, head_dim) - v_shape = (num_pages, page_size, head_num, v_head_dim) - k_stride = (stride_page, stride_tok_k, head_dim, 1) - v_stride = (stride_page, stride_tok_v, v_head_dim, 1) - k_buffer: List[torch.Tensor] = [] v_buffer: List[torch.Tensor] = [] for layer in range(layer_num): - # Layer L's K block starts at L * page_size * (k_row + v_row); V follows. - k_base_bytes = anchor_bytes + layer * page_size * (k_row_bytes + v_row_bytes) - v_base_bytes = k_base_bytes + page_size * k_row_bytes - assert k_base_bytes % itemsize == 0 - assert v_base_bytes % itemsize == 0 - k_buffer.append( - torch.as_strided( - as_dtype_view, - size=k_shape, - stride=k_stride, - storage_offset=k_base_bytes // itemsize, + k_base_bytes = anchor_bytes + (2 * layer) * page_size * row_bytes + v_base_bytes = k_base_bytes + page_size * row_bytes + for base_bytes, out in ((k_base_bytes, k_buffer), (v_base_bytes, v_buffer)): + assert base_bytes % itemsize == 0 + out.append( + torch.as_strided( + as_dtype_view, + size=(n_rows, head_num, head_dim), + stride=(row_elems, head_dim, 1), + storage_offset=base_bytes // itemsize, + ) ) - ) - v_buffer.append( - torch.as_strided( - as_dtype_view, - size=v_shape, - stride=v_stride, - storage_offset=v_base_bytes // itemsize, - ) - ) return k_buffer, v_buffer @@ -113,7 +109,7 @@ def mla_entry_bytes(*, layer_num: int, kv_cache_dim: int, itemsize: int) -> int: return layer_num * kv_cache_dim * itemsize -def build_dense_mla_views( +def build_mla_views( raw: torch.Tensor, *, layer_num: int, @@ -123,23 +119,23 @@ def build_dense_mla_views( num_pages: int, anchor_bytes: int = 0, ) -> List[torch.Tensor]: - """Per-layer DENSE views over ``raw`` for MLA in the page-major layout. + """Per-layer 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 + itself a valid 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 + kv_cache_dim)`` tensor, addressed by the layer-independent kernel-facing id - dense(t) = (t // ps) * (ps * layer_num) + t % ps (t = physical token) + kernel_id(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, + rows); that is safe because layer ``l`` is only ever indexed at kernel-facing 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 @@ -148,13 +144,13 @@ def build_dense_mla_views( 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 + n_rows = 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) + anchor_bytes + (layer_num - 1) * page_size * row_bytes + (n_rows * 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"build_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" @@ -168,7 +164,7 @@ def build_dense_mla_views( views.append( torch.as_strided( as_dtype_view, - size=(n_dense, 1, kv_cache_dim), + size=(n_rows, 1, kv_cache_dim), stride=(kv_cache_dim, kv_cache_dim, 1), storage_offset=base_bytes // itemsize, ) diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 03df6688b..abc2ae2dd 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -45,7 +45,6 @@ from sglang.kernels.ops.attention.dsa.quant_k_cache import ( from sglang.kernels.ops.kvcache.cache_move import ( copy_all_layer_kv_cache_func, set_kv_buffer_prefix_valid_tiled, - store_cache_4d, ) 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 @@ -62,9 +61,7 @@ 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, - build_page_major_mha_views, mamba_entry_bytes, - mha_entry_bytes, ) from sglang.srt.mem_cache.utils import ( get_mla_kv_buffer_triton, @@ -2866,9 +2863,8 @@ class MHATokenToKVPool(KVCache): 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). The 3-D - # per-layer buffers here ignore page_size in move_kv_cache_native. + # 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: @@ -3184,18 +3180,11 @@ class MHATokenToKVPoolFP4(MHATokenToKVPool): class PageMajorMHATokenToKVPool(MHATokenToKVPool): - """MHA pool with the page-major (layer-major within a page) page-granularity envelope layout. + """MHA pool with the page-major page-granularity envelope layout. - All layers/slots share one contiguous ``uint8`` ``_raw`` buffer; per-layer K/V - are 4-D strided views ``(num_pages, page_size, head_num, head_dim*)`` built by - ``mem_cache/layout/page_major.py``. Token id ``t`` -> page ``t // page_size``, - slot ``t % page_size``; the reserved padding slot 0 lives in page 0. At - ``page_size == 1`` a page is a single slot (token-granularity envelope). - - Supported: the standard CUDA Triton attention + native move path. The tiled KV - copy kernel, CPU offloading, and the spec-decode prefix-commit kernel all assume - the per-layer contiguous 3-D layout; here they fail loudly rather than silently - mis-indexing the strided views. + 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__( @@ -3221,89 +3210,12 @@ class PageMajorMHATokenToKVPool(MHATokenToKVPool): ) def _create_buffers(self): - # One contiguous byte buffer holds all layers/slots; per-layer K/V are - # 4-D strided views in the page-granularity envelope layout (see - # mem_cache/layout/page_major.py). - total_slots = self.size + self.page_size - assert total_slots % self.page_size == 0, ( - f"page_major_layer_major needs (size + page_size) divisible by " - f"page_size; got size={self.size}, page_size={self.page_size}" - ) - num_pages = total_slots // self.page_size - entry_bytes = mha_entry_bytes( - layer_num=self.layer_num, - head_num=self.head_num, - head_dim=self.head_dim, - v_head_dim=self.v_head_dim, - itemsize=self.store_dtype.itemsize, - ) - total_bytes = num_pages * self.page_size * entry_bytes - 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() - ): - # Unset slots read as zeros (matches the per-layer pool). - self._raw = torch.zeros( - total_bytes, dtype=torch.uint8, device=self.device - ) - self.k_buffer, self.v_buffer = build_page_major_mha_views( - self._raw, - layer_num=self.layer_num, - head_num=self.head_num, - head_dim=self.head_dim, - v_head_dim=self.v_head_dim, - store_dtype=self.store_dtype, - page_size=self.page_size, - num_pages=num_pages, - ) - # stride(0) * itemsize is the per-page byte stride; for these strided - # views np.prod(shape[1:]) would not equal it, so compute it directly. - 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( - [x.stride(0) * x.dtype.itemsize for x in (self.k_buffer + self.v_buffer)], - device=self.device, - ) - - def _store_kv_layer( - self, - layer_idx: int, - loc: torch.Tensor, - cache_k: torch.Tensor, - cache_v: torch.Tensor, - ): - # Single-launch Triton write into the 4-D envelope view. The parent's - # view(-1, row_dim) path can't merge the strided 4-D dims. - store_cache_4d( - self.k_buffer[layer_idx], - self.v_buffer[layer_idx], - cache_k, - cache_v, - loc, - page_size=self.page_size, - ) - - def _move_kv_cache_impl(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): - # Strided 4-D views: the tiled copy kernel assumes stride == row bytes, so - # always take the native move (it splits token ids into - # (page_id, slot_in_page) for the 4-D advanced index). - move_kv_cache_native( - self.k_buffer, - self.v_buffer, - tgt_loc, - src_loc, - page_size=self.page_size, + 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 @@ -3740,9 +3652,9 @@ 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 + # virtual->kernel-facing 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. + # identity for a static pool, `translate_kv_loc_for_kernel` for the unified pool. self._full_translate = lambda ids: ids self.use_mla = use_mla if full_kv_pool is not None: @@ -3985,7 +3897,7 @@ class HybridLinearKVPool(KVCache): ) else: # Mirror the MHA branch: `full_loc` is the unified pool's - # pre-translated (dense) loc; None for a static pool. + # 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( @@ -4029,16 +3941,16 @@ class HybridLinearKVPool(KVCache): loc: torch.Tensor, cache_k_nope: torch.Tensor, cache_k_rope: torch.Tensor, - loc_is_dense: bool = False, + loc_is_kernel_facing: bool = False, ): 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, so translate to the dense id space here. + # pool, so translate to the kernel-facing id space here. # - # `loc_is_dense`: the caller already translated `loc` (the unified-pool + # `loc_is_kernel_facing`: the caller already translated `loc` (the unified-pool # cuda-graph decode precomputes it out-of-graph into a capture-stable # buffer, so the in-graph write does not capture a translate allocation). - if not loc_is_dense: + if not loc_is_kernel_facing: 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) @@ -4631,17 +4543,11 @@ def move_kv_cache_native( v_buffer: List[torch.Tensor], tgt_loc: torch.Tensor, src_loc: torch.Tensor, - page_size: int = 1, ): """Move token-granular K/V rows from ``src_loc`` to ``tgt_loc``. - Supports two buffer shapes: - - - 3-D ``[max_slots, head_num, head_dim]`` (per-layer pool): direct advanced - indexing on dim 0; ``page_size`` is ignored. - - 4-D ``[num_pages, page_size, head_num, head_dim]`` (envelope layout): split - each token id into ``(page_id, slot_in_page)`` and use 2-D advanced - indexing. PyTorch resolves the strided byte address via the view's strides. + 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 @@ -4649,21 +4555,8 @@ def move_kv_cache_native( 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): - if k_cache.ndim == 4: - if page_size == 1: - # Degenerate (num_pages, 1, head, dim): token id == page id. - k_cache[tgt_loc_flat, 0] = k_cache[src_loc_flat, 0] - v_cache[tgt_loc_flat, 0] = v_cache[src_loc_flat, 0] - else: - tgt_page = tgt_loc_flat // page_size - tgt_tok = tgt_loc_flat % page_size - src_page = src_loc_flat // page_size - src_tok = src_loc_flat % page_size - k_cache[tgt_page, tgt_tok] = k_cache[src_page, src_tok] - v_cache[tgt_page, tgt_tok] = v_cache[src_page, src_tok] - else: - k_cache[tgt_loc_flat] = k_cache[src_loc_flat] - v_cache[tgt_loc_flat] = v_cache[src_loc_flat] + k_cache[tgt_loc_flat] = k_cache[src_loc_flat] + v_cache[tgt_loc_flat] = v_cache[src_loc_flat] @triton.jit diff --git a/python/sglang/srt/mem_cache/multi_ended_allocator.py b/python/sglang/srt/mem_cache/multi_ended_allocator.py index 154d1b6bf..a0674e757 100644 --- a/python/sglang/srt/mem_cache/multi_ended_allocator.py +++ b/python/sglang/srt/mem_cache/multi_ended_allocator.py @@ -114,7 +114,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): need_sort: bool = False, forward_stream: Optional[torch.cuda.Stream] = None, lazy_compaction: bool = False, - kernel_page_multiplier: int = 1, + kernel_page_multiplier: Optional[int] = None, ): spec = unified_buffer.spec(sub_pool_name) max_slots = unified_buffer.max_slots(sub_pool_name) @@ -134,11 +134,14 @@ 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 + # Kernel-facing page-stride scale, from the spec that owns the layout. + # `kernel_page_multiplier=` overrides it only for tests pinning the + # multiplier-1 collapse. + self.kernel_page_multiplier = ( + spec.blocks_per_page() + if kernel_page_multiplier is None + else kernel_page_multiplier + ) # Zero page envelopes on hand-out — see _maybe_zero_pages. self._zero_pages_on_alloc = isinstance(kvcache, UnifiedMLATokenToKVPool) # Overlap mode: `free` drops a wait_stream(forward_stream) barrier so its @@ -641,7 +644,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): return phys_pages def _maybe_zero_pages(self, phys_pages: torch.Tensor) -> None: - """Zero the page ENVELOPES on hand-out (MLA-dense full pool only): + """Zero the page ENVELOPES on hand-out (MLA full pool only): the MLA kernels arithmetically mask the rows beyond seq_len, so never-written page bytes must read as finite values. Runs on the schedule stream, ordered before the consuming forward by the @@ -713,58 +716,48 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): result = phys_pages * self.page_size + offsets return torch.clamp_min(result, 0) - def translate_kv_loc_dense( + def translate_kv_loc_for_kernel( self, virt_tokens: torch.Tensor, *, out: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """Translate virtual token ids to DENSE (kernel-facing) ids. + """Virtual token ids -> 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. + kernel_id(t) = (t // ps) * (ps * kernel_page_multiplier) + t % ps - 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. + Internal machinery (compaction, in-flight write sets) MUST keep using + `translate_kv_loc`: kernel-facing ids are for kernels only. Tombstones (-1) + clamp to kernel-facing id 0, the page-0 sink. int64 out; a consumer whose + kernel ABI wants int32 narrows where it fills that buffer. """ - if self.kernel_page_multiplier == 1: - return self.translate_kv_loc(virt_tokens, out=out) - if out is not None: + ps = self.page_size + stride = ps * self.kernel_page_multiplier + with record_function("MultiEndedAlloc.translate_kv_loc_for_kernel"): + pages = virt_tokens if ps == 1 else virt_tokens // ps + offsets = None if ps == 1 else virt_tokens % ps + if out is None: + phys = self.virtual_to_physical[pages] + ids = phys * stride if offsets is None else phys * stride + offsets + return ids.clamp_(min=0) assert out.dtype == torch.int64, ( - f"translate_kv_loc_dense: out= dtype must be int64 (matches v2p), " + f"translate_kv_loc_for_kernel: 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"translate_kv_loc_for_kernel: 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) + if pages.dtype != torch.int64: + pages = pages.to(torch.int64) + if pages is virt_tokens: + out.copy_(torch.take(self.virtual_to_physical, pages)) + else: + torch.take(self.virtual_to_physical, pages, out=out) + out.mul_(stride) + if offsets is not None: 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) + return out.clamp_(min=0) # -- alloc -- @@ -1722,7 +1715,6 @@ 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__( @@ -1750,7 +1742,6 @@ 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, @@ -1911,26 +1902,25 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): """Page-level virtual->physical table of the full sub-pool. Kernels that build the MLA block table directly from req_to_token (e.g. trtllm_mla, flashmla) gather through this to turn a VIRTUAL page into a physical one, - then scale by `kernel_page_multiplier` to reach the dense per-page block. + then scale by `kernel_page_multiplier` to reach the per-page block. """ return self.full_attn_allocator.virtual_to_physical - def translate_kv_loc_dense( + def translate_kv_loc_for_kernel( 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) + """Full-pool virtual TOKEN ids -> kernel-facing ids.""" + return self.full_attn_allocator.translate_kv_loc_for_kernel(loc, out=out) def translate_kv_indices_for_transfer( self, kv_indices: torch.Tensor ) -> torch.Tensor: """Virtual TOKEN ids -> PHYSICAL token ids for the PD transfer engine. - PHYSICAL, not dense: the transfer registers page ENVELOPES (see + PHYSICAL, not kernel-facing: the transfer registers page ENVELOPES (see `UnifiedMLATokenToKVPool.get_contiguous_buf_infos`). """ return self.full_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64)) @@ -2255,45 +2245,35 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): *, out: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """SWA-layer read path: virtual TOKEN ids -> swa-physical TOKEN ids (int32, - matching the non-shared API). Page math against the swa side's v2p table. - Supports ``out=`` (int32, same shape) for cuda-graph buffer stability. - """ - if out is not None: - assert out.dtype == torch.int32, ( - f"translate_loc_from_full_to_swa: out= dtype must be int32 " - f"(matches SWA Triton kernel contract), got {out.dtype}" - ) - assert out.shape == kv_indices.shape, ( - f"translate_loc_from_full_to_swa: out= shape " - f"{tuple(out.shape)} must match kv_indices shape " - f"{tuple(kv_indices.shape)}" - ) - # Tombstone-safety clamp (mirrors the full-side clamp): tombstoned (-1) - # v2p_swa entries must not reach `swa_k_buffer[-1]` (illegal under replay). - # Clamp to 0 routes them to the reserved padding sink (slot 0). - if self.swa_attn_allocator.page_size == 1: - if out is not None: - # Gather into a transient int64, then cast into out (`out.copy_`). - tmp = torch.index_select( - self.swa_attn_allocator.virtual_to_physical, 0, kv_indices - ) - tmp = torch.clamp_min(tmp, 0) - out.copy_(tmp.to(torch.int32)) - return out - result = self.swa_attn_allocator.virtual_to_physical[kv_indices] - result = torch.clamp_min(result, 0) - return result.to(torch.int32) - ps = self.swa_attn_allocator.page_size - virt_pages = kv_indices // ps - offsets = kv_indices % ps - swa_phys_pages = self.swa_attn_allocator.virtual_to_physical[virt_pages] - result = (swa_phys_pages * ps + offsets).to(torch.int32) - result = torch.clamp_min(result, 0) - if out is not None: - out.copy_(result) - return out - return result + """SWA-layer read path: virtual TOKEN ids -> swa kernel-facing ids.""" + return self.swa_attn_allocator.translate_kv_loc_for_kernel(kv_indices, out=out) + + @property + def kernel_page_multiplier(self) -> int: + return self.full_attn_allocator.kernel_page_multiplier + + @property + def full_v2p_page_table(self) -> torch.Tensor: + """Page-level virtual->physical table of the full sub-pool.""" + return self.full_attn_allocator.virtual_to_physical + + def translate_kv_loc_for_kernel( + self, + loc: torch.Tensor, + *, + out: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Full-pool virtual TOKEN ids -> kernel-facing ids.""" + return self.full_attn_allocator.translate_kv_loc_for_kernel(loc, out=out) + + @property + def swa_kernel_page_multiplier(self) -> int: + return self.swa_attn_allocator.kernel_page_multiplier + + @property + def swa_v2p_page_table(self) -> torch.Tensor: + """Page-level virtual->physical table of the SWA sub-pool.""" + return self.swa_attn_allocator.virtual_to_physical # -- alloc -- diff --git a/python/sglang/srt/mem_cache/unified_memory_pool.py b/python/sglang/srt/mem_cache/unified_memory_pool.py index 194a8a91d..f974b1f96 100644 --- a/python/sglang/srt/mem_cache/unified_memory_pool.py +++ b/python/sglang/srt/mem_cache/unified_memory_pool.py @@ -29,17 +29,15 @@ from dataclasses import dataclass from typing import Dict, List, NamedTuple, Optional, Tuple import torch -import triton from torch.profiler import record_function -from sglang.kernels.ops.kvcache.cache_move import store_cache_4d_kernel from sglang.kernels.ops.kvcache.zero_pages import zero_pages from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE from sglang.srt.environ import envs from sglang.srt.mem_cache.layout.page_major import ( - build_dense_mla_views, + build_mha_views, + build_mla_views, build_page_major_mamba_views, - build_page_major_mha_views, ) from sglang.srt.mem_cache.memory_pool import ( HybridLinearKVPool, @@ -47,7 +45,6 @@ from sglang.srt.mem_cache.memory_pool import ( MambaPool, MHATokenToKVPool, MLATokenToKVPool, - move_kv_cache_native, unwrap_write_loc, ) from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool @@ -96,6 +93,19 @@ class SubPoolSpec(ABC): """Storage dtype (informational). Multi-dtype subclasses return the dominant buffer's.""" raise NotImplementedError + def view_tail_pad_bytes(self, page_size: int) -> int: + """Bytes this sub-pool's views reach PAST its last page envelope.""" + return 0 + + def blocks_per_page(self) -> int: + """Row-blocks one page holds in this sub-pool's kernel-facing id space. + + The page envelope is a uniform array of equally wide row-blocks, so a + kernel-facing id is the physical page scaled by this count (see + `MultiEndedAllocator.translate_kv_loc_for_kernel`). 1 means the kernel-facing ids are the physical ones. + """ + return 1 + @dataclass(frozen=True, kw_only=True) class MHASubPoolSpec(SubPoolSpec): @@ -140,6 +150,13 @@ class MHASubPoolSpec(SubPoolSpec): + page_size * self.k_row_bytes() ) + def view_tail_pad_bytes(self, page_size: int) -> int: + return page_size * self.entry_bytes() + + def blocks_per_page(self) -> int: + """Row-blocks per page in the kernel-facing id space (one K + one V per layer).""" + return 2 * self.layer_num + def get_dtype(self) -> torch.dtype: return self.store_dtype @@ -174,6 +191,14 @@ class MLASubPoolSpec(SubPoolSpec): def entry_bytes(self) -> int: return self.layer_num * self.kv_cache_dim * self.store_dtype.itemsize + def view_tail_pad_bytes(self, page_size: int) -> int: + return page_size * self.entry_bytes() + + def blocks_per_page(self) -> int: + """One latent row per layer, so L blocks per page (MHA has 2L: a K + block and a V block per layer).""" + return self.layer_num + def get_dtype(self) -> torch.dtype: return self.store_dtype @@ -210,13 +235,28 @@ class MambaSubPoolSpec(SubPoolSpec): # --------------------------------------------------------------------------- -# UnifiedKVPool — the byte buffer + the strided per-sub-pool views +# UnifiedKVPool — the byte buffer + the per-sub-pool views # --------------------------------------------------------------------------- +def unified_memory_supported_for_model(model_config, *, use_mla_backend: bool) -> bool: + """Whether this model's KV geometry can back the unified memory pool.""" + return use_mla_backend or not model_config.has_asymmetric_kv + + +def _assert_kernel_id_bound(*, sub_pool_name: str, n_rows: int) -> None: + """Check if kernel-facing ids can flow through int32 read-index buffers.""" + assert n_rows < 2**31, ( + f"sub-pool {sub_pool_name!r}: kernel-facing id space has {n_rows} rows, " + f"exceeding the int32 bound (2^31) that read-index buffers assume. " + "Reduce max_total_num_tokens or the layer count." + ) + + class UnifiedKVPool: """One physical `uint8` byte buffer shared by 2 sub-pools, each exposing - strided per-layer views. Allocators keep byte ranges disjoint; no usage tracking here. + per-layer views over its own byte range (contiguous per layer for KV, + strided for the Mamba state). Allocators keep byte ranges disjoint; no usage tracking here. """ def __init__( @@ -227,7 +267,6 @@ 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, ( @@ -253,13 +292,12 @@ 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 + self.view_tail_pad_bytes = max( + spec.view_tail_pad_bytes(page_size) for spec in sub_pool_specs + ) with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): self._raw = torch.empty( - total_bytes + view_tail_pad_bytes, dtype=torch.uint8, device=device + total_bytes + self.view_tail_pad_bytes, dtype=torch.uint8, device=device ) if envs.SGLANG_DEBUG_POISON_POOL.get(): # Debug: bf16-NaN-fill so NaN-unsafe reads of never-written bytes @@ -275,7 +313,7 @@ class UnifiedKVPool: self._max_slots: Dict[str, int] = {} self._anchor_bytes: Dict[str, int] = {} self._min_slot_index: Dict[str, int] = {} - # MHA: (k_buffer, v_buffer); MLA: [per-layer dense views]; + # MHA: (k_buffer, v_buffer); MLA: [per-layer per-layer 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]] = {} @@ -405,7 +443,12 @@ class UnifiedKVPool: max_slots: int, page_size: int, ) -> Tuple[List[torch.Tensor], List[torch.Tensor]]: - return build_page_major_mha_views( + num_pages = max_slots // page_size + _assert_kernel_id_bound( + sub_pool_name=spec.name, + n_rows=num_pages * spec.blocks_per_page() * page_size, + ) + return build_mha_views( self._raw, layer_num=spec.layer_num, head_num=spec.head_num, @@ -413,7 +456,7 @@ class UnifiedKVPool: v_head_dim=spec.v_head_dim, store_dtype=spec.store_dtype, page_size=page_size, - num_pages=max_slots // page_size, + num_pages=num_pages, anchor_bytes=anchor_bytes, ) @@ -424,13 +467,18 @@ class UnifiedKVPool: max_slots: int, page_size: int, ) -> List[torch.Tensor]: - return build_dense_mla_views( + num_pages = max_slots // page_size + _assert_kernel_id_bound( + sub_pool_name=spec.name, + n_rows=num_pages * spec.blocks_per_page() * page_size, + ) + return build_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, + num_pages=num_pages, anchor_bytes=anchor_bytes, ) @@ -450,10 +498,16 @@ class UnifiedKVPool: class UnifiedMHATokenToKVPool(MHATokenToKVPool): - """MHA KV pool whose `k_buffer`/`v_buffer` are strided views into a `UnifiedKVPool`. + """MHA KV pool whose per-layer `k_buffer`/`v_buffer` are `build_mha_views` + views into a `UnifiedKVPool` (requires uniform K/V rows). - Relocation uses the native move (strided views break the tiled Triton kernel that - assumes stride == row bytes). `set_kv_buffer` gets PHYSICAL slot ids; never translates. + Views are contiguous `(n_rows, head_num, head_dim)`; locs are + + kernel_id(t) = (t // ps) * (ps * 2 * layer_num) + t % ps + + which is layer- and K/V-independent, each view's storage_offset folding in + its block origin (layer l's K at block 2l, V at 2l+1). `move_kv_cache` is + the exception: compaction passes REAL physical token ids. """ def __init__( @@ -467,17 +521,19 @@ class UnifiedMHATokenToKVPool(MHATokenToKVPool): enable_alt_stream: bool = True, ): spec = unified_buffer.mha_spec(sub_pool_name) - k_buffer, v_buffer = unified_buffer.mha_views_for(sub_pool_name) + k_views, v_views = unified_buffer.mha_views_for(sub_pool_name) max_slots = unified_buffer.max_slots(sub_pool_name) self._unified_buffer = unified_buffer self._sub_pool_name = sub_pool_name - self._k_views = k_buffer - self._v_views = v_buffer - self._page_size = page_size + self._k_views = k_views + self._v_views = v_views + self._num_pages = max_slots // page_size + self._page_bytes = page_size * spec.entry_bytes() + view_rows = self._num_pages * spec.blocks_per_page() * page_size super().__init__( - size=max_slots - 1, # -1 for reserved slot 0 + size=view_rows - page_size, page_size=page_size, dtype=spec.store_dtype, head_num=spec.head_num, @@ -489,126 +545,73 @@ class UnifiedMHATokenToKVPool(MHATokenToKVPool): start_layer=start_layer, end_layer=end_layer, enable_alt_stream=enable_alt_stream, - enable_kv_cache_copy=False, # strided views — force native move + enable_kv_cache_copy=False, + kv_cache_layout="page_major", ) def _create_buffers(self): self.k_buffer = self._k_views self.v_buffer = self._v_views - # For external inspectors only; the native move path doesn't consume them. - 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( - [x.stride(0) * x.dtype.itemsize for x in (self.k_buffer + self.v_buffer)], - device=self.device, - ) def _clear_buffers(self): # Lifetime owned by UnifiedKVPool; do not delete the views. pass - def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): - # tgt_loc/src_loc are PHYSICAL slot ids; native move only (strided views). - if tgt_loc.numel() == 0: - return - with record_function("UnifiedMHA.move_kv_cache"): - move_kv_cache_native( - self.k_buffer, - self.v_buffer, - tgt_loc, - src_loc, - page_size=self._page_size, - ) - def get_kv_size_bytes(self): return 0, 0 # UnifiedKVPool logs the total; per-sub-pool would double-count - def set_kv_buffer( - self, - layer, - loc: torch.Tensor, - cache_k: torch.Tensor, - cache_v: torch.Tensor, - k_scale=None, - v_scale=None, - layer_id_override: Optional[int] = None, - dcp_kv_mask: Optional[torch.Tensor] = None, - ): - # Decode context parallel (dcp_kv_mask) unsupported; fail loud. - assert dcp_kv_mask is None, ( - "UnifiedMHATokenToKVPool.set_kv_buffer: decode context parallel " - "(dcp_kv_mask) is not supported with --enable-unified-memory." - ) - # Bypass super().set_kv_buffer: the parent's `k_cache.view(-1, row_dim)` can't - # merge our 4-D layer-major view (stride[0]=page_bytes) at page_size>1. Call - # store_cache_4d_kernel directly. `loc` is PHYSICAL token ids — no v2p translate. - with record_function("UnifiedMHA.set_kv_buffer"): - 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) - - layer_id = ( - layer.layer_id if layer_id_override is None else layer_id_override - ) - self.start_layer - k_view = self.k_buffer[layer_id] - v_view = self.v_buffer[layer_id] - ps = self._page_size - N = loc.numel() - if N == 0: - return - head_num = k_view.shape[2] - head_dim = k_view.shape[3] - v_head_dim = v_view.shape[3] - K_ROW_DIM = head_num * head_dim - V_ROW_DIM = head_num * v_head_dim - BLOCK = 128 - row_dim_max = K_ROW_DIM if K_ROW_DIM > V_ROW_DIM else V_ROW_DIM - store_cache_4d_kernel[(N, triton.cdiv(row_dim_max, BLOCK), 2)]( - k_view, - v_view, - cache_k, - cache_v, - loc, - k_view.stride(0), - k_view.stride(1), - v_view.stride(0), - v_view.stride(1), - cache_k.stride(0), - cache_v.stride(0), - K_ROW_DIM=K_ROW_DIM, - V_ROW_DIM=V_ROW_DIM, - PAGE_SIZE=ps, - BLOCK=BLOCK, - num_warps=4, + def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): + """Relocate slots by whole page envelope. + `tgt_loc`/`src_loc` are REAL physical token ids, not kernel-facing ids. + """ + if tgt_loc.numel() == 0: + return + # The envelope view below starts at byte 0, so this sub-pool must be + # anchored there; a non-zero anchor moves another sub-pool's bytes, and + # the ids stay in range so nothing downstream notices. + assert self._unified_buffer.anchor_bytes(self._sub_pool_name) == 0 + 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("UnifiedMHA.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] + + def get_contiguous_buf_infos(self): + raise NotImplementedError( + "unified layout has no per-layer contiguous regions; " + "KV transfer / disaggregation is unsupported." + ) + + def get_cpu_copy(self, indices, mamba_indices=None): + raise NotImplementedError( + "CPU offloading is unsupported under the unified layout." + ) + + def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None): + raise NotImplementedError( + "CPU offloading is unsupported under the unified layout." + ) + + def set_kv_buffer_prefix_valid(self, *args, **kwargs): + raise NotImplementedError( + "prefix-valid commit is unsupported under the unified layout " + "(_set_kv_buffer_prefix_valid_impl assumes token-id indexing)." + ) class UnifiedMLATokenToKVPool(MLATokenToKVPool): - """MLA KV pool whose per-layer `kv_buffer` entries are DENSE views into a - `UnifiedKVPool` (see `build_dense_mla_views`). + """MLA KV pool whose per-layer `kv_buffer` entries are kernel-facing views into a + `UnifiedKVPool` (see `build_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 + `get_value_buffer`) is a kernel-facing id — the `translate_kv_loc_for_kernel` output - dense(t) = (t // ps) * (ps * layer_num) + t % ps + kernel_id(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 @@ -638,13 +641,12 @@ class UnifiedMLATokenToKVPool(MLATokenToKVPool): 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 + self._view_rows = 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, + # OOB checks bound locs by `size + page_size`; kernel-facing ids run to + # `_view_rows` (page 0 is the reserved padding sink). + size=self._view_rows - page_size, page_size=page_size, dtype=kv_cache_dtype, kv_lora_rank=spec.kv_lora_rank, @@ -669,8 +671,8 @@ class UnifiedMLATokenToKVPool(MLATokenToKVPool): ``raw_ptr + physical_page_id * page_envelope_bytes``. The transfer item is the whole page envelope (all layers of one page) - rather than a per-layer region, because the per-layer dense views - overlap and index in dense ids. Both sides must therefore build the + rather than a per-layer region, because the per-layer per-layer views + overlap and index in kernel-facing ids. Both sides must therefore build the pool with identical specs. """ # The address formula omits the anchor; a nonzero one would mis-address. @@ -681,7 +683,7 @@ class UnifiedMLATokenToKVPool(MLATokenToKVPool): 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 + `tgt_loc`/`src_loc` are REAL physical token ids (NOT kernel-facing 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. @@ -1164,16 +1166,12 @@ 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, @@ -1233,9 +1231,6 @@ 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). @@ -1249,9 +1244,9 @@ def init_unified_mamba_pools( 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 + # receive VIRTUAL locs and translate to the kernel-facing space internally # (eager-prefill-only paths; never captured in a cuda graph). - token_to_kv_pool._full_translate = allocator.translate_kv_loc_dense + token_to_kv_pool._full_translate = allocator.translate_kv_loc_for_kernel logger.info( "[unified-memory-pool] ============================================================" @@ -1263,7 +1258,7 @@ def init_unified_mamba_pools( 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, " + "qk_rope_head_dim=%d, page_size=%d (per-layer views, kernel_page_multiplier=%d, " "view_tail_pad=%d B)", len(full_attention_layer_ids), len(mamba_layer_ids), @@ -1271,18 +1266,20 @@ def init_unified_mamba_pools( qk_rope_head_dim, page_size, len(full_attention_layer_ids), - view_tail_pad_bytes, + shared_pool.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", + "page_size=%d, is_draft_worker=%s (%s)", len(full_attention_layer_ids), len(mamba_layer_ids), head_num, head_dim, page_size, is_draft_worker, + "per-layer views, kernel_page_multiplier=%d, view_tail_pad=%d B" + % (full_spec.blocks_per_page(), shared_pool.view_tail_pad_bytes), ) logger.info( "[unified-memory-pool] total_bytes=%d, max_total_num_tokens=%d, max_mamba_cache_size=%d, " @@ -1418,24 +1415,12 @@ class UnifiedSWAKVPool(SWAKVPool): return # no-op in shared mode (the swa-side v2p IS the mapping) def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor): - """Virtual token ids -> swa-physical token ids (int32).""" + """Virtual token ids -> swa kernel-facing ids (int64).""" assert self._swa_allocator is not None, ( "UnifiedSWAKVPool.translate_loc_from_full_to_swa called before " "attach_allocators" ) - ps = self._swa_allocator.page_size - # Tombstone-safety clamp, matching MultiEndedAllocator.translate_kv_loc: - # a tombstoned v2p entry (-1) must not reach the caller as a negative - # loc. Clamp to 0 routes it to the reserved padding sink instead. - if ps == 1: - swa_locs = self._swa_allocator.virtual_to_physical[kv_indices] - else: - virt_pages = kv_indices // ps - offsets = kv_indices % ps - swa_phys_pages = self._swa_allocator.virtual_to_physical[virt_pages] - # Tombstoned page: -1 * ps + offset lands in [-ps, -1]. - swa_locs = swa_phys_pages * ps + offsets - return swa_locs.clamp(min=0).to(torch.int32) + return self._swa_allocator.translate_kv_loc_for_kernel(kv_indices) def get_state_buf_infos(self): return self.swa_kv_pool.get_contiguous_buf_infos() @@ -1664,6 +1649,15 @@ def init_unified_swa_pools( "[unified-memory-pool] ============================================================" ) logger.info("[unified-memory-pool] UNIFIED MEMORY POOL ENABLED -- path=SWA hybrid") + logger.info( + "[unified-memory-pool] %s", + "per-layer views, kernel_page_multiplier full=%d swa=%d, view_tail_pad=%d B" + % ( + full_spec.blocks_per_page(), + swa_spec.blocks_per_page(), + shared_pool.view_tail_pad_bytes, + ), + ) logger.info( "[unified-memory-pool] full_layers=%d, swa_layers=%d, head_num=%d, head_dim=%d, " "v_head_dim=%d, swa_head_num=%d, swa_head_dim=%d, swa_v_head_dim=%d, " diff --git a/test/registered/page_major/test_page_major_gpt_oss.py b/test/registered/page_major/test_page_major_gpt_oss.py index f64e4fcaa..5b7ae0ebb 100644 --- a/test/registered/page_major/test_page_major_gpt_oss.py +++ b/test/registered/page_major/test_page_major_gpt_oss.py @@ -1,10 +1,10 @@ """ -End-to-end accuracy test for the page-major KV layout on a hybrid-SWA MoE model. +End-to-end accuracy test for the unified memory pool on a hybrid-SWA MoE model. -Launches gpt-oss-20b with ``--enable-page-major-kv-layout`` on the Triton -attention backend and checks that GSM8K accuracy holds. This exercises the -SWA + full-attention KV pools under the page-granularity envelope layout -(SWAKVPool routes both sub-pools through PageMajorMHATokenToKVPool). +Launches gpt-oss-20b with ``--enable-unified-memory`` on the Triton attention +backend and checks that GSM8K accuracy holds. This exercises the SWA + +full-attention KV sub-pools stored as per-layer views in the unified +page-major envelope. Registered to the label-gated ``run-ci-extra`` suite (opt-in, not per-commit). @@ -22,9 +22,17 @@ from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE register_cuda_ci(est_time=420, stage="extra-a", runner_config="1-gpu-large") +_UNIFIED_COMMON_ARGS = [ + "--enable-unified-memory", + "--mem-fraction-static", + "0.70", + "--cuda-graph-backend-prefill=disabled", +] -class TestPageMajorGptOss(DefaultServerBase): - """Page-major KV layout on gpt-oss-20b (hybrid-SWA MoE), Triton backend.""" + +class TestUnifiedGptOssTriton(DefaultServerBase): + """Unified pool on gpt-oss-20b (hybrid-SWA MoE), Triton pinned: dense + MHA/SWA views through the reference backend.""" model = DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE @@ -33,16 +41,7 @@ class TestPageMajorGptOss(DefaultServerBase): num_shots = 5 parallel = 32 - other_args = [ - "--enable-page-major-kv-layout", - # The envelope's strided 4-D K/V views are only read by the Triton - # attention kernels (the layout's validator enforces this). - "--attention-backend", - "triton", - "--mem-fraction-static", - "0.70", - "--cuda-graph-backend-prefill=disabled", - ] + other_args = _UNIFIED_COMMON_ARGS + ["--attention-backend", "triton"] def test_gsm8k(self): from sglang.test.few_shot_gsm8k import run_eval as run_few_shot_gsm8k diff --git a/test/registered/page_major/test_page_major_qwen_hybrid.py b/test/registered/page_major/test_page_major_qwen_hybrid.py index 61d55540f..1bd98f82e 100644 --- a/test/registered/page_major/test_page_major_qwen_hybrid.py +++ b/test/registered/page_major/test_page_major_qwen_hybrid.py @@ -1,12 +1,12 @@ """ -End-to-end accuracy test for the page-major KV layout on a GDN-hybrid model. +End-to-end accuracy test for the unified memory pool on a GDN-hybrid model. Launches Qwen3.5-4B (a gated-delta-net / linear-attention hybrid) with -``--enable-page-major-kv-layout`` on the Triton attention + linear-attn + Mamba -backends and checks that GSM8K accuracy holds. This exercises the page-major -path most prone to subtle bugs: the Mamba conv/SSM state stored as a strided -envelope view, plus the full-attention KV pool, both read/written by the GDN -prefill and decode kernels. +``--enable-unified-memory`` on the Triton attention + linear-attn + Mamba +backends and checks that GSM8K accuracy holds. This exercises the unified +envelope's most bug-prone path: the Mamba conv/SSM state stored as a strided +envelope view, plus the full-attention KV stored as per-layer views, +both read/written by the GDN prefill and decode kernels. Registered to the label-gated ``run-ci-extra`` suite (opt-in, not per-commit). @@ -24,35 +24,34 @@ from sglang.test.test_utils import DEFAULT_HYBRID_GDN_SMALL_MODEL_NAME_FOR_TEST register_cuda_ci(est_time=300, stage="extra-a", runner_config="1-gpu-large") +_UNIFIED_COMMON_ARGS = [ + "--trust-remote-code", + "--mem-fraction-static", + "0.85", + "--enable-unified-memory", + "--linear-attn-backend", + "triton", + "--mamba-backend", + "triton", +] -class TestPageMajorQwenHybrid(DefaultServerBase): - """Page-major KV layout on Qwen3.5-4B (GDN-hybrid), Triton backends.""" + +class TestUnifiedQwenHybridTriton(DefaultServerBase): + """Unified pool on Qwen3.5-4B (GDN-hybrid), Triton pinned: dense + full-attention views + strided conv/SSM state through the reference + backends.""" model = DEFAULT_HYBRID_GDN_SMALL_MODEL_NAME_FOR_TEST - # Measured in this harness: baseline (no page-major) and page-major both - # ~0.86; the 0.80 threshold leaves margin for run-to-run noise while still - # catching the prefill-state corruption that page-major hit before the - # gather/scatter fix in gdn_backend.forward_extend (which dropped it to ~0.61). + # Measured ~0.86 in this harness on both the static pools and the envelope + # layout; 0.80 leaves noise margin and still catches a corrupted prefill + # state, which reads ~0.61. gsm8k_threshold = 0.80 num_gsm8k_questions = 200 num_shots = 5 parallel = 32 - other_args = [ - "--trust-remote-code", - "--mem-fraction-static", - "0.85", - "--enable-page-major-kv-layout", - # Only the Triton attention / linear-attn / Mamba kernels read the - # strided envelope K/V and conv/SSM state (enforced by the validator). - "--attention-backend", - "triton", - "--linear-attn-backend", - "triton", - "--mamba-backend", - "triton", - ] + other_args = _UNIFIED_COMMON_ARGS + ["--attention-backend", "triton"] def test_gsm8k(self): from sglang.test.few_shot_gsm8k import run_eval as run_few_shot_gsm8k diff --git a/test/registered/unit/mem_cache/test_full_loc_fast_path.py b/test/registered/unit/mem_cache/test_full_loc_fast_path.py index 5d1c902a1..208c213ac 100644 --- a/test/registered/unit/mem_cache/test_full_loc_fast_path.py +++ b/test/registered/unit/mem_cache/test_full_loc_fast_path.py @@ -23,7 +23,7 @@ always PHYSICAL. Two routing contracts are pinned here: `UnifiedSWAKVPool` asserts it's present (the unified memory pool always precomputes it); `HybridLinearKVPool` falls back to `loc` for a static (non-shared) pool, where `loc` is itself already physical. -2. SWA. The swa-physical loc rides the backend `swa_out_cache_loc` rail +2. SWA. The swa-physical loc rides the backend `swa_out_cache_loc` slot (`KVWriteLoc.swa_loc`) and is written directly. Pure dispatch tests: the inner sub-pools are recording stubs, so no GPU / real @@ -127,7 +127,7 @@ class TestUnifiedSWARouting(unittest.TestCase): self.assertEqual(len(pool.swa_kv_pool.calls), 1) forwarded, kwargs = pool.swa_kv_pool.calls[0] - # SWA write rides the backend rail: forward the swa-physical loc directly. + # SWA write rides the backend slot: forward the swa-physical loc directly. self.assertIs(forwarded, swa_phys) self.assertNotIn("already_physical", kwargs) # Full pool untouched for an SWA layer. @@ -138,7 +138,7 @@ class TestUnifiedSWARouting(unittest.TestCase): virtual_loc = torch.tensor([10, 11, 12], dtype=torch.int64) layer = types.SimpleNamespace(layer_id=1) # SWA layer - # No swa_loc bundled -> the rail contract is violated; must assert + # No swa_loc bundled -> the write-loc contract is violated; must assert # rather than silently writing wrong (un-translated) locations. with self.assertRaises(AssertionError): pool.set_kv_buffer( @@ -149,6 +149,46 @@ class TestUnifiedSWARouting(unittest.TestCase): ) +class TestUnifiedSWATombstoneClamp(unittest.TestCase): + """`UnifiedSWAKVPool.translate_loc_from_full_to_swa` must clamp tombstoned + ids to the reserved padding sink (0). + + A token whose swa page was freed carries -1 in `virtual_to_physical`. Before + the clamp, that produced a negative id, which a captured graph stores at a + negative offset from the buffer base. The composite allocator's method of + the same name already clamped; this path did not. + """ + + def _make_bare_pool(self, page_size, v2p, multiplier=1): + from sglang.srt.mem_cache.multi_ended_allocator import MultiEndedAllocator + from sglang.srt.mem_cache.unified_memory_pool import UnifiedSWAKVPool + + # A real sub-allocator (not a stand-in): the translation reads its v2p + # table, and the pool reaches it through the allocator's own method. + swa_allocator = object.__new__(MultiEndedAllocator) + swa_allocator.page_size = page_size + swa_allocator.virtual_to_physical = v2p + swa_allocator.kernel_page_multiplier = multiplier + pool = object.__new__(UnifiedSWAKVPool) + pool._swa_allocator = swa_allocator + return pool + + def test_tombstoned_id_lands_on_sink(self): + for ps, mult in ((1, 1), (4, 1), (4, 6)): + v2p = torch.tensor([0, -1, 2], dtype=torch.int64) + pool = self._make_bare_pool(ps, v2p, multiplier=mult) + # Virtual ids covering the tombstoned page (index 1) and a live one. + kv_indices = torch.tensor([0, ps, 2 * ps], dtype=torch.int64) + out = pool.translate_loc_from_full_to_swa(kv_indices) + self.assertEqual(out.dtype, torch.int64) + self.assertTrue( + bool((out >= 0).all().item()), + f"tombstoned swa id stayed negative at page_size={ps}, " + f"multiplier={mult}: {out}", + ) + self.assertEqual(int(out[1].item()), 0) + + class TestHybridLinearFullLocRouting(unittest.TestCase): """`HybridLinearKVPool.set_kv_buffer` (non-MLA) writes the full-physical `full_loc` from the write metadata when present (unified memory pool), else the diff --git a/test/registered/unit/mem_cache/test_layout_compat.py b/test/registered/unit/mem_cache/test_layout_compat.py index 41ddbbe1e..93b002397 100644 --- a/test/registered/unit/mem_cache/test_layout_compat.py +++ b/test/registered/unit/mem_cache/test_layout_compat.py @@ -11,18 +11,18 @@ # See the License for the specific language governing permissions and # limitations under the License. # ============================================================================== -"""Unit tests for the page-major layer-major byte layout. +"""Unit tests for the page-major envelope byte layout. + +The subject here is the ENVELOPE — the byte layout the unified pool stores its +KV in — pinned through ``MHASubPoolSpec``'s offset math. The dense 3-D views +the pool exposes over the same bytes are covered by +``test_unified_mha_views.py``, which also pins the view addressing +against the envelope formula byte for byte. Verifies that: -1. The new 4-D ``_build_mha_views`` output exposes correct byte addresses - for each (layer, page, tok_in_page, head, dim) — under both the - degenerate ``page_size=1`` case (byte-identical to the old per-token - envelope) and the new ``page_size>1`` layer-major case. -2. ``MHASubPoolSpec.layer_k_offset_in_page`` / - ``layer_v_offset_in_page`` math matches the layout intent. -3. ``set_kv_buffer`` round-trips correctly for both page sizes. -4. Compaction (``move_kv_cache_native``) moves the right bytes for both - page sizes via the 4-D advanced indexing path. +1. ``MHASubPoolSpec.layer_k_offset_in_page`` / ``layer_v_offset_in_page`` math + matches the layout intent at ``page_size == 1`` and ``> 1``. +2. ``move_kv_cache_native`` (the stock per-layer 3-D move) stays byte-exact. CPU-only — no GPU / Triton needed. @@ -38,11 +38,7 @@ import unittest import torch from sglang.srt.mem_cache.memory_pool import move_kv_cache_native -from sglang.srt.mem_cache.unified_memory_pool import ( - MambaSubPoolSpec, - MHASubPoolSpec, - UnifiedKVPool, -) +from sglang.srt.mem_cache.unified_memory_pool import MHASubPoolSpec _DEV = "cpu" @@ -58,18 +54,6 @@ def _make_mha_spec(name, grow, layer_num=2, head_num=2, head_dim=4): ) -def _make_mamba_spec(name, grow, layer_num=2): - return MambaSubPoolSpec( - name=name, - layer_num=layer_num, - conv_state_shapes=((4, 3),), - conv_dtype=torch.float32, - temporal_state_shape=(2, 2, 2), - temporal_dtype=torch.float32, - grow_direction=grow, - ) - - class TestMHASpecLayerOffsets(unittest.TestCase): """Verify ``layer_k_offset_in_page`` / ``layer_v_offset_in_page`` math.""" @@ -113,175 +97,10 @@ class TestMHASpecLayerOffsets(unittest.TestCase): self.assertEqual(spec.page_bytes(ps), ps * spec.entry_bytes()) -class TestBuildMHAViews(unittest.TestCase): - """Verify the 4-D view shape + strides at both page sizes.""" - - def _build(self, page_size, layer_num=3, head_num=2, head_dim=4, n_full_slots=64): - full = _make_mha_spec( - "full", "up", layer_num=layer_num, head_num=head_num, head_dim=head_dim - ) - swa = _make_mha_spec( - "swa", "down", layer_num=2, head_num=head_num, head_dim=head_dim - ) - # Pad to ensure max_slots % page_size == 0 in both sub-pools. - # entry_bytes is fixed per spec; size accordingly. - total = full.entry_bytes() * n_full_slots + swa.entry_bytes() * n_full_slots - pool = UnifiedKVPool( - total_bytes=total, - sub_pool_specs=[full, swa], - device=_DEV, - enable_memory_saver=False, - page_size=page_size, - ) - return pool, full - - def test_view_shape_is_4d(self): - for ps in [1, 8]: - pool, spec = self._build(page_size=ps) - k_views, v_views = pool.mha_views_for("full") - self.assertEqual(len(k_views), spec.layer_num) - max_slots = pool.max_slots("full") - for L in range(spec.layer_num): - self.assertEqual(k_views[L].ndim, 4) - self.assertEqual( - tuple(k_views[L].shape), - (max_slots // ps, ps, spec.head_num, spec.head_dim), - ) - self.assertEqual( - tuple(v_views[L].shape), - (max_slots // ps, ps, spec.head_num, spec.v_head_dim), - ) - - def test_strides_at_page_size_1_match_envelope(self): - """At ps=1, the 4-D view's stride[0] equals what today's 3-D view's - stride[0] would have been (= entry_bytes / itemsize).""" - pool, spec = self._build(page_size=1, layer_num=4, head_num=3, head_dim=8) - k_views, _ = pool.mha_views_for("full") - itemsize = spec.store_dtype.itemsize - for L in range(spec.layer_num): - # stride[0] = page_bytes/itemsize = entry_bytes/itemsize at ps=1 - self.assertEqual(k_views[L].stride(0), spec.entry_bytes() // itemsize) - # stride[1] = k_row/itemsize (within-page token stride) - self.assertEqual(k_views[L].stride(1), spec.k_row_bytes() // itemsize) - # stride[2] = head_dim (head stride) - self.assertEqual(k_views[L].stride(2), spec.head_dim) - # stride[3] = 1 (innermost) - self.assertEqual(k_views[L].stride(3), 1) - - def test_strides_at_page_size_gt_1(self): - pool, spec = self._build(page_size=8, layer_num=4, head_num=3, head_dim=8) - k_views, _ = pool.mha_views_for("full") - itemsize = spec.store_dtype.itemsize - for L in range(spec.layer_num): - # page_bytes = 8 * 4 * (k_row + v_row); stride[0] = that / itemsize - self.assertEqual(k_views[L].stride(0), spec.page_bytes(8) // itemsize) - # token stride within layer L's K block = k_row/itemsize - self.assertEqual(k_views[L].stride(1), spec.k_row_bytes() // itemsize) - self.assertEqual(k_views[L].stride(2), spec.head_dim) - self.assertEqual(k_views[L].stride(3), 1) - - def test_distinct_layers_dont_alias_at_page_size_gt_1(self): - """Writes to layer 0 must not affect layer 1's K/V values (under - layer-major within-page layout).""" - pool, spec = self._build(page_size=8, layer_num=3, head_num=2, head_dim=4) - k_views, v_views = pool.mha_views_for("full") - # Set page 0, token 3, layer 0 K to a distinct pattern. - target_val = 0.5 - k_views[0][0, 3] = target_val - # Layer 1 K at the same (page, tok) should remain at default (0.0). - self.assertFalse(torch.all(k_views[1][0, 3] == target_val)) - self.assertTrue(torch.all(k_views[1][0, 3] == 0.0)) - # And layer 0 V at the same (page, tok) should remain at default. - self.assertFalse(torch.all(v_views[0][0, 3] == target_val)) - self.assertTrue(torch.all(v_views[0][0, 3] == 0.0)) - - def test_distinct_pages_dont_alias_at_page_size_gt_1(self): - """Writes to one page must not affect another page.""" - pool, spec = self._build(page_size=8, layer_num=3, head_num=2, head_dim=4) - k_views, _ = pool.mha_views_for("full") - # Set page 0, token 3, layer 0 K to a distinct pattern. - k_views[0][0, 3] = 1.25 - # Page 1, token 3, layer 0 K should remain at default. - self.assertTrue(torch.all(k_views[0][1, 3] == 0.0)) - - -class TestMoveKVCacheNative4D(unittest.TestCase): - """Verify ``move_kv_cache_native`` handles 4-D buffers at both - page_size=1 (degenerate envelope) and page_size>1 (layer-major).""" - - def _build_buffer( - self, page_size, layer_num=2, head_num=2, head_dim=4, n_full_slots=64 - ): - full = _make_mha_spec( - "full", "up", layer_num=layer_num, head_num=head_num, head_dim=head_dim - ) - swa = _make_mha_spec( - "swa", "down", layer_num=2, head_num=head_num, head_dim=head_dim - ) - total = full.entry_bytes() * n_full_slots + swa.entry_bytes() * n_full_slots - pool = UnifiedKVPool( - total_bytes=total, - sub_pool_specs=[full, swa], - device=_DEV, - enable_memory_saver=False, - page_size=page_size, - ) - return pool - - def test_move_kv_cache_page_size_1(self): - pool = self._build_buffer(page_size=1, layer_num=2, head_num=2, head_dim=4) - k_views, v_views = pool.mha_views_for("full") - # Write distinct markers at source slots 5, 6. - for L in range(2): - k_views[L][5, 0] = float(L + 1) - v_views[L][5, 0] = -float(L + 1) - k_views[L][6, 0] = float(L + 10) - v_views[L][6, 0] = -float(L + 10) - # Move 5 -> 8 and 6 -> 9. - move_kv_cache_native( - k_views, - v_views, - tgt_loc=torch.tensor([8, 9], dtype=torch.int64), - src_loc=torch.tensor([5, 6], dtype=torch.int64), - page_size=1, - ) - for L in range(2): - self.assertTrue(torch.all(k_views[L][8, 0] == float(L + 1))) - self.assertTrue(torch.all(v_views[L][8, 0] == -float(L + 1))) - self.assertTrue(torch.all(k_views[L][9, 0] == float(L + 10))) - self.assertTrue(torch.all(v_views[L][9, 0] == -float(L + 10))) - - def test_move_kv_cache_page_size_gt_1(self): - ps = 8 - pool = self._build_buffer(page_size=ps, layer_num=2, head_num=2, head_dim=4) - k_views, v_views = pool.mha_views_for("full") - # Write markers at token ids 5 and 14 (different pages). - for L in range(2): - # token 5 = (page 0, tok 5) - k_views[L][0, 5] = float(L + 1) - v_views[L][0, 5] = -float(L + 1) - # token 14 = (page 1, tok 6) - k_views[L][1, 6] = float(L + 10) - v_views[L][1, 6] = -float(L + 10) - # Move token 5 -> token 23 (page 2, tok 7) and 14 -> 31 (page 3, tok 7). - move_kv_cache_native( - k_views, - v_views, - tgt_loc=torch.tensor([23, 31], dtype=torch.int64), - src_loc=torch.tensor([5, 14], dtype=torch.int64), - page_size=ps, - ) - for L in range(2): - # 23 = page 2, tok 7 - self.assertTrue(torch.all(k_views[L][2, 7] == float(L + 1))) - self.assertTrue(torch.all(v_views[L][2, 7] == -float(L + 1))) - # 31 = page 3, tok 7 - self.assertTrue(torch.all(k_views[L][3, 7] == float(L + 10))) - self.assertTrue(torch.all(v_views[L][3, 7] == -float(L + 10))) - - def test_move_kv_cache_3d_legacy_path_unchanged(self): - """move_kv_cache_native(3-D, page_size=1) must take the legacy - else-branch and be byte-identical to today.""" +class TestMoveKVCacheNative(unittest.TestCase): + def test_move_kv_cache_3d_path_unchanged(self): + """The stock per-layer 3-D move must relocate exactly the named token + rows, byte-identically — compaction on static pools rides on it.""" k = [torch.zeros((32, 2, 4), dtype=torch.float16) for _ in range(2)] v = [torch.zeros((32, 2, 4), dtype=torch.float16) for _ in range(2)] for L in range(2): @@ -292,64 +111,11 @@ class TestMoveKVCacheNative4D(unittest.TestCase): v, tgt_loc=torch.tensor([7], dtype=torch.int64), src_loc=torch.tensor([5], dtype=torch.int64), - page_size=1, ) for L in range(2): self.assertTrue(torch.all(k[L][7] == float(L + 1))) self.assertTrue(torch.all(v[L][7] == -float(L + 1))) -class TestByteIdentityAtPageSize1(unittest.TestCase): - """Verify that at page_size=1 the new 4-D view describes the SAME - physical bytes as the old 3-D view would have. The view - semantics differ (4-D vs 3-D shape) but the underlying byte layout is - identical — confirmed by manually computing expected byte offsets and - matching them against the 4-D view's strides + storage_offset. - """ - - def test_byte_addresses_match_envelope(self): - spec = _make_mha_spec("full", "up", layer_num=4, head_num=2, head_dim=4) - ps = 1 - # Build pool. - total = spec.entry_bytes() * 64 + spec.entry_bytes() * 32 - pool = UnifiedKVPool( - total_bytes=total, - sub_pool_specs=[ - spec, - _make_mha_spec("swa", "down", layer_num=2), - ], - device=_DEV, - enable_memory_saver=False, - page_size=ps, - ) - k_views, v_views = pool.mha_views_for("full") - # For each (layer, slot), compute the expected byte address under - # the envelope layout and verify the 4-D view's data_ptr + - # advanced indexing agrees. - max_slots = pool.max_slots("full") - itemsize = spec.store_dtype.itemsize - base_addr = pool._raw.data_ptr() - for L in range(spec.layer_num): - for s in range(0, max_slots, max(1, max_slots // 4)): - # Envelope: bytes for slot s, layer L's K start at: - # s * entry_bytes + L * (k_row + v_row) - expected_k_byte_offset = s * spec.entry_bytes() + L * ( - spec.k_row_bytes() + spec.v_row_bytes() - ) - # 4-D view: k_views[L][page=s, tok=0, head=0, dim=0] - # storage_offset of the element [s, 0, 0, 0]: - view_offset_elems = ( - k_views[L].storage_offset() - + s * k_views[L].stride(0) - + 0 * k_views[L].stride(1) - + 0 * k_views[L].stride(2) - + 0 * k_views[L].stride(3) - ) - view_byte_offset = view_offset_elems * itemsize - # 4-D view sits over `_raw.view(spec.store_dtype)`, which - # has data_ptr == _raw.data_ptr() (same backing storage). - self.assertEqual(view_byte_offset, expected_k_byte_offset) - - if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_multi_ended_allocator.py b/test/registered/unit/mem_cache/test_multi_ended_allocator.py index 06fa88555..3329c0de9 100644 --- a/test/registered/unit/mem_cache/test_multi_ended_allocator.py +++ b/test/registered/unit/mem_cache/test_multi_ended_allocator.py @@ -438,8 +438,7 @@ class TestMultiEndedAllocator(unittest.TestCase): def test_translate_kv_loc_dtype_assertion(self): """REGRESSION: wrong-dtype `out=` (int32 instead of int64) raises - AssertionError. Guards against the copy/paste hazard where someone - might allocate the full-physical buffer with the SWA int32 pattern.""" + AssertionError -- `out=` must match the v2p dtype the gather writes.""" _, full_alloc, _, full_kv, _ = self._build_pair() v = self._alloc(full_alloc, full_kv, 5) wrong_dtype = torch.empty(v.shape, dtype=torch.int32, device=_DEV) @@ -911,28 +910,28 @@ class TestUnifiedSWATokenToKVPoolAllocator(unittest.TestCase): def test_swa_translate_loc_from_full_to_swa_with_out_writes_inplace(self): """REGRESSION: `translate_loc_from_full_to_swa(v, out=buf)` must modify `buf` in place AND preserve `buf.data_ptr()`. `out=` - buffer MUST be int32 (matches SWA Triton kernel contract).""" + buffer is int64 — every id the allocator emits is.""" _, allocator, _ = self._build() v = allocator.alloc(4) self.assertIsNotNone(v) - buf = torch.empty(v.shape, dtype=torch.int32, device=_DEV) + buf = torch.empty(v.shape, dtype=torch.int64, device=_DEV) ptr_before = buf.data_ptr() ret = allocator.translate_loc_from_full_to_swa(v, out=buf) self.assertIs(ret, buf) self.assertEqual(buf.data_ptr(), ptr_before) # Byte-identical to the no-out form: no_out = allocator.translate_loc_from_full_to_swa(v) - self.assertEqual(no_out.dtype, torch.int32) + self.assertEqual(no_out.dtype, torch.int64) self.assertTrue(bool((buf == no_out).all().item())) def test_swa_translate_loc_from_full_to_swa_dtype_assertion(self): - """REGRESSION: wrong-dtype `out=` (int64 instead of int32) - raises AssertionError. Guards against accidentally reusing the int64 - full-physical buffer pattern for the SWA precompute.""" + """REGRESSION: wrong-dtype `out=` (int32 instead of int64) raises + AssertionError. Guards against reintroducing a narrowed SWA write loc: the + allocator emits int64 and consumers narrow at their own buffer.""" _, allocator, _ = self._build() v = allocator.alloc(4) self.assertIsNotNone(v) - wrong_dtype = torch.empty(v.shape, dtype=torch.int64, device=_DEV) + wrong_dtype = torch.empty(v.shape, dtype=torch.int32, device=_DEV) with self.assertRaises(AssertionError): allocator.translate_loc_from_full_to_swa(v, out=wrong_dtype) @@ -947,21 +946,48 @@ class TestUnifiedSWATokenToKVPoolAllocator(unittest.TestCase): # Inject a tombstone on the swa side at one of the live virtual ids. v_tomb = int(v[1].item()) allocator.swa_attn_allocator.virtual_to_physical[v_tomb] = -1 - # No-out form: result must be int32 AND every entry >= 0. + # No-out form: result must be int64 AND every entry >= 0. out = allocator.translate_loc_from_full_to_swa(v) - self.assertEqual(out.dtype, torch.int32) + self.assertEqual(out.dtype, torch.int64) self.assertTrue( bool((out >= 0).all().item()), "translate_loc_from_full_to_swa must clamp tombstoned to >=0", ) self.assertEqual(int(out[1].item()), 0) - # out= form (int32 buffer) must also clamp. - buf = torch.empty(v.shape, dtype=torch.int32, device=_DEV) + # out= form must also clamp. + buf = torch.empty(v.shape, dtype=torch.int64, device=_DEV) ret = allocator.translate_loc_from_full_to_swa(v, out=buf) self.assertIs(ret, buf) self.assertTrue(bool((buf >= 0).all().item())) self.assertEqual(int(buf[1].item()), 0) + def test_swa_slot_zero_sink_invariant_survives_churn(self): + """PINNED INVARIANT (swa side of the physical-loc contract): BOTH maps + send virtual 0 to physical 0 — `translate_kv_loc(zeros) == zeros` AND + `translate_loc_from_full_to_swa(zeros) == zeros` — after init and + after alloc/free/free_swa churn. Cuda-graph capture replaced the + capture-time translate with zero-fill/copy of the zero-filled static + buffers; that is only equivalent while slot 0 stays the sink in both + sub-pools.""" + _, allocator, kvcache = self._build() + zeros64 = torch.zeros(4, dtype=torch.int64) + + def check(): + self.assertTrue(torch.equal(allocator.translate_kv_loc(zeros64), zeros64)) + self.assertTrue( + torch.equal(allocator.translate_loc_from_full_to_swa(zeros64), zeros64) + ) + + check() + a = self._alloc(allocator, kvcache, 5) + b = self._alloc(allocator, kvcache, 5) + allocator.free_swa(a) # tombstone swa side only + self._free(allocator, kvcache, b) # full free (compaction on both) + self._free(allocator, kvcache, a) + c = self._alloc(allocator, kvcache, 3) + self._free(allocator, kvcache, c) + check() + # --------------------------------------------------------------------------- # page_size > 1 — paged unit tests @@ -1912,13 +1938,17 @@ class TestPagedMultiEndedAllocator(unittest.TestCase): "v2p_page[virt_pages] * page_size + offsets.", ) - # And the composite allocator's translate method must produce the - # same token-granular result (same page math). + # The composite emits KERNEL-FACING ids, not the physical token ids + # this helper returns; they coincide only at multiplier 1, which no + # sub-pool uses. + swa_mult = allocator.swa_kernel_page_multiplier + self.assertEqual(swa_mult, 2 * swa_spec.layer_num) composite_out = allocator.translate_loc_from_full_to_swa(v_tokens) + expected_dense = swa_phys_pages_direct * (PS * swa_mult) + offsets_in self.assertTrue( - bool((swa_phys.long() == composite_out.long()).all().item()), - "REGRESSION: the UnifiedSWAKVPool helper and the composite " - "allocator's translate_loc_from_full_to_swa must agree.", + bool((composite_out.long() == expected_dense.long()).all().item()), + "REGRESSION: translate_loc_from_full_to_swa must emit the swa " + "sub-pool's kernel-facing ids (phys_page * ps * blocks_per_page + offset).", ) @@ -2565,5 +2595,131 @@ class TestO3FusedAllocBind(unittest.TestCase): self.assertEqual(int(sa.physical_to_virtual[p].item()), v) +class TestSWACompositeDenseSurface(unittest.TestCase): + """The SWA composite's dense (kernel-facing) id surface. + + Presence of `translate_kv_loc_for_kernel` / `full_v2p_page_table` is what flips + the attention backends' kernel-facing-first probes, and the `page_stride` scale in + `translate_loc_from_full_to_swa` is what carries the swa kernel-facing space. + Everything must collapse + byte-identically at multiplier 1 — the strided arm every existing SWA model + runs — and follow `kernel_id(t) = v2p[t//ps]*(ps*mult) + t%ps` otherwise. + """ + + PS = 4 + FULL_L = 4 + SWA_L = 2 + + def _build(self): + full_spec = MHASubPoolSpec( + name="full", + layer_num=self.FULL_L, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="up", + ) + swa_spec = MHASubPoolSpec( + name="swa", + layer_num=self.SWA_L, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="down", + ) + n_full, n_swa = 64, 32 # tokens = 16 / 8 pages at PS=4 + total = n_full * full_spec.entry_bytes() + n_swa * swa_spec.entry_bytes() + pool = UnifiedKVPool( + total_bytes=total, + sub_pool_specs=[full_spec, swa_spec], + device=_DEV, + enable_memory_saver=False, + page_size=self.PS, + ) + kvcache = _FakeUnifiedSWAKVPool(pool) + return UnifiedSWATokenToKVPoolAllocator( + unified_buffer=pool, + kvcache=kvcache, + device=_DEV, + full_max_total_num_tokens=n_full, + swa_max_total_num_tokens=n_swa, + page_size=self.PS, + need_sort=False, + forward_stream=None, + ) + + def test_multipliers_come_from_the_specs(self): + """Both sides scale by their OWN sub-pool's block count, and the + composite exposes the raw v2p tables unwrapped. Nothing injects the + scale: a spec whose views are dense cannot be paired with a + physical-id multiplier, which is the state that writes physical ids + into view rows.""" + a = self._build() + self.assertEqual(a.kernel_page_multiplier, 2 * self.FULL_L) + self.assertEqual(a.swa_kernel_page_multiplier, 2 * self.SWA_L) + self.assertIs(a.full_v2p_page_table, a.full_attn_allocator.virtual_to_physical) + self.assertIs(a.swa_v2p_page_table, a.swa_attn_allocator.virtual_to_physical) + + def test_full_dense_translate_matches_formula(self): + mult = 2 * self.FULL_L + a = self._build() + v = a.alloc(3 * self.PS) + self.assertIsNotNone(v) + v2p = a.full_attn_allocator.virtual_to_physical + expected = v2p[v // self.PS] * (self.PS * mult) + v % self.PS + self.assertTrue(torch.equal(a.translate_kv_loc_for_kernel(v), expected)) + # The PHYSICAL translate must stay unscaled — compaction and the byte + # machinery depend on it staying in physical space. + phys = v2p[v // self.PS] * self.PS + v % self.PS + self.assertTrue(torch.equal(a.translate_kv_loc(v), phys)) + + def test_dense_translate_accepts_an_int32_page_table(self): + """REGRESSION: fa3 translates its own page table, which is int32 and + 2-D. A gather that requires an int64 index (`torch.take`) crashes the + scheduler there while every int64 caller stays green. Both page sizes: + at ps == 1 the index IS the caller's tensor, at ps > 1 it is derived.""" + for ps in (1, 4): + with self.subTest(page_size=ps): + self.PS = ps + mult = 2 * self.FULL_L + a = self._build() + v = a.alloc(4 * ps) + self.assertIsNotNone(v) + v2p = a.full_attn_allocator.virtual_to_physical + expected = v2p[v // ps] * (ps * mult) + v % ps + page_table = v.to(torch.int32).view(2, -1) + got = a.translate_kv_loc_for_kernel(page_table) + self.assertEqual(got.shape, page_table.shape) + self.assertTrue(torch.equal(got.reshape(-1), expected)) + # `out=` takes the same int32 index; the buffer stays int64. + dst = torch.empty(page_table.shape, dtype=torch.int64, device=_DEV) + a.translate_kv_loc_for_kernel(page_table, out=dst) + self.assertTrue(torch.equal(dst.reshape(-1), expected)) + + def test_swa_translate_scales_page_stride(self): + mult = 2 * self.SWA_L + a = self._build() + v = a.alloc(3 * self.PS) + self.assertIsNotNone(v) + v2p_swa = a.swa_attn_allocator.virtual_to_physical + expected = v2p_swa[v // self.PS] * (self.PS * mult) + v % self.PS + self.assertTrue(torch.equal(a.translate_loc_from_full_to_swa(v), expected)) + + def test_swa_dense_tombstone_still_lands_on_sink(self): + """The scaled stride must not break the tombstone clamp: a tombstoned + page's ids (v2p == -1 -> -stride + offset, negative for every in-page + offset) still land on the sink, never negative.""" + mult = 2 * self.SWA_L + a = self._build() + v = a.alloc(2 * self.PS) + self.assertIsNotNone(v) + tomb_page = int(v[0].item()) // self.PS + a.swa_attn_allocator.virtual_to_physical[tomb_page] = -1 + got = a.translate_loc_from_full_to_swa(v) + self.assertTrue(bool((got >= 0).all().item())) + in_tomb = v // self.PS == tomb_page + self.assertTrue(bool((got[in_tomb] == 0).all().item())) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_page_major_layout.py b/test/registered/unit/mem_cache/test_page_major_layout.py index 9bfd50895..002ff97b3 100644 --- a/test/registered/unit/mem_cache/test_page_major_layout.py +++ b/test/registered/unit/mem_cache/test_page_major_layout.py @@ -1,12 +1,9 @@ -"""CPU correctness tests for the page-major layer-major envelope layout. +"""CPU correctness tests for the page-major envelope Mamba state views. -Covers the standalone view builders (no allocator / shared pool): - - - ``build_page_major_mha_views``: 4-D K/V views with correct addressing at - page_size 1 (token-granularity envelope) and > 1 (layer-major within a page), - and no aliasing across layers / slots. - - ``build_page_major_mamba_views``: conv / temporal state views. - - ``move_kv_cache_native`` 4-D branch: relocating token rows preserves data. +Covers the standalone ``build_page_major_mamba_views`` builder (no allocator / +shared pool): conv / temporal state views with correct shapes and no aliasing +across layers / slots. The unified pool stores its Mamba/KDA state through +these views. Runs on CPU — pure-torch advanced indexing, no Triton. @@ -23,107 +20,10 @@ import torch from sglang.srt.mem_cache.layout.page_major import ( build_page_major_mamba_views, - build_page_major_mha_views, mamba_entry_bytes, - mha_entry_bytes, ) -from sglang.srt.mem_cache.memory_pool import move_kv_cache_native _DEV = "cpu" -_DT = torch.float32 - - -def _make_mha_views(layer_num, head_num, head_dim, v_head_dim, page_size, num_pages): - entry = mha_entry_bytes( - layer_num=layer_num, - head_num=head_num, - head_dim=head_dim, - v_head_dim=v_head_dim, - itemsize=_DT.itemsize, - ) - raw = torch.zeros(num_pages * page_size * entry, dtype=torch.uint8, device=_DEV) - k, v = build_page_major_mha_views( - raw, - layer_num=layer_num, - head_num=head_num, - head_dim=head_dim, - v_head_dim=v_head_dim, - store_dtype=_DT, - page_size=page_size, - num_pages=num_pages, - ) - return raw, k, v - - -class TestPageMajorMHAViews(unittest.TestCase): - def test_view_shapes(self): - _, k, v = _make_mha_views(3, 2, 4, 4, page_size=2, num_pages=4) - self.assertEqual(len(k), 3) - for t in k: - self.assertEqual(tuple(t.shape), (4, 2, 2, 4)) - for t in v: - self.assertEqual(tuple(t.shape), (4, 2, 2, 4)) - - def test_no_aliasing_ps1(self): - # Every (layer, slot) cell must be independently addressable. - layer_num, slots = 3, 5 - _, k, v = _make_mha_views(layer_num, 2, 4, 4, page_size=1, num_pages=slots) - for L in range(layer_num): - for s in range(slots): - k[L][s, 0] = float(100 + L * 10 + s) - v[L][s, 0] = float(200 + L * 10 + s) - for L in range(layer_num): - for s in range(slots): - self.assertTrue(torch.all(k[L][s, 0] == float(100 + L * 10 + s))) - self.assertTrue(torch.all(v[L][s, 0] == float(200 + L * 10 + s))) - - def test_page_slot_addressing_ps_gt1(self): - # token id t -> page t // ps, slot t % ps; no aliasing across tokens. - ps, pages = 2, 4 - total = ps * pages - _, k, _ = _make_mha_views(2, 1, 2, 2, page_size=ps, num_pages=pages) - for L in range(2): - for t in range(total): - k[L][t // ps, t % ps, 0] = float(1000 + L * 100 + t) - for L in range(2): - for t in range(total): - self.assertEqual( - float(k[L][t // ps, t % ps, 0, 0].item()), 1000 + L * 100 + t - ) - - def test_asymmetric_v_head_dim(self): - _, k, v = _make_mha_views(2, 2, 6, 4, page_size=1, num_pages=3) - self.assertEqual(tuple(k[0].shape), (3, 1, 2, 6)) - self.assertEqual(tuple(v[0].shape), (3, 1, 2, 4)) - - -class TestPageMajorMove(unittest.TestCase): - def test_move_ps1(self): - slots = 6 - _, k, v = _make_mha_views(2, 1, 4, 4, page_size=1, num_pages=slots) - for L in range(2): - for s in range(slots): - k[L][s, 0] = float(s + 1) - v[L][s, 0] = float(-(s + 1)) - tgt = torch.tensor([0, 1], dtype=torch.int64) - src = torch.tensor([4, 5], dtype=torch.int64) - move_kv_cache_native(k, v, tgt, src, page_size=1) - for L in range(2): - self.assertTrue(torch.all(k[L][0, 0] == 5.0)) - self.assertTrue(torch.all(k[L][1, 0] == 6.0)) - self.assertTrue(torch.all(v[L][0, 0] == -5.0)) - - def test_move_ps_gt1(self): - ps, pages = 2, 4 - total = ps * pages - _, k, v = _make_mha_views(1, 1, 2, 2, page_size=ps, num_pages=pages) - for t in range(total): - k[0][t // ps, t % ps, 0] = float(t + 1) - tgt = torch.tensor([0, 3], dtype=torch.int64) # page0 slot0, page1 slot1 - src = torch.tensor([6, 7], dtype=torch.int64) # page3 slot0, page3 slot1 - move_kv_cache_native(k, v, tgt, src, page_size=ps) - self.assertEqual(float(k[0][0, 0, 0, 0].item()), 7.0) - self.assertEqual(float(k[0][1, 1, 0, 0].item()), 8.0) class TestMambaEnvelopeViews(unittest.TestCase): diff --git a/test/registered/unit/mem_cache/test_pd_envelope_transfer_layout.py b/test/registered/unit/mem_cache/test_pd_envelope_transfer_layout.py index f357c2aa9..07147ccc5 100644 --- a/test/registered/unit/mem_cache/test_pd_envelope_transfer_layout.py +++ b/test/registered/unit/mem_cache/test_pd_envelope_transfer_layout.py @@ -18,7 +18,7 @@ import unittest import torch from sglang.srt.mem_cache.layout.page_major import ( - build_dense_mla_views, + build_mla_views, build_page_major_mamba_views, mamba_entry_bytes, mla_entry_bytes, @@ -31,7 +31,7 @@ register_cpu_ci(est_time=60, suite="base-a-test-cpu") class TestMLAEnvelopeTransferAddressing(CustomTestCase): def test_page_envelope_matches_dense_views(self): - """Every (page, layer, slot) row written through the dense MLA views + """Every (page, layer, slot) row written through the MLA views must land at raw_ptr + page * page_envelope_bytes + layer-block offset, i.e. inside the page's transfer envelope.""" layer_num, page_size, kv_dim, num_pages = 3, 4, 8, 6 @@ -49,7 +49,7 @@ class TestMLAEnvelopeTransferAddressing(CustomTestCase): ) # +1 page envelope of tail pad, as UnifiedKVPool allocates for MLA. raw = torch.zeros((num_pages + 1) * page_bytes, dtype=torch.uint8) - views = build_dense_mla_views( + views = build_mla_views( raw, layer_num=layer_num, kv_cache_dim=kv_dim, diff --git a/test/registered/unit/mem_cache/test_store_cache_4d.py b/test/registered/unit/mem_cache/test_store_cache_4d.py deleted file mode 100644 index d899f1b0c..000000000 --- a/test/registered/unit/mem_cache/test_store_cache_4d.py +++ /dev/null @@ -1,412 +0,0 @@ -"""Parity tests for the `store_cache_4d` Triton kernel. - -The kernel writes K/V into the 4-D page-major envelope view. These tests prove -it produces byte-identical output to the legacy advanced-indexing path on -representative fixtures: - - - ``page_size = 1`` (envelope-degenerate, the critical compatibility case) - - ``page_size > 1`` (layer-major within page) - - both int32 and int64 ``loc`` dtypes - - bf16 and fp8_e5m2 view dtypes - - asymmetric ``head_dim != v_head_dim`` - - empty ``loc`` (no-op) - -Skipped on CPU — Triton requires a GPU. - - python -m pytest test/registered/unit/mem_cache/test_store_cache_4d.py -v -""" - -import importlib.util -import unittest - -import torch - -from sglang.test.ci.ci_register import register_cuda_ci - -_HAS_CUDA = torch.cuda.is_available() -# The set_kv_buffer integration test needs UnifiedMHATokenToKVPool, which only -# exists once the shared-KV-pool feature lands; skip it where absent. -_HAS_SHARED_POOL = ( - importlib.util.find_spec("sglang.srt.mem_cache.unified_memory_pool") is not None -) - -register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-small") - - -def _legacy_advanced_indexing_write( - k_view: torch.Tensor, - v_view: torch.Tensor, - cache_k: torch.Tensor, - cache_v: torch.Tensor, - loc: torch.Tensor, - page_size: int, -) -> None: - """Reference implementation: the legacy bypass-super() advanced-indexing - path that the Triton kernel replaces. Used as the byte-identity oracle - for the parity tests below. - """ - if page_size == 1: - k_view[loc, 0] = cache_k - v_view[loc, 0] = cache_v - else: - page_id = loc // page_size - tok_in_p = loc % page_size - k_view[page_id, tok_in_p] = cache_k - v_view[page_id, tok_in_p] = cache_v - - -@unittest.skipUnless(_HAS_CUDA, "Triton kernels require CUDA") -class TestStoreCache4D(unittest.TestCase): - """Byte-identity parity vs the legacy advanced-indexing write path.""" - - def _make_view_and_cache( - self, - num_pages: int, - page_size: int, - head_num: int, - head_dim: int, - v_head_dim: int, - N: int, - dtype: torch.dtype = torch.bfloat16, - loc_dtype: torch.dtype = torch.int64, - seed: int = 0xC0FFEE, - ): - torch.manual_seed(seed) - # The unified memory pool's views are 4-D `(num_pages, page_size, head_num, - # head_dim)` with the trailing two dims contiguous. We allocate two - # independent contiguous buffers (one for the kernel-under-test, - # one as the legacy-path target) so we can compare them. - k_view = torch.zeros( - (num_pages, page_size, head_num, head_dim), - dtype=dtype, - device="cuda", - ) - v_view = torch.zeros( - (num_pages, page_size, head_num, v_head_dim), - dtype=dtype, - device="cuda", - ) - cache_k = torch.randn( - (N, head_num, head_dim), dtype=torch.float32, device="cuda" - ).to(dtype) - cache_v = torch.randn( - (N, head_num, v_head_dim), dtype=torch.float32, device="cuda" - ).to(dtype) - # Valid loc values in [0, num_pages * page_size); generate without - # duplicates so the comparison is unambiguous (advanced-indexing - # with duplicates is order-undefined for both paths). - total_slots = num_pages * page_size - assert N <= total_slots - loc = torch.randperm(total_slots, device="cuda")[:N].to(loc_dtype) - return k_view, v_view, cache_k, cache_v, loc - - def _check_parity( - self, - num_pages: int, - page_size: int, - head_num: int, - head_dim: int, - v_head_dim: int, - N: int, - dtype: torch.dtype = torch.bfloat16, - loc_dtype: torch.dtype = torch.int64, - ): - from sglang.kernels.ops.kvcache.cache_move import store_cache_4d - - # Two independent target buffers — one for the kernel, one for the - # legacy reference path. - k_kernel, v_kernel, cache_k, cache_v, loc = self._make_view_and_cache( - num_pages, - page_size, - head_num, - head_dim, - v_head_dim, - N, - dtype=dtype, - loc_dtype=loc_dtype, - ) - k_legacy = k_kernel.clone() - v_legacy = v_kernel.clone() - - # Kernel-under-test - store_cache_4d(k_kernel, v_kernel, cache_k, cache_v, loc, page_size) - # Legacy reference - _legacy_advanced_indexing_write( - k_legacy, v_legacy, cache_k, cache_v, loc, page_size - ) - - # Byte-identical comparison — the kernel must reproduce the - # advanced-indexing path bit-for-bit, NOT just numerically close. - # For fp8 dtypes, torch.equal works on the integer bit pattern. - self.assertTrue( - torch.equal(k_kernel, k_legacy), - f"K view mismatch: ps={page_size}, dtype={dtype}, " - f"loc_dtype={loc_dtype}, N={N}", - ) - self.assertTrue( - torch.equal(v_kernel, v_legacy), - f"V view mismatch: ps={page_size}, dtype={dtype}, " - f"loc_dtype={loc_dtype}, N={N}", - ) - - # ---- Test 1: ps=1 envelope-degenerate (the critical compat case) ---- - - def test_store_cache_4d_ps1_byte_identical(self): - """At page_size=1 the kernel constexpr-folds to the slot-major - envelope view. Output must be byte-identical to advanced indexing. - This protects against byte-layout regression.""" - self._check_parity( - num_pages=64, - page_size=1, - head_num=4, - head_dim=128, - v_head_dim=128, - N=16, - ) - - # ---- Test 2: ps>1 layer-major within page ---- - - def test_store_cache_4d_ps_gt1_byte_identical(self): - """At page_size > 1 the kernel splits loc into (page_id, tok_in_p) - and writes via the 4-D stride. Output must match the equivalent - advanced-indexing write.""" - self._check_parity( - num_pages=8, - page_size=64, - head_num=4, - head_dim=128, - v_head_dim=128, - N=128, - ) - - # ---- Test 3: int32 loc dtype ---- - - def test_store_cache_4d_int32_loc(self): - """The SWA-side path passes int32 loc (matches the SWA Triton - kernel contract). PyTorch advanced indexing tolerates either - int32 or int64; the kernel must too.""" - self._check_parity( - num_pages=32, - page_size=1, - head_num=4, - head_dim=64, - v_head_dim=64, - N=10, - loc_dtype=torch.int32, - ) - - # ---- Test 4: int64 loc dtype (already exercised, explicit) ---- - - # ---- Test 5: bf16 dtype (the production case) ---- - - # ---- Test 6: fp8_e5m2 dtype ---- - - def test_store_cache_4d_dtype_fp8_e5m2(self): - """fp8_e5m2 is used for KV-cache quantization. Caller is responsible - for the cast; the kernel sees same-dtype source and destination.""" - self._check_parity( - num_pages=16, - page_size=64, - head_num=4, - head_dim=128, - v_head_dim=128, - N=64, - dtype=torch.float8_e5m2, - ) - - # ---- Test 7: empty loc (no-op) ---- - - def test_store_cache_4d_empty_loc(self): - """N=0 must be a no-op: no kernel launch, no exception, no buffer - mutation.""" - from sglang.kernels.ops.kvcache.cache_move import store_cache_4d - - k_view = torch.zeros((8, 4, 4, 64), dtype=torch.bfloat16, device="cuda") - v_view = torch.zeros((8, 4, 4, 64), dtype=torch.bfloat16, device="cuda") - k_before = k_view.clone() - v_before = v_view.clone() - cache_k = torch.empty((0, 4, 64), dtype=torch.bfloat16, device="cuda") - cache_v = torch.empty((0, 4, 64), dtype=torch.bfloat16, device="cuda") - loc = torch.empty((0,), dtype=torch.int64, device="cuda") - - store_cache_4d(k_view, v_view, cache_k, cache_v, loc, page_size=4) - - # Buffers must be unchanged. - self.assertTrue(torch.equal(k_view, k_before)) - self.assertTrue(torch.equal(v_view, v_before)) - - # ---- Test 8: head_dim != v_head_dim (asymmetric, e.g. MLA-style) ---- - - def test_store_cache_4d_v_head_dim_differs(self): - """When v_head_dim != head_dim, the kernel's K and V branches use - different per-token strides. Exercises the stride_k_tok ≠ - stride_v_tok branch.""" - self._check_parity( - num_pages=8, - page_size=16, - head_num=2, - head_dim=128, - v_head_dim=64, - N=16, - ) - - -@unittest.skipUnless(_HAS_CUDA, "Triton kernels require CUDA") -class TestStoreCache4DAssertions(unittest.TestCase): - """The wrapper's contract assertions must fire on bad inputs.""" - - def test_rejects_non_contiguous_view_trailing_dim(self): - """Wrapper requires `stride[-1] == 1` and `stride[-2] == head_dim` - (the trailing two dims must be contiguous). A permutation that - breaks this should trigger AssertionError.""" - from sglang.kernels.ops.kvcache.cache_move import store_cache_4d - - # Build a 4-D view, then permute the last two dims → trailing - # contiguity violated. - k_view = torch.zeros( - (4, 4, 4, 64), dtype=torch.bfloat16, device="cuda" - ).permute( - 0, 1, 3, 2 - ) # now shape (4, 4, 64, 4); strides broken - v_view = torch.zeros((4, 4, 4, 64), dtype=torch.bfloat16, device="cuda") - cache_k = torch.zeros((2, 4, 64), dtype=torch.bfloat16, device="cuda") - cache_v = torch.zeros((2, 4, 64), dtype=torch.bfloat16, device="cuda") - loc = torch.arange(2, dtype=torch.int64, device="cuda") - with self.assertRaises(AssertionError): - store_cache_4d(k_view, v_view, cache_k, cache_v, loc, page_size=4) - - def test_rejects_dtype_mismatch(self): - """All four tensors must share a dtype; the caller is responsible - for any cast before the call.""" - from sglang.kernels.ops.kvcache.cache_move import store_cache_4d - - k_view = torch.zeros((4, 4, 4, 64), dtype=torch.bfloat16, device="cuda") - v_view = torch.zeros((4, 4, 4, 64), dtype=torch.bfloat16, device="cuda") - cache_k = torch.zeros((2, 4, 64), dtype=torch.float16, device="cuda") - cache_v = torch.zeros((2, 4, 64), dtype=torch.bfloat16, device="cuda") - loc = torch.arange(2, dtype=torch.int64, device="cuda") - with self.assertRaises(AssertionError): - store_cache_4d(k_view, v_view, cache_k, cache_v, loc, page_size=4) - - -@unittest.skipUnless( - _HAS_CUDA and _HAS_SHARED_POOL, - "Triton kernels require CUDA; UnifiedMHATokenToKVPool required", -) -class TestStoreCache4DThroughSetKVBuffer(unittest.TestCase): - """Integration parity test — exercises the kernel through the FULL - ``UnifiedMHATokenToKVPool.set_kv_buffer`` path (the direct PHYSICAL write + - the dtype cast; the pool no longer translates). Confirms it produces - bit-identical output to a PyTorch advanced-indexing reference write. - """ - - def _build_pool(self, page_size: int): - """Build a small UnifiedMHATokenToKVPool. The pool writes PHYSICAL locs - directly (no allocator / v2p translate), so `set_kv_buffer` receives the - already-physical write location.""" - import torch as _t - - from sglang.srt.mem_cache.unified_memory_pool import ( - MHASubPoolSpec, - UnifiedKVPool, - UnifiedMHATokenToKVPool, - ) - - spec = MHASubPoolSpec( - name="full", - layer_num=2, - head_num=4, - head_dim=64, - store_dtype=_t.bfloat16, - grow_direction="up", - ) - total = spec.entry_bytes() * 64 - # Use a peer to satisfy the two-sub-pool contract. - peer = MHASubPoolSpec( - name="swa", - layer_num=1, - head_num=4, - head_dim=64, - store_dtype=_t.bfloat16, - grow_direction="down", - ) - pool = UnifiedKVPool( - total_bytes=total + peer.entry_bytes() * 16, - sub_pool_specs=[spec, peer], - device="cuda", - enable_memory_saver=False, - page_size=page_size, - ) - kv_pool = UnifiedMHATokenToKVPool( - unified_buffer=pool, - sub_pool_name="full", - page_size=page_size, - start_layer=0, - end_layer=2, - enable_alt_stream=False, - ) - - return kv_pool - - def _run_set_kv_buffer_and_compare(self, page_size: int): - import torch as _t - - kv_pool = self._build_pool(page_size) - - # A fake `layer` object with the minimum interface - # `set_kv_buffer` reads: `.layer_id`. - class _FakeLayer: - layer_id = 0 - - layer = _FakeLayer() - head_num, head_dim = 4, 64 - N = 16 - # Generate valid loc in range [0, num_pages * page_size). - num_pages = kv_pool.k_buffer[0].shape[0] - total = num_pages * page_size - assert N <= total - loc = _t.randperm(total, device="cuda")[:N].to(_t.int64) - cache_k = _t.randn((N, head_num, head_dim), dtype=_t.bfloat16, device="cuda") - cache_v = _t.randn((N, head_num, head_dim), dtype=_t.bfloat16, device="cuda") - - # Production path: the Triton `store_cache_4d` kernel via set_kv_buffer. - kv_pool.set_kv_buffer(layer, loc, cache_k.clone(), cache_v.clone()) - k_kernel = kv_pool.k_buffer[0].clone() - v_kernel = kv_pool.v_buffer[0].clone() - - # Reference: PyTorch advanced-indexing into a fresh view at the same - # (physical) loc, with no dtype cast (store_dtype == dtype) — the exact - # write the kernel performs. - kv_pool.k_buffer[0].zero_() - kv_pool.v_buffer[0].zero_() - k_view = kv_pool.k_buffer[0] - v_view = kv_pool.v_buffer[0] - if page_size == 1: - k_view[loc, 0] = cache_k - v_view[loc, 0] = cache_v - else: - page_id = loc // page_size - tok_in_p = loc % page_size - k_view[page_id, tok_in_p] = cache_k - v_view[page_id, tok_in_p] = cache_v - k_ref = kv_pool.k_buffer[0].clone() - v_ref = kv_pool.v_buffer[0].clone() - - self.assertTrue( - _t.equal(k_kernel, k_ref), - f"K view mismatch through set_kv_buffer at ps={page_size}", - ) - self.assertTrue( - _t.equal(v_kernel, v_ref), - f"V view mismatch through set_kv_buffer at ps={page_size}", - ) - - def test_integration_ps1(self): - self._run_set_kv_buffer_and_compare(page_size=1) - - def test_integration_ps64(self): - self._run_set_kv_buffer_and_compare(page_size=64) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/mem_cache/test_triton_kernel_layout.py b/test/registered/unit/mem_cache/test_triton_kernel_layout.py deleted file mode 100644 index 9d4b56816..000000000 --- a/test/registered/unit/mem_cache/test_triton_kernel_layout.py +++ /dev/null @@ -1,190 +0,0 @@ -"""Triton-kernel parity test for the page-aware decode / extend kernels. - -Verifies that the modified decode / extend Triton kernels produce -bit-identical output when called against: - - (a) the legacy 3-D ``[N, head, dim]`` KV view (PAGE_SIZE=1 default), - (b) the new 4-D ``[num_pages, page_size, head, dim]`` view with - ``page_size=1`` (degenerate envelope — same physical bytes as (a)), - (c) the new 4-D view with ``page_size>1`` (layer-major), using the - same logical KV data but routed via page-aware address math. - -Output for (a) vs (b) must be bit-identical at PAGE_SIZE=1 (the kernel -specializes to the legacy branch). Output for (c) must match a hand- -computed reference SDPA result (same logical attention; different byte -layout). - -Skipped on CPU — Triton requires a GPU. - - python -m pytest test/registered/unit/mem_cache/test_triton_kernel_layout.py -v -""" - -import unittest - -import torch - -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -_HAS_CUDA = torch.cuda.is_available() - -register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-small") -register_amd_ci(est_time=5, stage="stage-b", runner_config="1-gpu-small-amd") - - -@unittest.skipUnless(_HAS_CUDA, "Triton kernels require CUDA") -class TestTritonKernelLayoutParity(unittest.TestCase): - """Decode + extend kernel parity across (3-D, 4-D ps=1, 4-D ps>1).""" - - def _setup_decode_inputs( - self, bs=2, head_num=2, head_dim=8, num_slots=64, dtype=torch.float16 - ): - torch.manual_seed(0xC0FFEE) - # Logical KV: shape [num_slots, head_num, head_dim] - logical_kv_k = torch.randn( - num_slots, head_num, head_dim, dtype=dtype, device="cuda" - ) - logical_kv_v = torch.randn( - num_slots, head_num, head_dim, dtype=dtype, device="cuda" - ) - q = torch.randn(bs, head_num, head_dim, dtype=dtype, device="cuda") - # All requests use the first `seq_len` slots. - seq_len = 16 - kv_indices_per_req = torch.arange(seq_len, dtype=torch.int64, device="cuda") - kv_indices = kv_indices_per_req.repeat(bs) # [bs * seq_len] - kv_indptr = torch.tensor( - [i * seq_len for i in range(bs + 1)], dtype=torch.int32, device="cuda" - ) - return q, logical_kv_k, logical_kv_v, kv_indptr, kv_indices, seq_len - - def _run_decode(self, q, k_buf, v_buf, kv_indptr, kv_indices, page_size): - from sglang.kernels.ops.attention.decode_attention import ( - decode_attention_fwd, - ) - - bs, head_num, head_dim = q.shape - max_kv_splits = 4 - attn_logits = torch.empty( - (bs, head_num, max_kv_splits, head_dim), - dtype=torch.float32, - device="cuda", - ) - attn_lse = torch.empty( - (bs, head_num, max_kv_splits), - dtype=torch.float32, - device="cuda", - ) - o = torch.empty_like(q) - num_kv_splits = torch.full( - (bs,), max_kv_splits, dtype=torch.int32, device="cuda" - ) - decode_attention_fwd( - q, - k_buf, - v_buf, - o, - kv_indptr, - kv_indices, - attn_logits, - attn_lse, - num_kv_splits, - max_kv_splits, - sm_scale=1.0 / (head_dim**0.5), - k_scale=1.0, - v_scale=1.0, - logit_cap=0.0, - page_size=page_size, - ) - return o - - def test_decode_3d_vs_4d_ps1_byte_identical(self): - """(a) vs (b): same physical bytes, different view shape. - Triton specializes PAGE_SIZE=1 to the legacy branch; output must - be bit-identical (modulo non-deterministic FP add ordering, which - we sidestep here since the kernels use deterministic reductions - for fixed input + grid).""" - q, k, v, kv_indptr, kv_indices, seq_len = self._setup_decode_inputs() - # (a) legacy 3-D view - o_3d = self._run_decode(q, k, v, kv_indptr, kv_indices, page_size=1) - # (b) 4-D view: reshape SAME physical bytes to (num_pages=N, 1, head, dim) - num_slots = k.shape[0] - k_4d = k.view(num_slots, 1, *k.shape[1:]) - v_4d = v.view(num_slots, 1, *v.shape[1:]) - o_4d_ps1 = self._run_decode(q, k_4d, v_4d, kv_indptr, kv_indices, page_size=1) - # bit-identical (same byte layout, same PAGE_SIZE specialization) - self.assertTrue(torch.equal(o_3d, o_4d_ps1)) - - def test_extend_3d_vs_4d_ps1_byte_identical(self): - """Same parity check for extend kernel.""" - from sglang.kernels.ops.attention.extend_attention import ( - extend_attention_fwd, - ) - - torch.manual_seed(0xDEADBEEF) - # head_dim must be >= 16: the extend kernel's QK^T tl.dot requires the - # contraction dim K (= head_dim) >= 16 on modern GPU archs (Hopper+). - head_num, head_dim = 2, 32 - num_slots = 32 - dtype = torch.float16 - bs = 2 - prefix_len = 8 - extend_len = 4 - - k_buffer = torch.randn( - num_slots, head_num, head_dim, dtype=dtype, device="cuda" - ) - v_buffer = torch.randn( - num_slots, head_num, head_dim, dtype=dtype, device="cuda" - ) - q_extend = torch.randn( - bs * extend_len, head_num, head_dim, dtype=dtype, device="cuda" - ) - k_extend = torch.randn( - bs * extend_len, head_num, head_dim, dtype=dtype, device="cuda" - ) - v_extend = torch.randn( - bs * extend_len, head_num, head_dim, dtype=dtype, device="cuda" - ) - o = torch.empty_like(q_extend) - - qo_indptr = torch.tensor( - [i * extend_len for i in range(bs + 1)], dtype=torch.int32, device="cuda" - ) - kv_indptr = torch.tensor( - [i * prefix_len for i in range(bs + 1)], dtype=torch.int32, device="cuda" - ) - kv_indices = torch.arange(prefix_len, dtype=torch.int64, device="cuda").repeat( - bs - ) - - def run(k_buf, v_buf, page_size): - o_out = torch.empty_like(q_extend) - extend_attention_fwd( - q_extend, - k_extend, - v_extend, - o_out, - k_buf, - v_buf, - qo_indptr, - kv_indptr, - kv_indices, - custom_mask=None, - is_causal=True, - mask_indptr=None, - max_len_extend=extend_len, - k_scale=1.0, - v_scale=1.0, - sm_scale=1.0 / (head_dim**0.5), - page_size=page_size, - ) - return o_out - - o_3d = run(k_buffer, v_buffer, page_size=1) - k_4d = k_buffer.view(num_slots, 1, *k_buffer.shape[1:]) - v_4d = v_buffer.view(num_slots, 1, *v_buffer.shape[1:]) - o_4d_ps1 = run(k_4d, v_4d, page_size=1) - self.assertTrue(torch.equal(o_3d, o_4d_ps1)) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_handout_zeroing.py b/test/registered/unit/mem_cache/test_unified_handout_zeroing.py index 6349f29e8..fa456bd87 100644 --- a/test/registered/unit/mem_cache/test_unified_handout_zeroing.py +++ b/test/registered/unit/mem_cache/test_unified_handout_zeroing.py @@ -47,7 +47,6 @@ def _build(device, page_size=1, kernel_page_multiplier=None): device=device, enable_memory_saver=False, page_size=page_size, - view_tail_pad_bytes=page_size * full_spec.entry_bytes(), ) kvcache = UnifiedMLATokenToKVPool( unified_buffer=buf, diff --git a/test/registered/unit/mem_cache/test_unified_mha_views.py b/test/registered/unit/mem_cache/test_unified_mha_views.py new file mode 100644 index 000000000..a28fcbe05 --- /dev/null +++ b/test/registered/unit/mem_cache/test_unified_mha_views.py @@ -0,0 +1,531 @@ +# Copyright 2023-2026 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. +# ============================================================================== +"""Dense MHA K/V views for the unified memory pool (uniform-row hybrid models). + +Covers, CPU-only (pure torch — no GPU / Triton kernels): + - `build_mha_views` refuses an asymmetric-KV spec: its addressing + assumes one uniform row width, so it is the boundary that checks; + - `build_mha_views` addressing: view_l[kernel_id(t)] must land exactly at + the page-major envelope byte offset the STRIDED builder assigns to the same + (page, slot, layer, K|V) cell — the two builders are views over one truth; + - K and V of one token share ONE kernel-facing id (per-layer origin shift does the + disambiguation), with no aliasing across the 2*L overlapping views; + - the missing-tail-pad and asymmetric-dims cases fail loud at construction. + +Addressing law under test (the derived property everything else builds on): + + kernel_id(t) = (t // ps) * (ps * 2L) + t % ps + K of layer l at block 2l, V at block 2l+1, blocks are ps rows of + head_num*head_dim elements — offsets identical to + MHASubPoolSpec.layer_k/v_offset_in_page when rows are uniform. + + python -m pytest test/registered/unit/mem_cache/test_unified_mha_views.py -v +""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=8, suite="base-a-test-cpu") + +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.environ import envs +from sglang.srt.mem_cache.layout.page_major import ( + build_mha_views, + mha_entry_bytes, +) +from sglang.srt.mem_cache.unified_memory_pool import ( + MHASubPoolSpec, + UnifiedKVPool, + UnifiedMHATokenToKVPool, +) + +_DEV = "cpu" +# `set_kv_buffer` dispatches on the PLATFORM (memory_pool._is_cuda, resolved at +# import), not on the tensors it is handed, so cases driving it must build on +# the platform's device. The rest of this file is byte arithmetic, so CPU. +_STORE_DEV = "cuda" if torch.cuda.is_available() else "cpu" + +# Small-but-nontrivial MHA geometry: L=2 layers, H=2 heads, D=4, so every byte +# offset is hand-checkable. blocks = 2L = 4 per page. +_L = 2 +_H = 2 +_D = 4 +_ROW = _H * _D # row elements +_DTYPE = torch.bfloat16 +_ITEM = _DTYPE.itemsize +_BLOCKS = 2 * _L + + +def _mha_spec(head_dim=_D, v_head_dim=None, layer_num=_L, grow="down"): + return MHASubPoolSpec( + name="full", + layer_num=layer_num, + head_num=_H, + head_dim=head_dim, + v_head_dim=v_head_dim, + store_dtype=_DTYPE, + grow_direction=grow, + ) + + +def _kernel_id(t, ps): + return (t // ps) * (ps * _BLOCKS) + t % ps + + +def _make_raw(ps, num_pages, pad_pages=1): + page_bytes = ps * _BLOCKS * _ROW * _ITEM + raw = torch.zeros( + (num_pages + pad_pages) * page_bytes, dtype=torch.uint8, device=_DEV + ) + return raw + + +def _build_views(raw, ps, num_pages, head_dim=_D, v_head_dim=_D, layer_num=_L): + return build_mha_views( + raw, + layer_num=layer_num, + head_num=_H, + head_dim=head_dim, + v_head_dim=v_head_dim, + store_dtype=_DTYPE, + page_size=ps, + num_pages=num_pages, + ) + + +def _reference_strided_views(raw, *, page_size, num_pages, anchor_bytes=0): + """Independent 4-D strided description of the page-major envelope. + + This is the retired production strided builder, kept here as the oracle: + per-layer ``(num_pages, page_size, head_num, head_dim)`` views addressed by + ``(page, slot)``, so the builder's addressing can be cross-checked + against a second, independently-derived description of the same bytes. + """ + k_row_bytes = _ROW * _ITEM + v_row_bytes = _ROW * _ITEM + page_bytes = page_size * _L * (k_row_bytes + v_row_bytes) + as_dtype_view = raw.view(_DTYPE) + k_stride = (page_bytes // _ITEM, k_row_bytes // _ITEM, _D, 1) + v_stride = (page_bytes // _ITEM, v_row_bytes // _ITEM, _D, 1) + shape = (num_pages, page_size, _H, _D) + k_views, v_views = [], [] + for layer in range(_L): + k_base = anchor_bytes + layer * page_size * (k_row_bytes + v_row_bytes) + v_base = k_base + page_size * k_row_bytes + k_views.append( + torch.as_strided( + as_dtype_view, + size=shape, + stride=k_stride, + storage_offset=k_base // _ITEM, + ) + ) + v_views.append( + torch.as_strided( + as_dtype_view, + size=shape, + stride=v_stride, + storage_offset=v_base // _ITEM, + ) + ) + return k_views, v_views + + +class TestMHADenseSpecSurface(unittest.TestCase): + def test_asymmetric_rows_refused_by_the_view_builder(self): + """The row-block array exists only for uniform rows, so the builder + whose addressing depends on it is the one that refuses (the MiMoV2 + shape, scaled down). ServerArgs screens such models out of + --enable-unified-memory long before we get here; this is the check for + a caller that reaches the builder directly.""" + spec = _mha_spec() + raw = torch.zeros(1 << 16, dtype=torch.uint8) + with self.assertRaises(AssertionError): + build_mha_views( + raw, + layer_num=spec.layer_num, + head_num=spec.head_num, + head_dim=6, + v_head_dim=4, + store_dtype=spec.store_dtype, + page_size=1, + num_pages=4, + ) + + def test_spec_offsets_equal_block_origins(self): + """The spec's byte math and the view builder's origins are two + independent derivations of the envelope; under uniform rows they must + agree: layer_k_offset(l) == (2l)*ps*row, layer_v_offset(l) == (2l+1)*ps*row.""" + spec = _mha_spec() + for ps in (1, 4): + row = spec.k_row_bytes() + for l in range(_L): + self.assertEqual(spec.layer_k_offset_in_page(l, ps), (2 * l) * ps * row) + self.assertEqual( + spec.layer_v_offset_in_page(l, ps), (2 * l + 1) * ps * row + ) + + def test_entry_bytes_matches_layout_helper(self): + spec = _mha_spec() + self.assertEqual( + spec.entry_bytes(), + mha_entry_bytes( + layer_num=_L, head_num=_H, head_dim=_D, v_head_dim=_D, itemsize=_ITEM + ), + ) + + +class TestDenseMHAViews(unittest.TestCase): + def test_view_shapes_are_stock_mha(self): + ps, num_pages = 4, 6 + k_views, v_views = _build_views(_make_raw(ps, num_pages), ps, num_pages) + n_rows = num_pages * _BLOCKS * ps + self.assertEqual(len(k_views), _L) + self.assertEqual(len(v_views), _L) + for v in (*k_views, *v_views): + # The stock MHATokenToKVPool per-layer signature: 3-D, packed rows. + self.assertEqual(tuple(v.shape), (n_rows, _H, _D)) + self.assertEqual(v.stride(), (_ROW, _D, 1)) + + def test_addressing_matches_strided_reference(self): + """Cross-readback: bytes written through the reference STRIDED views at + (page, slot) must be read back through the views at kernel_id(t), + for both K and V of every layer — and vice versa. This pins that the + view builder and the independent strided description agree on the + same physical envelope.""" + for ps in (1, 4): + num_pages = 5 + raw = _make_raw(ps, num_pages) + sk, sv = _reference_strided_views(raw, page_size=ps, num_pages=num_pages) + dk, dv = _build_views(raw, ps, num_pages) + probes = [(0, 0, 0), (1, 1, ps - 1), (4, 0, ps // 2), (3, 1, 0)] + # strided-write -> view-read + for p, l, s in probes: + t = p * ps + s + d = _kernel_id(t, ps) + sk[l][p, s] = float(p * 100 + l * 10 + s + 1) + sv[l][p, s] = float(p * 100 + l * 10 + s + 2) + self.assertTrue( + torch.all(dk[l][d] == float(p * 100 + l * 10 + s + 1)), + f"K (p={p}, l={l}, s={s}, ps={ps}) view readback off-formula", + ) + self.assertTrue( + torch.all(dv[l][d] == float(p * 100 + l * 10 + s + 2)), + f"V (p={p}, l={l}, s={s}, ps={ps}) view readback off-formula", + ) + # view-write -> strided-read + for p, l, s in probes: + t = p * ps + s + d = _kernel_id(t, ps) + dk[l][d] = float(p * 100 + l * 10 + s + 3) + dv[l][d] = float(p * 100 + l * 10 + s + 4) + self.assertTrue( + torch.all(sk[l][p, s] == float(p * 100 + l * 10 + s + 3)) + ) + self.assertTrue( + torch.all(sv[l][p, s] == float(p * 100 + l * 10 + s + 4)) + ) + + def test_byte_addresses_match_envelope_formula(self): + """The per-layer view's byte address for token ``t``, layer ``L`` must equal + the hand-computed envelope formula: page origin + layer-block origin + + slot offset. Independent of any view builder — this is the raw layout + contract every envelope consumer (moves, sizing, transfer math) relies + on.""" + k_row = _ROW * _ITEM + v_row = _ROW * _ITEM + for ps in (1, 4): + num_pages = 5 + page_bytes = ps * _L * (k_row + v_row) + dk, dv = _build_views(_make_raw(ps, num_pages), ps, num_pages) + for t in (0, 1, ps, 3 * ps + (ps - 1), 4 * ps): + d = _kernel_id(t, ps) + for L in range(_L): + expected_k = ( + (t // ps) * page_bytes + + L * ps * (k_row + v_row) + + (t % ps) * k_row + ) + expected_v = ( + (t // ps) * page_bytes + + L * ps * (k_row + v_row) + + ps * k_row + + (t % ps) * v_row + ) + got_k = (dk[L].storage_offset() + d * dk[L].stride(0)) * _ITEM + got_v = (dv[L].storage_offset() + d * dv[L].stride(0)) * _ITEM + self.assertEqual(got_k, expected_k, f"K t={t} L={L} ps={ps}") + self.assertEqual(got_v, expected_v, f"V t={t} L={L} ps={ps}") + + def test_k_and_v_share_one_kernel_id_without_aliasing(self): + """One kernel-facing id, 2L distinct cells (K and V of every layer): writes + through all 2L views at the SAME id must not clobber each other.""" + ps, num_pages = 4, 4 + dk, dv = _build_views(_make_raw(ps, num_pages), ps, num_pages) + t = 2 * ps + 1 # page 2, slot 1 + d = _kernel_id(t, ps) + for l in range(_L): + dk[l][d] = float(2 * l + 1) + dv[l][d] = float(2 * l + 2) + for l in range(_L): + self.assertTrue(torch.all(dk[l][d] == float(2 * l + 1))) + self.assertTrue(torch.all(dv[l][d] == float(2 * l + 2))) + + def test_missing_tail_pad_fails_loud(self): + ps, num_pages = 2, 4 + raw = _make_raw(ps, num_pages, pad_pages=0) + with self.assertRaises(AssertionError): + _build_views(raw, ps, num_pages) + + def test_asymmetric_dims_rejected(self): + ps, num_pages = 2, 4 + raw = _make_raw(ps, num_pages) + with self.assertRaises(AssertionError): + _build_views(raw, ps, num_pages, head_dim=6, v_head_dim=4) + + +# ---- pool level ---- + +_N_FULL = 32 # full-attn token slots per pool in the fixtures below +_N_SWA = 16 + + +def _swa_spec(grow="up", head_dim=_D, v_head_dim=None): + return MHASubPoolSpec( + name="swa", + layer_num=_L, + head_num=_H, + head_dim=head_dim, + v_head_dim=v_head_dim, + store_dtype=_DTYPE, + grow_direction=grow, + ) + + +def _make_pool(ps=1, full_spec=None, device=_DEV): + full = full_spec if full_spec is not None else _mha_spec() + swa = _swa_spec() + total = full.entry_bytes() * _N_FULL + swa.entry_bytes() * _N_SWA + return UnifiedKVPool( + total_bytes=total, + sub_pool_specs=[full, swa], + device=device, + enable_memory_saver=False, + page_size=ps, + ) + + +class TestUnifiedKVPoolDenseViews(unittest.TestCase): + def test_every_mha_sub_pool_is_per_layer_contiguous(self): + """The unified pool has ONE MHA layout: both sub-pools come back as + stock 3-D per-layer views, whatever their page size.""" + for ps in (1, 4): + pool = _make_pool(ps=ps) + for name in ("full", "swa"): + k, v = pool.mha_views_for(name) + self.assertEqual(k[0].dim(), 3, f"{name} K at ps={ps}") + self.assertEqual(v[0].dim(), 3, f"{name} V at ps={ps}") + self.assertTrue(k[0].is_contiguous()) + + def test_tail_pad_is_derived_from_the_specs(self): + """The per-layer views hang past the last page envelope, so the pool + over-allocates one envelope of the widest sub-pool. Derived here, not + passed in, so no construction site can under-allocate it.""" + for ps in (1, 4): + kv = _make_pool(ps) + full, swa = _mha_spec(), _swa_spec() + self.assertEqual( + kv.view_tail_pad_bytes, + ps * max(full.entry_bytes(), swa.entry_bytes()), + f"tail pad at ps={ps}", + ) + self.assertEqual( + kv._raw.numel(), + full.entry_bytes() * _N_FULL + + swa.entry_bytes() * _N_SWA + + kv.view_tail_pad_bytes, + "the pad extends the allocation only", + ) + + +def _layer(l): + return SimpleNamespace(layer_id=l) + + +def _make_pool_and_kv(ps, device=_DEV): + kv = _make_pool(ps=ps, device=device) + return kv, UnifiedMHATokenToKVPool( + unified_buffer=kv, + sub_pool_name="full", + page_size=ps, + enable_alt_stream=False, + ) + + +class TestUnifiedMHATokenToKVPool(unittest.TestCase): + def test_size_is_view_row_bound(self): + """`size` drives BOTH the python OOB check and the store kernel's + device-side size_limit; it must be the view row bound, not slot count.""" + for ps in (1, 4): + unified_kv, pool_under_test = _make_pool_and_kv(ps) + n_rows = (unified_kv.max_slots("full") // ps) * _BLOCKS * ps + self.assertEqual(pool_under_test.size, n_rows - ps) + + def test_stock_write_lands_on_envelope_truth(self): + """Byte-identity: the pool's stock inherited `set_kv_buffer` at kernel-facing + locs must produce exactly the bytes that direct writes through STRIDED + views over the same envelope produce at the same (page, slot, layer) + cells. The strided views are built here purely as the independent + description of the envelope — pins the whole write path (loc -> view -> + raw bytes) end to end.""" + for ps in (1, 4): + kv, pool = _make_pool_and_kv(ps, device=_STORE_DEV) + # An independent strided view of the SAME sub-pool region. + sk, sv = _reference_strided_views( + kv._raw, + page_size=ps, + num_pages=kv.max_slots("full") // ps, + anchor_bytes=kv.anchor_bytes("full"), + ) + probes = [(1, 0), (2, ps - 1), (5, ps // 2)] + for l in range(_L): + toks = torch.tensor( + [p * ps + s for (p, s) in probes], device=_STORE_DEV + ) + kernel_locs = (toks // ps) * (ps * _BLOCKS) + toks % ps + k = torch.full( + (len(probes), _H, _D), + float(l + 1), + dtype=_DTYPE, + device=_STORE_DEV, + ) + v = torch.full( + (len(probes), _H, _D), + float(l + 101), + dtype=_DTYPE, + device=_STORE_DEV, + ) + pool.set_kv_buffer(_layer(l), kernel_locs, k, v) + for p, s in probes: + self.assertTrue( + torch.all(sk[l][p, s] == float(l + 1)), + f"K (l={l}, p={p}, s={s}, ps={ps}) not at the envelope cell", + ) + self.assertTrue( + torch.all(sv[l][p, s] == float(l + 101)), + f"V (l={l}, p={p}, s={s}, ps={ps}) not at the envelope cell", + ) + + def test_move_kv_cache_relocates_whole_envelopes(self): + """Compaction hands PHYSICAL token runs, not kernel-facing ids. The override + must relocate exactly the page envelopes those runs name — red if it is + lost, since the inherited per-layer move would apply physical ids to + the row space.""" + ps = 4 + kv, pool = _make_pool_and_kv(ps) + live = kv._raw.numel() - kv.view_tail_pad_bytes + seed = (torch.arange(live, dtype=torch.float32) % 251).to(torch.uint8) + kv._raw[:live] = seed + page_bytes = ps * _mha_spec().entry_bytes() + + src_pages, tgt_pages = torch.tensor([5, 6]), torch.tensor([2, 3]) + offs = torch.arange(ps) + run = lambda p: (p[:, None] * ps + offs).reshape(-1) + pool.move_kv_cache(run(tgt_pages), run(src_pages)) + + want = seed.clone() + for sp, tp in zip(src_pages.tolist(), tgt_pages.tolist()): + want[tp * page_bytes : (tp + 1) * page_bytes] = seed[ + sp * page_bytes : (sp + 1) * page_bytes + ] + self.assertTrue( + torch.equal(kv._raw[:live], want), + "envelope move did not relocate exactly the named pages", + ) + + def test_transfer_entry_points_fail_loud(self): + """PD / CPU-copy entry points assume per-layer buffers indexed by TOKEN + id; against the row space they would silently mis-index (or hit a + missing-attr AttributeError). Every one of them must raise.""" + _, pool = _make_pool_and_kv(1) + with self.assertRaises(NotImplementedError): + pool.get_contiguous_buf_infos() + with self.assertRaises(NotImplementedError): + pool.get_cpu_copy(torch.tensor([1])) + with self.assertRaises(NotImplementedError): + pool.load_cpu_copy(None, torch.tensor([1])) + with self.assertRaises(NotImplementedError): + pool.set_kv_buffer_prefix_valid() + + def test_hnd_env_cannot_hijack_layout(self): + """SGLANG_USE_HND_KVCACHE=1 used to flip the inherited env-driven + layout selector, putting the pool in a mode whose code paths do not + match its buffers (HND indexes 4-D; the per-layer views are 3-D). The + pinned label must win.""" + with envs.SGLANG_USE_HND_KVCACHE.override(True): + _, pool = _make_pool_and_kv(1) + self.assertFalse(pool.use_hnd) + self.assertEqual(pool.kv_cache_layout, "page_major") + + +class TestFactoryDenseViews(unittest.TestCase): + """The real SWA factory builds both sub-pools and wires the matching + kernel-facing multipliers into the composite allocator.""" + + # _swa_factory geometry: L_full = L_swa = 2, uniform 8/8 dims, ps = 1. + FULL_MULT = 4 # 2 * L_full + SWA_MULT = 4 # 2 * L_swa + + def _bundle(self): + # Self-contained tiny SWA-factory bundle (L_full = L_swa = 2, uniform + # 8/8 dims, ps = 1) — small enough that per-layer views build on CPU. + from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools + + return init_unified_swa_pools( + device="cpu", + kv_cache_dtype=torch.float16, + head_num=2, + head_dim=8, + v_head_dim=8, + swa_head_num=2, + swa_head_dim=8, + swa_v_head_dim=8, + page_size=1, + start_layer=0, + end_layer=4, + swa_attention_layer_ids=[1, 3], + full_attention_layer_ids=[0, 2], + full_max_total_num_tokens=64, + swa_max_total_num_tokens=32, + enable_memory_saver=False, + need_sort=False, + ) + + def test_factory_wires_matching_multipliers(self): + b = self._bundle() + pool = b.unified_memory_pool + alloc = b.token_to_kv_pool_allocator + self.assertEqual(alloc.kernel_page_multiplier, self.FULL_MULT) + self.assertEqual(alloc.swa_kernel_page_multiplier, self.SWA_MULT) + # Sub-pools are the dense class exposing stock 3-D per-layer views. + self.assertEqual(b.token_to_kv_pool.full_kv_pool.k_buffer[0].dim(), 3) + self.assertEqual(b.token_to_kv_pool.swa_kv_pool.k_buffer[0].dim(), 3) + self.assertGreater(pool.view_tail_pad_bytes, 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_mla_dense_block_table.py b/test/registered/unit/mem_cache/test_unified_mla_block_table.py similarity index 93% rename from test/registered/unit/mem_cache/test_unified_mla_dense_block_table.py rename to test/registered/unit/mem_cache/test_unified_mla_block_table.py index 196f265f1..e4c582fe7 100644 --- a/test/registered/unit/mem_cache/test_unified_mla_dense_block_table.py +++ b/test/registered/unit/mem_cache/test_unified_mla_block_table.py @@ -15,15 +15,15 @@ memory pool (Kimi-Linear). `req_to_token` holds VIRTUAL token ids, while the per-layer MLA views are dense -(`build_dense_mla_views`). The paged MLA backends therefore need their page-level -block table filled with DENSE page ids: +(`build_mla_views`). The paged MLA backends therefore need their page-level +block table filled with kernel-facing page ids: dense_page(virtual_page) = v2p[virtual_page] * layer_num Three backend families reach that same formula by different routes: - `create_flashmla_kv_indices_triton` in-kernel via `v2p_ptr` / `PAGE_MULT` (trtllm_mla / cutedsl_mla / tokenspeed_mla); - - the flashinfer_mla updaters, post-gathering `translate_kv_loc_dense` over the + - the flashinfer_mla updaters, post-gathering `translate_kv_loc_for_kernel` over the token-level kv_indices; - `normal_decode_set_metadata` in-kernel, for fa3's captured-decode page table. @@ -32,13 +32,13 @@ Covered here: - kernel dense mapping against the python reference, for several page sizes, ragged sequence lengths and a non-identity v2p permutation; - padded block-table lanes never index the v2p table out of bounds; - - the token-level dense translate the flashinfer updaters apply agrees with the + - the token-level kernel-facing translate the flashinfer updaters apply agrees with the page-level block table the trtllm path builds; - fa3's fused metadata kernels agree with the same reference, on both the page_size == 1 fast path (which is what Kimi-Linear takes: fa3 imposes no page-size constraint) and the general path. - python -m pytest test/registered/unit/mem_cache/test_unified_mla_dense_block_table.py -v + python -m pytest test/registered/unit/mem_cache/test_unified_mla_block_table.py -v """ import unittest @@ -145,7 +145,7 @@ class TestDenseBlockTable(unittest.TestCase): def test_single_full_attention_layer_still_maps_v2p(self): """A config with exactly ONE full-attention layer (e.g. a PP rank owning a single MLA layer) has `kernel_page_multiplier == 1`, but its req_to_token - still holds VIRTUAL ids. The dense id collapses onto the physical id, so + still holds VIRTUAL ids. The kernel-facing id collapses onto the physical id, so the v2p gather alone IS the whole translation -- it must not be skipped. Regression guard for detecting the unified pool via `multiplier > 1`: @@ -180,7 +180,7 @@ class TestDenseBlockTable(unittest.TestCase): def test_agrees_with_token_level_dense_translate(self): """The flashinfer updaters translate TOKEN ids with - `translate_kv_loc_dense`; the trtllm path builds PAGE ids in-kernel. Both + `translate_kv_loc_for_kernel`; the trtllm path builds PAGE ids in-kernel. Both must address the same dense page block.""" page_size = 64 rt, rpi, sl, v2p = self._make_batch(page_size) @@ -190,12 +190,12 @@ class TestDenseBlockTable(unittest.TestCase): for r in range(rt.shape[0]): n = int(sl[r].item()) virt_tokens = rt[r, :n].long() - # translate_kv_loc_dense's formula, applied to token ids. + # translate_kv_loc_for_kernel's formula, applied to token ids. dense_tokens = ( v2p[virt_tokens // page_size] * (page_size * _LAYERS) + virt_tokens % page_size ) - # The block-table entry scaled by page_size must be the dense id of + # The block-table entry scaled by page_size must be the kernel-facing id of # each page's first token. first_of_page = dense_tokens[::page_size] n_pages = (n + page_size - 1) // page_size @@ -330,19 +330,19 @@ class TestUnifiedMLAHookDetection(unittest.TestCase): hooks = self._probe() self.assertFalse(hooks.enabled) self.assertIsNone(hooks.v2p_page_table) - self.assertIsNone(hooks.translate_kv_loc_dense) + self.assertIsNone(hooks.translate_kv_loc_for_kernel) self.assertEqual(hooks.kernel_page_multiplier, 1) def test_multi_layer_unified_pool(self): table = torch.arange(8) hooks = self._probe( full_v2p_page_table=table, - translate_kv_loc_dense=lambda x, **kw: x, + translate_kv_loc_for_kernel=lambda x, **kw: x, kernel_page_multiplier=_LAYERS, ) self.assertTrue(hooks.enabled) self.assertIs(hooks.v2p_page_table, table) - self.assertIsNotNone(hooks.translate_kv_loc_dense) + self.assertIsNotNone(hooks.translate_kv_loc_for_kernel) self.assertEqual(hooks.kernel_page_multiplier, _LAYERS) def test_single_full_attention_layer_pool_is_still_unified(self): @@ -356,13 +356,13 @@ class TestUnifiedMLAHookDetection(unittest.TestCase): table = torch.arange(8) hooks = self._probe( full_v2p_page_table=table, - translate_kv_loc_dense=lambda x, **kw: x, + translate_kv_loc_for_kernel=lambda x, **kw: x, kernel_page_multiplier=1, ) self.assertTrue(hooks.enabled, "single-layer unified pool read as static") self.assertIs(hooks.v2p_page_table, table) - self.assertIsNotNone(hooks.translate_kv_loc_dense) - # Multiplier stays 1: dense id == physical id, so the v2p gather alone is + self.assertIsNotNone(hooks.translate_kv_loc_for_kernel) + # Multiplier stays 1: kernel-facing id == physical id, so the v2p gather alone is # the whole translation and PAGE_MULT must not scale it. self.assertEqual(hooks.kernel_page_multiplier, 1) @@ -409,7 +409,6 @@ class TestInPlaceKvIndicesTranslate(unittest.TestCase): device=_DEV, enable_memory_saver=False, page_size=page_size, - view_tail_pad_bytes=page_size * full.entry_bytes(), ) class _Stub: @@ -438,7 +437,7 @@ class TestInPlaceKvIndicesTranslate(unittest.TestCase): def test_int32_buffer_prefix_translated_tail_untouched(self): """Mirrors the updater: an int32 capture-stable buffer holding VIRTUAL - ids in [:n] gets the dense ids written back in place, narrowed to int32, + ids in [:n] gets the kernel-facing ids written back in place, narrowed to int32, with the stale tail left alone (it must never index the v2p table).""" alloc = self._allocator() virt = alloc.alloc(64) @@ -452,13 +451,13 @@ class TestInPlaceKvIndicesTranslate(unittest.TestCase): tail_before = buf[n:].clone() valid = buf[:n] - valid.copy_(alloc.translate_kv_loc_dense(valid)) + valid.copy_(alloc.translate_kv_loc_for_kernel(valid)) - expected = alloc.translate_kv_loc_dense(virt) + expected = alloc.translate_kv_loc_for_kernel(virt) self.assertEqual(buf.dtype, torch.int32) self.assertTrue( torch.equal(buf[:n].long(), expected), - "in-place translate did not land dense ids in the stable buffer", + "in-place translate did not land kernel-facing ids in the stable buffer", ) self.assertTrue( torch.equal(buf[n:], tail_before), @@ -471,8 +470,8 @@ class TestInPlaceKvIndicesTranslate(unittest.TestCase): virt = alloc.alloc(64) self.assertIsNotNone(virt) self.assertFalse( - torch.equal(alloc.translate_kv_loc_dense(virt), virt), - "dense ids coincide with virtual ids; pick a different allocation", + torch.equal(alloc.translate_kv_loc_for_kernel(virt), virt), + "kernel-facing ids coincide with virtual ids; pick a different allocation", ) diff --git a/test/registered/unit/mem_cache/test_unified_mla_gpu_parity.py b/test/registered/unit/mem_cache/test_unified_mla_gpu_parity.py index f331df3ce..a6d3cd669 100644 --- a/test/registered/unit/mem_cache/test_unified_mla_gpu_parity.py +++ b/test/registered/unit/mem_cache/test_unified_mla_gpu_parity.py @@ -11,16 +11,16 @@ # See the License for the specific language governing permissions and # limitations under the License. # ============================================================================== -"""GPU parity of the dense-view `UnifiedMLATokenToKVPool` against the stock +"""GPU parity of the per-layer-view `UnifiedMLATokenToKVPool` against the stock `MLATokenToKVPool` on real K3 MLA geometry (L=24, D=512+64). -The unified pool receives DENSE locs (dense(t) = (t//ps)*(ps*L) + t%ps); the +The unified pool receives kernel-facing locs (kernel_id(t) = (t//ps)*(ps*L) + t%ps); the reference pool receives the raw token ids. Every (layer, token) cell must hold identical bytes afterwards. Covers: - `set_mla_kv_buffer` under BOTH kernel paths — the Triton fallback (n_loc < 768) and the TMA JIT fast path (n_loc >= 768, which flattens the - buffer via `.view(shape[0], -1)`, only legal because dense views are + buffer via `.view(shape[0], -1)`, only legal because per-layer views are contiguous); - `set_kv_buffer` (combined pre-concatenated write, the Triton-backend path); - `get_mla_kv_buffer` roundtrip; @@ -48,7 +48,7 @@ _D = _LORA + _ROPE _DTYPE = torch.bfloat16 -def _dense(t: torch.Tensor, ps: int) -> torch.Tensor: +def _kernel_id(t: torch.Tensor, ps: int) -> torch.Tensor: return (t // ps) * (ps * _L) + t % ps @@ -85,7 +85,6 @@ def _make_pools(ps: int, n_tokens: int = 4096): device=_DEV, enable_memory_saver=False, page_size=ps, - view_tail_pad_bytes=ps * full.entry_bytes(), ) unified = UnifiedMLATokenToKVPool( unified_buffer=pool, @@ -118,7 +117,7 @@ def _rand_locs(max_tokens: int, ps: int, n: int) -> torch.Tensor: class TestUnifiedMLAPoolGPUParity(unittest.TestCase): def _assert_parity(self, unified, ref, locs, ps, layers=range(_L)): for l in layers: - got = unified.get_key_buffer(l)[_dense(locs, ps)] + got = unified.get_key_buffer(l)[_kernel_id(locs, ps)] want = ref.get_key_buffer(l)[locs] torch.testing.assert_close(got, want, rtol=0, atol=0) @@ -130,7 +129,7 @@ class TestUnifiedMLAPoolGPUParity(unittest.TestCase): layer = types.SimpleNamespace(layer_id=l) nope = torch.randn(n_loc, 1, _LORA, dtype=_DTYPE, device=_DEV) rope = torch.randn(n_loc, 1, _ROPE, dtype=_DTYPE, device=_DEV) - unified.set_mla_kv_buffer(layer, _dense(locs, ps), nope, rope) + unified.set_mla_kv_buffer(layer, _kernel_id(locs, ps), nope, rope) ref.set_mla_kv_buffer(layer, locs, nope, rope) torch.cuda.synchronize() self._assert_parity(unified, ref, locs, ps) @@ -156,7 +155,7 @@ class TestUnifiedMLAPoolGPUParity(unittest.TestCase): for l in (0, _L // 2, _L - 1): layer = types.SimpleNamespace(layer_id=l) k = torch.randn(n_loc, 1, _D, dtype=_DTYPE, device=_DEV) - unified.set_kv_buffer(layer, _dense(locs, ps), k, None) + unified.set_kv_buffer(layer, _kernel_id(locs, ps), k, None) ref.set_kv_buffer(layer, locs, k, None) torch.cuda.synchronize() self._assert_parity(unified, ref, locs, ps, layers=(0, _L // 2, _L - 1)) @@ -170,8 +169,8 @@ class TestUnifiedMLAPoolGPUParity(unittest.TestCase): layer = types.SimpleNamespace(layer_id=3) nope = torch.randn(n_loc, 1, _LORA, dtype=_DTYPE, device=_DEV) rope = torch.randn(n_loc, 1, _ROPE, dtype=_DTYPE, device=_DEV) - unified.set_mla_kv_buffer(layer, _dense(locs, ps), nope, rope) - got_nope, got_rope = unified.get_mla_kv_buffer(layer, _dense(locs, ps)) + unified.set_mla_kv_buffer(layer, _kernel_id(locs, ps), nope, rope) + got_nope, got_rope = unified.get_mla_kv_buffer(layer, _kernel_id(locs, ps)) torch.cuda.synchronize() torch.testing.assert_close(got_nope, nope, rtol=0, atol=0) torch.testing.assert_close(got_rope, rope, rtol=0, atol=0) @@ -188,14 +187,15 @@ class TestUnifiedMLAPoolGPUParity(unittest.TestCase): for l in range(_L): layer = types.SimpleNamespace(layer_id=l) k = torch.randn(n_loc, 1, _D, dtype=_DTYPE, device=_DEV) - unified.set_kv_buffer(layer, _dense(src_t, ps), k, None) + unified.set_kv_buffer(layer, _kernel_id(src_t, ps), k, None) before = [ - unified.get_key_buffer(l)[_dense(src_t, ps)].clone() for l in range(_L) + unified.get_key_buffer(l)[_kernel_id(src_t, ps)].clone() + for l in range(_L) ] unified.move_kv_cache(dst_t, src_t) torch.cuda.synchronize() for l in range(_L): - got = unified.get_key_buffer(l)[_dense(dst_t, ps)] + got = unified.get_key_buffer(l)[_kernel_id(dst_t, ps)] torch.testing.assert_close(got, before[l], rtol=0, atol=0) diff --git a/test/registered/unit/mem_cache/test_unified_mla_views.py b/test/registered/unit/mem_cache/test_unified_mla_views.py index 173f69dcc..32b632ba9 100644 --- a/test/registered/unit/mem_cache/test_unified_mla_views.py +++ b/test/registered/unit/mem_cache/test_unified_mla_views.py @@ -15,15 +15,15 @@ Covers, CPU-only (pure torch — no GPU / Triton kernels): - `MLASubPoolSpec` byte math; - - `build_dense_mla_views` addressing: view_l[dense(t)] must land exactly at + - `build_mla_views` addressing: view_l[kernel_id(t)] must land exactly at the page-major envelope byte offset `p*(L*ps*D) + l*(ps*D) + s*D`, the - overlapping per-layer views must not alias at equal dense ids, and the + overlapping per-layer views must not alias at equal kernel-facing ids, and the missing-tail-pad case must fail loud; - `UnifiedKVPool` MLA plumbing: `view_tail_pad_bytes` extends the allocation only, and the reserved sink floor covers the whole page-0 envelope; - `UnifiedMLATokenToKVPool`: buffer wiring, V-as-prefix-slice, and the page-envelope `move_kv_cache` (REAL physical token ids, page-major runs); - - `MultiEndedAllocator.translate_kv_loc_dense`: dense = v2p-page * (ps*L) + + - `MultiEndedAllocator.translate_kv_loc_for_kernel`: dense = v2p-page * (ps*L) + offset, tombstone clamp to the sink, `out=` contract, multiplier-1 fallback, and correctness across eager compaction. @@ -42,7 +42,7 @@ import unittest import torch from sglang.srt.mem_cache.layout.page_major import ( - build_dense_mla_views, + build_mla_views, mla_entry_bytes, ) from sglang.srt.mem_cache.multi_ended_allocator import MultiEndedAllocator @@ -98,12 +98,11 @@ def _make_unified(page_size=1, n_full_tokens=64, n_mamba_slots=8): device=_DEV, enable_memory_saver=False, page_size=page_size, - view_tail_pad_bytes=page_size * full.entry_bytes(), ) return pool, full, mamba -def _dense(t, ps, layer_num): +def _kernel_id(t, ps, layer_num): return (t // ps) * (ps * layer_num) + t % ps @@ -142,7 +141,7 @@ class TestDenseMLAViews(unittest.TestCase): for ps in (1, 4): num_pages = 6 raw, _ = self._make_raw(ps, num_pages) - views = build_dense_mla_views( + views = build_mla_views( raw, layer_num=_L, kv_cache_dim=_D, @@ -151,9 +150,9 @@ class TestDenseMLAViews(unittest.TestCase): num_pages=num_pages, ) self.assertEqual(len(views), _L) - n_dense = num_pages * _L * ps + n_rows = num_pages * _L * ps for v in views: - self.assertEqual(tuple(v.shape), (n_dense, 1, _D)) + self.assertEqual(tuple(v.shape), (n_rows, 1, _D)) # contiguous in the (row, dim) sense — .view(-1, ps, D) legality self.assertEqual(v.stride(0), _D) self.assertEqual(v.stride(2), 1) @@ -161,7 +160,7 @@ class TestDenseMLAViews(unittest.TestCase): for p, l, s in [(0, 0, 0), (1, 2, ps - 1), (4, 1, ps // 2), (5, 2, 0)]: t = p * ps + s marker = float(p * 100 + l * 10 + s + 1) - views[l][_dense(t, ps, _L)] = marker + views[l][_kernel_id(t, ps, _L)] = marker # envelope formula, in elements elem = p * (_L * ps * _D) + l * (ps * _D) + s * _D self.assertTrue( @@ -173,7 +172,7 @@ class TestDenseMLAViews(unittest.TestCase): ps = 4 num_pages = 4 raw, _ = self._make_raw(ps, num_pages) - views = build_dense_mla_views( + views = build_mla_views( raw, layer_num=_L, kv_cache_dim=_D, @@ -182,7 +181,7 @@ class TestDenseMLAViews(unittest.TestCase): num_pages=num_pages, ) t = 2 * ps + 1 # page 2, slot 1 - d = _dense(t, ps, _L) + d = _kernel_id(t, ps, _L) for l in range(_L): views[l][d] = float(l + 1) for l in range(_L): @@ -193,7 +192,7 @@ class TestDenseMLAViews(unittest.TestCase): num_pages = 4 raw, _ = self._make_raw(ps, num_pages, pad_pages=0) with self.assertRaises(AssertionError): - build_dense_mla_views( + build_mla_views( raw, layer_num=_L, kv_cache_dim=_D, @@ -279,7 +278,7 @@ class TestUnifiedMLATokenToKVPool(unittest.TestCase): # write through the views at src, expect it at dst after the move for l in range(_L): for s in range(ps): - kv_pool.kv_buffer[l][_dense(src_page * ps + s, ps, _L)] = float( + kv_pool.kv_buffer[l][_kernel_id(src_page * ps + s, ps, _L)] = float( l * ps + s + 1 ) offsets = torch.arange(ps, dtype=torch.int64) @@ -289,7 +288,7 @@ class TestUnifiedMLATokenToKVPool(unittest.TestCase): ) for l in range(_L): for s in range(ps): - got = kv_pool.kv_buffer[l][_dense(dst_page * ps + s, ps, _L)] + got = kv_pool.kv_buffer[l][_kernel_id(dst_page * ps + s, ps, _L)] self.assertTrue( torch.all(got == float(l * ps + s + 1)), f"(l={l}, s={s})" ) @@ -331,7 +330,7 @@ class TestTranslateKvLocDense(unittest.TestCase): v = alloc.alloc(8) self.assertIsNotNone(v) phys = alloc.translate_kv_loc(v) - dense = alloc.translate_kv_loc_dense(v) + dense = alloc.translate_kv_loc_for_kernel(v) self.assertTrue(torch.all(dense == phys * _L)) def test_dense_matches_formula_paged(self): @@ -340,15 +339,15 @@ class TestTranslateKvLocDense(unittest.TestCase): v = alloc.alloc(3 * ps) self.assertIsNotNone(v) phys = alloc.translate_kv_loc(v) - dense = alloc.translate_kv_loc_dense(v) + dense = alloc.translate_kv_loc_for_kernel(v) expected = (phys // ps) * (ps * _L) + phys % ps self.assertTrue(torch.all(dense == expected)) def test_tombstone_clamps_to_sink(self): alloc = self._build(ps=1) - # never-allocated virtual ids -> v2p == -1 -> dense id 0 + # never-allocated virtual ids -> v2p == -1 -> kernel-facing id 0 virt = torch.tensor([alloc.min_slot_index + 1], dtype=torch.int64) - dense = alloc.translate_kv_loc_dense(virt) + dense = alloc.translate_kv_loc_for_kernel(virt) self.assertTrue(torch.all(dense == 0)) def test_out_matches_and_aliases(self): @@ -356,14 +355,14 @@ class TestTranslateKvLocDense(unittest.TestCase): alloc = self._build(ps=ps) v = alloc.alloc(2 * ps) self.assertIsNotNone(v) - no_out = alloc.translate_kv_loc_dense(v) + no_out = alloc.translate_kv_loc_for_kernel(v) out = torch.empty_like(v) - ret = alloc.translate_kv_loc_dense(v, out=out) + ret = alloc.translate_kv_loc_for_kernel(v, out=out) self.assertIs(ret, out) self.assertTrue(torch.all(out == no_out)) # canonical in-place aliasing: translate(x, out=x) x = v.clone() - alloc.translate_kv_loc_dense(x, out=x) + alloc.translate_kv_loc_for_kernel(x, out=x) self.assertTrue(torch.all(x == no_out)) def test_multiplier_one_falls_back_to_physical(self): @@ -371,7 +370,7 @@ class TestTranslateKvLocDense(unittest.TestCase): v = alloc.alloc(4) self.assertIsNotNone(v) self.assertTrue( - torch.all(alloc.translate_kv_loc_dense(v) == alloc.translate_kv_loc(v)) + torch.all(alloc.translate_kv_loc_for_kernel(v) == alloc.translate_kv_loc(v)) ) def test_dense_follows_compaction(self): @@ -383,8 +382,8 @@ class TestTranslateKvLocDense(unittest.TestCase): alloc.free(b) # eager compaction relocates survivors phys_a = alloc.translate_kv_loc(a) phys_c = alloc.translate_kv_loc(c) - self.assertTrue(torch.all(alloc.translate_kv_loc_dense(a) == phys_a * _L)) - self.assertTrue(torch.all(alloc.translate_kv_loc_dense(c) == phys_c * _L)) + self.assertTrue(torch.all(alloc.translate_kv_loc_for_kernel(a) == phys_a * _L)) + self.assertTrue(torch.all(alloc.translate_kv_loc_for_kernel(c) == phys_c * _L)) if __name__ == "__main__": diff --git a/test/registered/unit/server_args/test_page_major_backend_allowlist.py b/test/registered/unit/server_args/test_page_major_backend_allowlist.py index 5cf9e3717..3c1752017 100644 --- a/test/registered/unit/server_args/test_page_major_backend_allowlist.py +++ b/test/registered/unit/server_args/test_page_major_backend_allowlist.py @@ -15,9 +15,9 @@ The page-major envelope K/V views are strided, which only the Triton attention kernels read. The one exception is the unified-memory MLA pool: it exposes each -layer as a DENSE contiguous view (`build_dense_mla_views`), so the paged MLA +layer as a contiguous view (`build_mla_views`), so the paged MLA backends can read it directly once their kv_indices / block tables are remapped -to dense ids -- `fa3`, `flashinfer`'s MLA backend, and `trtllm_mla` with its +to kernel-facing ids -- `fa3`, `flashinfer`'s MLA backend, and `trtllm_mla` with its `cutedsl_mla` / `tokenspeed_mla` subclasses. Pinned here so the exception cannot silently widen to a backend that has no @@ -47,6 +47,7 @@ def _accepts( unified: bool = True, linear_decode: str | None = None, linear_prefill: str | None = None, + has_asymmetric_kv: bool = False, ) -> bool: """Run just `_handle_page_major_kv_layout` against a minimal stand-in. @@ -72,7 +73,12 @@ def _accepts( sa, "_model_config", SimpleNamespace( - attention_arch=AttentionArch.MLA if use_mla else AttentionArch.MHA + attention_arch=AttentionArch.MLA if use_mla else AttentionArch.MHA, + has_asymmetric_kv=has_asymmetric_kv, + head_dim=192 if has_asymmetric_kv else 128, + v_head_dim=128, + swa_head_dim=128, + swa_v_head_dim=128, ), ) try: @@ -106,7 +112,7 @@ class TestPageMajorBackendAllowlist(unittest.TestCase): ) def test_dense_mla_backends_rejected_for_mha(self): - """The dense-view exception is MLA-only -- MHA sub-pools stay strided.""" + """The per-layer-view exception is MLA-only -- MHA sub-pools stay strided.""" for backend in self.DENSE_MLA_BACKENDS: self.assertFalse( _accepts(backend, use_mla=False), @@ -122,6 +128,40 @@ class TestPageMajorBackendAllowlist(unittest.TestCase): f"{backend} must stay rejected without --enable-unified-memory", ) + def test_plain_page_major_arm_is_gated_at_boot(self): + """The strided views were removed: --enable-page-major-kv-layout + without --enable-unified-memory must be rejected up front for EVERY + backend, Triton included, until the per-layer-view reimplementation.""" + for backend in ("triton",) + self.DENSE_MLA_BACKENDS: + for use_mla in (True, False): + self.assertFalse( + _accepts(backend, use_mla=use_mla, unified=False), + f"{backend} must be rejected on the static page-major arm", + ) + + def test_asymmetric_kv_mha_model_cannot_use_unified_memory(self): + """head_dim != v_head_dim (MiMoV2): no uniform rows, so no per-layer views + and no unified pool. The rejection is the POOL's, not a backend's, so + it must fire on every backend -- Triton included.""" + for backend in ("triton",) + self.DENSE_MLA_BACKENDS: + self.assertFalse( + _accepts(backend, use_mla=False, has_asymmetric_kv=True), + f"--enable-unified-memory + {backend} must be rejected for an " + "asymmetric-K/V model", + ) + + def test_asymmetric_dims_do_not_screen_out_mla(self): + """MLA stores one latent row per layer, so its K/V head dims never have + to agree -- and real MLA configs report them as unequal (Kimi-Linear: + head_dim 72, v_head_dim 128). Screening on `has_asymmetric_kv` alone + would lock every one of them out of the unified pool.""" + for backend in ("triton",) + self.DENSE_MLA_BACKENDS: + self.assertTrue( + _accepts(backend, use_mla=True, has_asymmetric_kv=True), + f"{backend} must stay allowed for an MLA model with asymmetric " + "K/V head dims", + ) + def test_unwired_backends_always_rejected(self): for backend in self.UNWIRED_BACKENDS: for use_mla in (True, False): @@ -131,15 +171,10 @@ class TestPageMajorBackendAllowlist(unittest.TestCase): ) def test_helion_linear_attention_is_kda_only(self): - for unified in (True, False): - for phase in ("decode", "prefill"): - kwargs = {f"linear_{phase}": "helion"} - self.assertTrue( - _accepts("triton", use_mla=True, unified=unified, **kwargs) - ) - self.assertFalse( - _accepts("triton", use_mla=False, unified=unified, **kwargs) - ) + for phase in ("decode", "prefill"): + kwargs = {f"linear_{phase}": "helion"} + self.assertTrue(_accepts("triton", use_mla=True, **kwargs)) + self.assertFalse(_accepts("triton", use_mla=False, **kwargs)) if __name__ == "__main__":