diff --git a/python/sglang/kernels/jit/csrc/inkling/causal_conv1d.cuh b/python/sglang/kernels/jit/csrc/inkling/causal_conv1d.cuh index 818470fb2..281401ae1 100644 --- a/python/sglang/kernels/jit/csrc/inkling/causal_conv1d.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/causal_conv1d.cuh @@ -175,7 +175,7 @@ struct CausalConv1dKernel { // x may be a non-contiguous row view (stride_t arbitrary) but must be // channel-contiguous. cache_mask is torch-bool (verify shape/device only). TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(x); - TensorMatcher({-1, Km1, D}).with_dtype().with_device(dev).verify(cache); + TensorMatcher({-1, Km1, D}).with_strides({-1, -1, 1}).with_dtype().with_device(dev).verify(cache); TensorMatcher({NS}).with_dtype().with_device(dev).verify(safe_idx); TensorMatcher({NS, 1, 1}).with_device(dev).verify(cache_mask); TensorMatcher({D, Wd}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(weight); diff --git a/python/sglang/kernels/jit/csrc/inkling/draft_extend_sconv.cuh b/python/sglang/kernels/jit/csrc/inkling/draft_extend_sconv.cuh index d15fdfb06..e67be4839 100644 --- a/python/sglang/kernels/jit/csrc/inkling/draft_extend_sconv.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/draft_extend_sconv.cuh @@ -114,7 +114,7 @@ struct DraftExtendSconvKernel { W1s.set_value(W1); TensorMatcher({BT, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(hidden); - TensorMatcher({-1, W1s, D}).with_dtype().with_device(dev).verify(cache); + TensorMatcher({-1, W1s, D}).with_strides({-1, -1, 1}).with_dtype().with_device(dev).verify(cache); TensorMatcher({B}).with_dtype().with_device(dev).verify(cache_indices); TensorMatcher({B}).with_dtype().with_device(dev).verify(num_accepted); RuntimeCheck(sizeof(DType) == 2, "draft_extend: bf16x2 kernel requires 16-bit dtype"); diff --git a/python/sglang/kernels/jit/csrc/inkling/fused_decode_update.cuh b/python/sglang/kernels/jit/csrc/inkling/fused_decode_update.cuh index ba605d3e7..dab83778f 100644 --- a/python/sglang/kernels/jit/csrc/inkling/fused_decode_update.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/fused_decode_update.cuh @@ -149,7 +149,7 @@ struct FusedDecodeUpdateKernel { W1s.set_value(W - 1); TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(x); - TensorMatcher({-1, W1s, D}).with_dtype().with_device(dev).verify(cache); + TensorMatcher({-1, W1s, D}).with_strides({-1, -1, 1}).with_dtype().with_device(dev).verify(cache); TensorMatcher({T}).with_dtype().with_device(dev).verify(cache_indices); TensorMatcher({T}).with_device(dev).verify(cache_mask); TensorMatcher({D, Wd}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(weight); diff --git a/python/sglang/kernels/jit/csrc/inkling/gather_scatter_sconv.cuh b/python/sglang/kernels/jit/csrc/inkling/gather_scatter_sconv.cuh index d5c91e211..82ce0ad12 100644 --- a/python/sglang/kernels/jit/csrc/inkling/gather_scatter_sconv.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/gather_scatter_sconv.cuh @@ -75,7 +75,7 @@ struct GatherScatterSconvKernel { W1s.set_value(W1); TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(hidden); - TensorMatcher({-1, W1s, D}).with_dtype().with_device(dev).verify(cache); + TensorMatcher({-1, W1s, D}).with_strides({-1, -1, 1}).with_dtype().with_device(dev).verify(cache); TensorMatcher({B, W1s}).with_dtype().with_device(dev).verify(track_idx); TensorMatcher({B}).with_device(dev).verify(mask); TensorMatcher({B}).with_dtype().with_device(dev).verify(dst); diff --git a/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh b/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh index 034a6e47f..99de47c25 100644 --- a/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh @@ -627,7 +627,7 @@ struct ArSconvNormKernel { TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(residual_out); TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(hs_out); TensorMatcher({D}).with_dtype().with_device(dev).verify(norm_weight); - TensorMatcher({-1, W1s, D}).with_dtype().with_device(dev).verify(cache); + TensorMatcher({-1, W1s, D}).with_strides({-1, -1, 1}).with_dtype().with_device(dev).verify(cache); TensorMatcher({T}).with_dtype().with_device(dev).verify(cache_indices); TensorMatcher({T}).with_device(dev).verify(cache_mask); TensorMatcher({D, Wd}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(conv_weight); @@ -752,7 +752,7 @@ struct ArSconvNormVerifyKernel { TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(residual_out); TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(hs_out); TensorMatcher({D}).with_dtype().with_device(dev).verify(norm_weight); - TensorMatcher({-1, W1s, D}).with_dtype().with_device(dev).verify(cache); + TensorMatcher({-1, W1s, D}).with_strides({-1, -1, 1}).with_dtype().with_device(dev).verify(cache); TensorMatcher({B}).with_dtype().with_device(dev).verify(cache_indices); TensorMatcher({B}).with_device(dev).verify(cache_mask); TensorMatcher({D, Wd}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(conv_weight); diff --git a/python/sglang/kernels/jit/csrc/inkling/update_sconv_cache.cuh b/python/sglang/kernels/jit/csrc/inkling/update_sconv_cache.cuh index c178ebada..9422ae9a0 100644 --- a/python/sglang/kernels/jit/csrc/inkling/update_sconv_cache.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/update_sconv_cache.cuh @@ -106,7 +106,7 @@ struct UpdateSconvCacheKernel { // x channel-contiguous (may be a non-contiguous row view); cache contiguous // [slots, W1, D]. cache_indices/qsl int32, has_state torch-bool (shape/device only). TensorMatcher({T, D}).with_strides({-1, 1}).with_dtype().with_device(dev).verify(x); - TensorMatcher({-1, W1s, D}).with_dtype().with_device(dev).verify(cache); + TensorMatcher({-1, W1s, D}).with_strides({-1, -1, 1}).with_dtype().with_device(dev).verify(cache); TensorMatcher({B}).with_dtype().with_device(dev).verify(cache_indices); TensorMatcher({B}).with_device(dev).verify(has_state); TensorMatcher({-1}).with_dtype().with_device(dev).verify(qsl); diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index f83f2c637..a81b8c069 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -530,8 +530,6 @@ class Envs: # fall back to the per-free eager compaction. Used for production # A/B and quick rollback. Default False (lazy compaction on). SGLANG_DISABLE_LAZY_COMPACTION = EnvBool(False) - # Sort the multi-ended allocator's free list after a merge (perf A/B knob). - SGLANG_SORT_FREE_LIST_AFTER_MERGE = EnvBool(False) # Periodically log lazy-compaction stats per sub-pool (observability only). SGLANG_LOG_LAZY_COMPACTION_STATS = EnvBool(False) SGLANG_LOG_LAZY_COMPACTION_STATS_INTERVAL_SEC = EnvInt(30) @@ -1724,6 +1722,9 @@ _DEPRECATED_ENVS: Dict[str, _DeprecatedEnv] = { # Superseded by the unified JIT per_token_group_quant, the default CUDA path. "SGLANG_OPT_USE_JIT_PER_TOKEN_GROUP_QUANT": _DeprecatedEnv(), "SGLANG_MASKED_GEMM_FAST_ACT": _DeprecatedEnv(), + # The unified free list is kept unsorted between flushes by design; the + # sort-after-merge A/B knob never left its off default and is gone. + "SGLANG_SORT_FREE_LIST_AFTER_MERGE": _DeprecatedEnv(), "SGLANG_OPT_SWA_EVICT_DROP_PAGE_MARGIN": _DeprecatedEnv(), # sconv-family kernels always use the CUDA-JIT ports when supported; no toggle. "SGLANG_OPT_USE_CUDA_SCONV": _DeprecatedEnv(), diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 0e73867d2..42892e8cf 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -1030,7 +1030,6 @@ class AiterAttnBackend(AttentionBackend): req_pool_indices=forward_batch.req_pool_indices, seq_lens=forward_batch.seq_lens, seq_lens_sum=None if in_capture else forward_batch.seq_lens_sum, - encoder_lens=forward_batch.encoder_lens, forward_mode=forward_batch.forward_mode, spec_info=forward_batch.spec_info, seq_lens_cpu=seq_lens_cpu, @@ -1667,7 +1666,6 @@ class AiterAttnBackend(AttentionBackend): req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 2c81c5937..34cdbfae0 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -533,7 +533,6 @@ class FlashAttentionBackend(AttentionBackend): bs=bs, req_pool_indices=req_pool_indices, seq_lens=seq_lens, - seq_lens_sum=None, encoder_lens=encoder_lens, forward_mode=forward_mode, spec_info=spec_info, @@ -579,7 +578,6 @@ class FlashAttentionBackend(AttentionBackend): bs=bs, req_pool_indices=req_pool_indices, seq_lens=seq_lens, - seq_lens_sum=forward_batch.seq_lens_sum, encoder_lens=encoder_lens, forward_mode=forward_mode, spec_info=spec_info, @@ -2720,7 +2718,6 @@ class FlashAttentionBackend(AttentionBackend): bs: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - seq_lens_sum: int, encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], diff --git a/python/sglang/srt/mem_cache/multi_ended_allocator.py b/python/sglang/srt/mem_cache/multi_ended_allocator.py index 93a6c39e2..7b89ea1be 100644 --- a/python/sglang/srt/mem_cache/multi_ended_allocator.py +++ b/python/sglang/srt/mem_cache/multi_ended_allocator.py @@ -25,7 +25,7 @@ from __future__ import annotations import inspect import logging import os -from typing import Callable, Dict, List, Optional, Set, Tuple +from typing import Callable, Dict, List, Optional, Sequence, Set, Tuple import torch from torch.profiler import record_function @@ -47,10 +47,6 @@ from sglang.srt.utils.common import get_num_new_pages, next_power_of_2 logger = logging.getLogger(__name__) -# OFF (default): cat unsorted, `_flush` sorts once. ON: sort after each cat. -_SORT_FREE_LIST_AFTER_MERGE = envs.SGLANG_SORT_FREE_LIST_AFTER_MERGE.get() - - import atexit import signal import time as _time_mod # local alias so tests can patch @@ -292,6 +288,10 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): else: self.free_virtual_ids = None self.free_group = None + # Segment frees buffer page REPRESENTATIVES here, not whole token + # ranges: `torch.cat` of the ranges destroys the per-segment shape the + # stride derivation needs, forcing the position-less dedup back on. + self.free_page_reps_group: Optional[List[torch.Tensor]] = None self._inverse_history.clear() self._free_phys_pages = torch.empty(0, dtype=torch.int64, device=self.device) self._pending_reuse.clear() @@ -452,16 +452,8 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): return None if n_drain > 0: - if _SORT_FREE_LIST_AFTER_MERGE: - if self.grow_direction == "up": - drained_t = self._free_phys_pages[:n_drain] - self._free_phys_pages = self._free_phys_pages[n_drain:] - else: - drained_t = self._free_phys_pages[-n_drain:].flip(0) - self._free_phys_pages = self._free_phys_pages[:-n_drain] - else: - drained_t = self._free_phys_pages[:n_drain] - self._free_phys_pages = self._free_phys_pages[n_drain:] + drained_t = self._free_phys_pages[:n_drain] + self._free_phys_pages = self._free_phys_pages[n_drain:] else: drained_t = None @@ -958,10 +950,17 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): # -- free with eager compaction -- - def free(self, free_index: torch.Tensor) -> None: + def free( + self, free_index: torch.Tensor, *, _pages: Optional[torch.Tensor] = None + ) -> None: """Free virtual TOKEN ids: recover virtual PAGE ids, un-map v2p/p2v, (if id-owner) recycle the page ids, trigger eager compaction. + `_pages` carries virtual PAGE ids already derived by `free_segment` + from `start_pos` arithmetic; when given, the data-dependent dedup is + skipped. Dropped on the free-group path, which has its own + representative buffer. + `free_index` is token-granular and need not be page-aligned. EAGER mode drops one `wait_stream(forward_stream)` barrier so v2p/p2v writes and the compaction move serialize with the in-flight forward. LAZY mode needs no @@ -976,7 +975,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): self.free_group.append(self._copy_for_free_group(free_index)) return if self.lazy_compaction: - self._free_lazy(free_index) + self._free_lazy(free_index, pages=_pages) return # --- EAGER path --- # Near-no-op in normal mode (sampling's CPU sync already drained @@ -986,8 +985,12 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): with record_function("MultiEndedAlloc.free.wait_stream"): torch.cuda.current_stream().wait_stream(self.forward_stream) with record_function("MultiEndedAlloc.free.v2p_lookup"): - free_v_pages = torch.unique( - free_index.detach().to(torch.int64) // self.page_size + free_v_pages = ( + _pages + if _pages is not None + else torch.unique( + free_index.detach().to(torch.int64) // self.page_size + ) ) freed_p_pages = self.virtual_to_physical[free_v_pages] with record_function("MultiEndedAlloc.free.sync_check"): @@ -996,12 +999,51 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): self._raise_stale_slot_assertion( free_v=free_v_pages, freed_p=freed_p_pages ) - self.virtual_to_physical[free_v_pages] = -1 + self.virtual_to_physical.index_fill_(0, free_v_pages, -1) if self.is_id_owner: self.free_virtual_ids = torch.cat([self.free_virtual_ids, free_v_pages]) self._compact_pending(freed_p_pages) - def _free_lazy(self, free_index: torch.Tensor) -> None: + def _page_reps_pieces( + self, free_index: torch.Tensor, start_pos: int + ) -> Tuple[torch.Tensor, ...]: + """Page-representative TOKEN slices of one kv-row segment. + + Mirrors `PagedTokenToKVPoolAllocator.free_segment`: a page's tokens sit + consecutively in the kv row, so with `start_pos` known on the host the + representatives are stride slices -- no `torch.unique`, whose + data-dependent output shape forces a device sync. + + Exact for any segment shape: a partial head page is the `[:1]` term, a + partial tail page the final stride step. + """ + ps = self.page_size + offset = start_pos % ps + if offset == 0: + return (free_index[::ps],) + return (free_index[:1], free_index[ps - offset :: ps]) + + def free_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None: + """Fixed-shape counterpart of `free()`; see `_page_reps_pieces`. + + Contract: see base; a page must be freed by only one call per group. + """ + if free_index is None or free_index.numel() == 0: + return + if self.page_size == 1: + # token == page: nothing to dedup, the plain path is already exact. + self.free(free_index) + return + pieces = self._page_reps_pieces(free_index.detach().to(torch.int64), start_pos) + if self.free_page_reps_group is None: + reps = pieces[0] if len(pieces) == 1 else torch.cat(pieces) + self.free(reps, _pages=reps // self.page_size) + else: + self.free_page_reps_group.extend(pieces) + + def _free_lazy( + self, free_index: torch.Tensor, pages: Optional[torch.Tensor] = None + ) -> None: """Lazy free path: disjoint-element scatters + ONE `torch.cat` onto `_free_phys_pages`. No sort, no boundary absorb, no watermark mutation, no D2H sync. Boundary absorption is deferred to `_flush`. @@ -1015,24 +1057,31 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): with record_function("MultiEndedAlloc._free_lazy"): with record_function("MultiEndedAlloc._free_lazy.v2p_lookup"): free_v_pages_raw = free_index.detach().to(torch.int64) - if self.page_size == 1: + if pages is not None: + # `free_segment` already derived these by stride slicing. + free_v_pages = pages + elif self.page_size == 1: free_v_pages = free_v_pages_raw else: free_v_pages = torch.unique(free_v_pages_raw // self.page_size) freed_p_pages = self.virtual_to_physical[free_v_pages] # Disjoint-element scatters — no barrier (a freed v has no live reader; # per-element scatter writes are atomic). - self.virtual_to_physical[free_v_pages] = -1 - self.physical_to_virtual[freed_p_pages] = -1 + # `index_fill_`, NOT `t[idx] = -1`: the scalar form makes torch + # materialise -1 as a CPU tensor and copy it H2D, and a pageable + # H2D copy is host-BLOCKING -- the scheduler parks behind the + # in-flight forward until the stream drains (~16 ms per free on an + # 8192-token prefill). `index_fill_` takes the scalar through the + # ATen Scalar overload: one device kernel, no host sync. + self.virtual_to_physical.index_fill_(0, free_v_pages, -1) + self.physical_to_virtual.index_fill_(0, freed_p_pages, -1) if self.is_id_owner: self.free_virtual_ids = torch.cat([self.free_virtual_ids, free_v_pages]) self._free_phys_pages = torch.cat([self._free_phys_pages, freed_p_pages]) - if _SORT_FREE_LIST_AFTER_MERGE: - self._free_phys_pages, _ = torch.sort(self._free_phys_pages) self.live_page_count -= int(freed_p_pages.shape[0]) def _release_phys_pages_batch(self, pages: torch.Tensor) -> None: - """Cat `pages` onto `_free_phys_pages` (+ optional sort). Called by `_flush` + """Cat `pages` onto `_free_phys_pages`. Called by `_flush` at END to merge event-fired compaction-srcs (`released_fired`) AFTER the trailing dst-slice, keeping `_free_phys_pages == holes_cpu` during the walk. @@ -1044,8 +1093,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): self._stats_n_release_batch += 1 with record_function("MultiEndedAlloc._release_phys_pages_batch"): self._free_phys_pages = torch.cat([self._free_phys_pages, pages]) - if _SORT_FREE_LIST_AFTER_MERGE: - self._free_phys_pages, _ = torch.sort(self._free_phys_pages) def _compact_pending(self, freed_physical_pages: torch.Tensor) -> None: """Eager compaction over the freed PHYSICAL pages: move survivors from the @@ -1221,7 +1268,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): self._stats_peak_pending_pages = max( self._stats_peak_pending_pages, cur_pending ) - sort_tag = "ON" if _SORT_FREE_LIST_AFTER_MERGE else "OFF" logger.info( f"[lazy-stats sub={self.sub_pool_name!r}] " f"free_lazy={self._stats_n_free_lazy} " @@ -1230,7 +1276,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): f"moves={self._stats_n_flush_moves} " f"abs={self._stats_n_pages_absorbed}) " f"drain={self._stats_n_drain_did_work}/{self._stats_n_drain_calls} " - f"sort={sort_tag} " f"peak_holes={self._stats_peak_free_list_len} " f"peak_pending={self._stats_peak_pending_pages} " f"cur_holes={cur_holes} cur_pending={cur_pending} " @@ -1254,7 +1299,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): self._stats_peak_pending_pages = max( self._stats_peak_pending_pages, cur_pending ) - sort_tag = "ON" if _SORT_FREE_LIST_AFTER_MERGE else "OFF" self._stats_final_emitted = True logger.info( f"[lazy-stats FINAL sub={self.sub_pool_name!r} reason={reason}] " @@ -1264,7 +1308,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): f"moves={self._stats_n_flush_moves} " f"abs={self._stats_n_pages_absorbed}) " f"drain={self._stats_n_drain_did_work}/{self._stats_n_drain_calls} " - f"sort={sort_tag} " f"peak_holes={self._stats_peak_free_list_len} " f"peak_pending={self._stats_peak_pending_pages} " f"cur_holes={cur_holes} cur_pending={cur_pending} " @@ -1312,8 +1355,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): self._stats_n_drained_pages_total += sum( t.numel() for t in ready_tensors ) - if _SORT_FREE_LIST_AFTER_MERGE: - self._free_phys_pages, _ = torch.sort(self._free_phys_pages) def maybe_drain_pending_reuse(self) -> None: """Public scheduler hook (once per step): flow fired compaction-src pages @@ -1452,8 +1493,8 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): with record_function("MultiEndedAlloc._flush"): self._drain_pending_reuse(urgent=urgent) - # Sort ASCENDING (skip if the env knob keeps the list always-sorted). - if not _SORT_FREE_LIST_AFTER_MERGE and self._free_phys_pages.numel() > 1: + # Sort ASCENDING. + if self._free_phys_pages.numel() > 1: self._free_phys_pages, _ = torch.sort(self._free_phys_pages) all_cpu = self._free_phys_pages.tolist() # one batched D2H sync @@ -1486,8 +1527,8 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): # deferred to AFTER the trailing dst-slice, keeping `_free_phys_pages` # byte-identical to `holes_cpu` for the whole walk. That invariant is # what makes the directional dst-slice correct in both directions - # (catting srcs mid-flush would chop the wrong end / scramble under - # sort=ON, leaving ghost p2v=-1 pages + double-bound dsts). Event- + # (catting srcs mid-flush would chop the wrong end, leaving ghost + # p2v=-1 pages + double-bound dsts). Event- # PENDING srcs still route to `_pending_reuse` (read-race gating). released_fired: List[torch.Tensor] = [] @@ -1655,7 +1696,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): # ONE bulk remap (single-writer on schedule_stream). self.virtual_to_physical[v_moveds_t] = dst_pages_t self.physical_to_virtual[dst_pages_t] = v_moveds_t - self.physical_to_virtual[src_pages_t] = -1 + self.physical_to_virtual.index_fill_(0, src_pages_t, -1) self._inverse_history.append((src_pages_t, dst_pages_t, v_moveds_t)) # Src disposition — ONE entry per batch. `src_pages_t` is reused as the # `_pending_reuse` GPU tensor (no second H2D at drain). @@ -1695,6 +1736,19 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): f"Caller: {callers}." ) + # -- free-group -- + + def free_group_begin(self) -> None: + super().free_group_begin() + self.free_page_reps_group = [] + + def free_group_end(self) -> None: + pending, self.free_page_reps_group = self.free_page_reps_group, None + super().free_group_end() + if pending: + reps = torch.cat(pending) + self.free(reps, _pages=reps // self.page_size) + class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): """Composite allocator for the MHA (full-attn) + Mamba hybrid pair. @@ -1764,6 +1818,7 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): # write locations are resolved in the attention metadata. self.free_group = None + self.free_page_reps_group: Optional[List[torch.Tensor]] = None # Base init left these None; we use watermark math, not free-lists. self.free_pages = torch.empty(0, dtype=torch.int64, device=device) self.release_pages = torch.empty(0, dtype=torch.int64, device=device) @@ -1962,6 +2017,47 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): self.mamba_allocator.clear() self.free_group = None + def free_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None: + """Fixed-shape counterpart of `free()`; see + `MultiEndedAllocator._page_reps_pieces`. The mamba sub-pool is + slot-granular and untouched by a token free, so only the full side + needs the representatives. + """ + if free_index is None or free_index.numel() == 0: + return + if self.page_size == 1: + self.free(free_index) + return + pieces = self.full_attn_allocator._page_reps_pieces( + free_index.detach().to(torch.int64), start_pos + ) + if self.free_page_reps_group is None: + self._release_page_reps(pieces) + else: + self.free_page_reps_group.extend(pieces) + + def _release_page_reps(self, pieces: Sequence[torch.Tensor]) -> None: + reps = pieces[0] if len(pieces) == 1 else torch.cat(tuple(pieces)) + self.full_attn_allocator.free(reps, _pages=reps // self.page_size) + self.full_attn_allocator.clear_inverse_history() + self.mamba_allocator.clear_inverse_history() + + def free_group_begin(self) -> None: + super().free_group_begin() + self.free_page_reps_group = [] + + def free_group_end(self) -> None: + pending, self.free_page_reps_group = self.free_page_reps_group, None + super().free_group_end() + if pending: + self._release_page_reps(pending) + + def clear(self) -> None: + self.full_attn_allocator.clear() + self.mamba_allocator.clear() + self.free_group = None + self.free_page_reps_group = None + # -- Lazy compaction hooks -- def set_latest_forward_done_event(self, event: Optional[torch.cuda.Event]) -> None: @@ -2098,6 +2194,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): ) self.free_group = None + self.free_page_reps_group: Optional[List[torch.Tensor]] = None # Empty (not None) for the leak checker. self.free_pages = torch.empty(0, dtype=torch.int64, device=device) self.release_pages = torch.empty(0, dtype=torch.int64, device=device) @@ -2462,10 +2559,54 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): # Paired with set_full_to_swa_mapping: shared mode has no mapping tensor. return + # -- free-group -- + + def free_group_begin(self) -> None: + super().free_group_begin() + self.free_page_reps_group = [] + + def free_group_end(self) -> None: + pending, self.free_page_reps_group = self.free_page_reps_group, None + super().free_group_end() + if pending: + self._release_page_reps(pending) + + def free_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None: + """Fixed-shape counterpart of `free()`; see + `MultiEndedAllocator._page_reps_pieces`. Both sides share one + derivation -- neither repeats the position-less dedup. + """ + if free_index is None or free_index.numel() == 0: + return + if self.page_size == 1: + self.free(free_index) + return + pieces = self.full_attn_allocator._page_reps_pieces( + free_index.detach().to(torch.int64), start_pos + ) + if self.free_page_reps_group is None: + self._release_page_reps(pieces) + else: + self.free_page_reps_group.extend(pieces) + + def _release_page_reps(self, pieces: Sequence[torch.Tensor]) -> None: + reps = pieces[0] if len(pieces) == 1 else torch.cat(tuple(pieces)) + v_pages = reps // self.page_size + # Same tombstone filter as `free`, but at PAGE granularity (page_size + # times smaller): `> 0` strict -- -1 = tombstoned, 0 = padding sink. + swa_v2p_pages = self.swa_attn_allocator.virtual_to_physical[v_pages] + live_pages = v_pages[swa_v2p_pages > 0] + if live_pages.numel() > 0: + self.swa_attn_allocator.free(live_pages * self.page_size, _pages=live_pages) + self.full_attn_allocator.free(reps, _pages=v_pages) + self.full_attn_allocator.clear_inverse_history() + self.swa_attn_allocator.clear_inverse_history() + def clear(self) -> None: self.full_attn_allocator.clear() self.swa_attn_allocator.clear() self.free_group = None + self.free_page_reps_group = None # -- Lazy compaction hooks -- diff --git a/python/sglang/srt/mem_cache/swa_memory_pool.py b/python/sglang/srt/mem_cache/swa_memory_pool.py index e934baf6d..fb05c77cd 100644 --- a/python/sglang/srt/mem_cache/swa_memory_pool.py +++ b/python/sglang/srt/mem_cache/swa_memory_pool.py @@ -209,6 +209,15 @@ class SWAKVPool(BaseSWAKVPool): else: return self.full_kv_pool.get_value_buffer(layer_id_pool) + def get_v_head_dim(self): + # The FULL side's dim, as HybridLinearKVPool.get_v_head_dim(): a caller + # asking a pool for "the" v_head_dim wants the full-attention geometry. + # `start_layer`, not 0, so pipeline parallelism (start_layer > 0) works, + # and because layer 0 need not be a full-attention layer. + return self.full_kv_pool.get_value_buffer(self.full_kv_pool.start_layer).shape[ + -1 + ] + def get_kv_buffer(self, layer_id: int): self._wait_for_layer(layer_id) layer_id_pool, is_swa_layer = self.layers_mapping[layer_id] diff --git a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py index 5f1a521bd..6e9dccfc7 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py @@ -178,6 +178,74 @@ class SWAComponent(TreeComponent): full_indices ) + def _unified_allocator(self): + """The unified SWA composite, or None when running on the static pool.""" + from sglang.srt.mem_cache.multi_ended_allocator import ( + UnifiedSWATokenToKVPoolAllocator, + ) + + allocator = self.cache.token_to_kv_pool_allocator + if isinstance(allocator, UnifiedSWATokenToKVPoolAllocator): + return allocator + return None + + def _page_pairs( + self, full_value: torch.Tensor, incoming_full_value: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + """Page ids of two token ranges that address the SAME logical tokens. + + Dedupes by FIRST OCCURRENCE with one shared mask rather than + `torch.unique`: unique sorts by id value, and allocation hands out + virtual ids in no particular order, so sorting would pair page k of one + range with an unrelated page of the other. One mask keeps the pairing + positional, hence logical. + """ + page_size = self.tree_core.page_size + kept = full_value.detach().to(torch.int64) // page_size + incoming = incoming_full_value.detach().to(torch.int64) // page_size + assert kept.numel() == incoming.numel(), ( + f"locked-full recovery needs a 1:1 token correspondence, got " + f"{kept.numel()} kept vs {incoming.numel()} incoming" + ) + starts = torch.ones_like(kept, dtype=torch.bool) + starts[1:] = kept[1:] != kept[:-1] + incoming_starts = torch.ones_like(incoming, dtype=torch.bool) + incoming_starts[1:] = incoming[1:] != incoming[:-1] + assert torch.equal(starts, incoming_starts), ( + "the two ranges break into pages at different offsets, so no " + "page-granular ownership transfer expresses the token mapping" + ) + return kept[starts], incoming[starts] + + def _transfer_swa_pages( + self, + allocator, + full_value: torch.Tensor, + incoming_full_value: torch.Tensor, + ) -> None: + """Move swa page OWNERSHIP from the incoming ids onto the node's ids. + + The static recipe re-points the node's locked full ids at the incoming + swa pages through `full_to_swa_index_mapping`. Under the unified pool + the swa sub-pool's v2p IS that mapping, so the same move is a rebind: + give the node's virtual pages the incoming pages' physical pages, then + tombstone the incoming ones. No page is allocated or freed, so no + capacity changes — only ownership does. + """ + swa = allocator.swa_attn_allocator + kept_pages, incoming_pages = self._page_pairs(full_value, incoming_full_value) + physical = swa.virtual_to_physical[incoming_pages] + # `> 0` strict: -1 = tombstoned, 0 = the padding sink. The incoming ids + # were just allocated by the in-flight request, so every page must be + # live; a violation means we would hand the node the sink and serve + # zeros, which is worth a hard failure rather than silent corruption. + assert bool( + (physical > 0).all() + ), f"incoming swa pages must all be live, got {physical.tolist()}" + swa.bind(kept_pages, physical) + swa.virtual_to_physical.index_fill_(0, incoming_pages, -1) + swa.clear_inverse_history() + def refresh_lru( self, phase: LRURefreshPhase, @@ -1280,8 +1348,24 @@ class SWAComponent(TreeComponent): alloc.set_full_to_swa_mapping(full, swa) return if isinstance(action, RecoverSWAWithLockedFull): - # Keep the locked full; remap it onto the incoming full's SWA translation, + # Keep the locked full; hand the node the INCOMING ids' swa pages, # freeing only the incoming full, then store the swa on the node. + unified = self._unified_allocator() + if unified is not None: + # No `full_to_swa_index_mapping` here: the swa sub-pool's v2p IS + # the mapping. Rebind page ownership, then free through the + # composite -- its `swa_v2p_pages > 0` filter skips the + # just-tombstoned swa side, releasing only the full one. + self._transfer_swa_pages( + unified, action.kept_full, action.incoming_full + ) + unified.free(action.incoming_full) + self.tree_core.set_component_device_value( + action.node_id, + self.component_type, + self._translate_full_to_swa(action.kept_full), + ) + return swa_value = self._translate_full_to_swa(action.incoming_full) alloc.set_full_to_swa_mapping(action.kept_full, swa_value) alloc.clear_full_to_swa_mapping(action.incoming_full) diff --git a/test/registered/unit/mem_cache/test_inkling_sconv_strided_conv_state.py b/test/registered/unit/mem_cache/test_inkling_sconv_strided_conv_state.py new file mode 100644 index 000000000..88ff13cc1 --- /dev/null +++ b/test/registered/unit/mem_cache/test_inkling_sconv_strided_conv_state.py @@ -0,0 +1,177 @@ +# 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. +# ============================================================================== +"""Inkling SConv kernels must accept a STRIDED (page-major / unified) conv-state. + +Bug regression (fixed by relaxing 7 TensorMatcher sites): every conv-state cache +matcher in ``kernels/jit/csrc/inkling/*.cuh`` used the bare +``TensorMatcher({-1, W1s, D})`` form, whose default is a hard +``view.is_contiguous()`` RuntimeCheck (``sgl_kernel/tensor.h``). The kernel +BODIES are already stride-aware — they index via ``cache.stride(0)`` / +``cache.stride(1)`` and only require the channel dim contiguous — so the matcher +was strictly stronger than the kernel's real contract. Under the unified +tri-pool the conv-state is served as a page-major envelope view (slot pitch +spans all layers), which is non-contiguous, and the matcher rejection kills the +forward. + +The fix chains ``.with_strides({-1, -1, 1})``: slot/window strides wildcarded, +channel stride pinned to 1 (the one contract the vectorized loads rely on). + +Two layers of guard: + + 1. SOURCE SCAN (CPU, always runs — the portable red/green, same precedent as + ``test_unified_free_no_host_sync.py``): every ``.verify(cache)`` matcher in + the inkling kernel sources must carry the stride relaxation. Fails the + moment a site is reverted to the contiguity-default form or a new + conv-state matcher lands without it. + + 2. FUNCTIONAL (CUDA + JIT, skipped elsewhere): drive the real + ``update_sconv_cache`` kernel with a page-major strided cache view; on + pre-fix sources this raises ``Tensor is not contiguous as expected``; + post-fix it must run AND be bit-identical to the same op on a contiguous + clone. + + python -m pytest test/registered/unit/mem_cache/test_inkling_sconv_strided_conv_state.py -v +""" + +import re +import unittest +from pathlib import Path + +import torch + +import sglang.kernels.jit as _jit_pkg +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +_KERNEL_DIR = Path(_jit_pkg.__file__).parent / "csrc" / "inkling" + +# The 7 known conv-state cache matcher sites (file -> expected count). A new +# file/site is still caught: the scan sweeps every *.cuh, and any +# `.verify(cache)` without the relaxation fails regardless of this table. +_KNOWN_SITES = { + "update_sconv_cache.cuh": 1, + "causal_conv1d.cuh": 1, + "draft_extend_sconv.cuh": 1, + "fused_decode_update.cuh": 1, + "gather_scatter_sconv.cuh": 1, + "inkling_ar_fused_decode.cuh": 2, +} + +_RELAXATION = "with_strides({-1, -1, 1})" + + +def _cache_matcher_lines(): + """Every TensorMatcher line that verifies a tensor named `cache`.""" + hits = [] + for cuh in sorted(_KERNEL_DIR.glob("*.cuh")): + for lineno, line in enumerate(cuh.read_text().splitlines(), 1): + if "TensorMatcher" in line and re.search(r"\.verify\(cache\)", line): + hits.append((cuh.name, lineno, line.strip())) + return hits + + +class TestConvStateMatchersAcceptStrided(unittest.TestCase): + def test_every_cache_matcher_carries_the_stride_relaxation(self): + hits = _cache_matcher_lines() + bad = [(f, n, l) for f, n, l in hits if _RELAXATION not in l] + self.assertEqual( + bad, + [], + msg=( + "conv-state cache matcher(s) without the stride relaxation " + f"{_RELAXATION!r} — the TensorMatcher default enforces " + "is_contiguous(), which rejects the unified/page-major " + f"conv-state view the stride-aware kernel bodies accept: {bad}" + ), + ) + + def test_all_known_sites_still_present(self): + """Completeness guard: the relaxation must not be 'fixed' by deleting + the matcher (losing shape/dtype/device verification entirely).""" + by_file = {} + for f, _, _ in _cache_matcher_lines(): + by_file[f] = by_file.get(f, 0) + 1 + for fname, expected in _KNOWN_SITES.items(): + self.assertGreaterEqual( + by_file.get(fname, 0), + expected, + msg=f"{fname}: conv-state matcher site(s) disappeared", + ) + + def test_channel_dim_stays_pinned_contiguous(self): + """The relaxation must wildcard ONLY slot/window: a fully-wildcarded + stride spec ({-1, -1, -1}) would drop the channel-contiguity contract + the vectorized state loads rely on.""" + for f, n, line in _cache_matcher_lines(): + self.assertNotIn( + "with_strides({-1, -1, -1})", + line, + msg=f"{f}:{n} wildcards the channel stride", + ) + + +@unittest.skipUnless(torch.cuda.is_available(), "needs CUDA + JIT for the real kernel") +class TestUpdateSconvCacheStridedFunctional(unittest.TestCase): + """The real kernel on a page-major strided view == on a contiguous clone. + + Red on pre-fix sources: the matcher raises + 'Tensor is not contiguous as expected' for the strided view. + """ + + _SLOTS, _LAYERS, _W1, _D = 4, 2, 3, 64 + + def _run(self, cache: torch.Tensor) -> torch.Tensor: + from sglang.kernels.ops.mamba.inkling_sconv import update_sconv_cache + + torch.manual_seed(0) + dev = cache.device + tokens = 10 + x = torch.randn(tokens, self._D, dtype=cache.dtype, device=dev) + # 2 sequences: [0:6) -> slot 1 (has state), [6:10) -> slot 3 (fresh) + cache_indices = torch.tensor([1, 3], dtype=torch.int32, device=dev) + has_initial_state = torch.tensor([True, False], device=dev) + query_start_loc = torch.tensor([0, 6, tokens], dtype=torch.int32, device=dev) + update_sconv_cache(x, cache, cache_indices, has_initial_state, query_start_loc) + return cache + + def test_strided_view_matches_contiguous(self): + dev = "cuda" + torch.manual_seed(1) + # Page-major envelope: (slots, LAYERS, W1, D); the per-layer view + # cache = env[:, 1] has stride(0) = LAYERS*W1*D != W1*D -> non-contiguous. + env = torch.randn( + self._SLOTS, + self._LAYERS, + self._W1, + self._D, + dtype=torch.bfloat16, + device=dev, + ) + strided = env[:, 1] + self.assertFalse(strided.is_contiguous(), "precondition: view is strided") + contiguous = strided.clone() + self.assertTrue(contiguous.is_contiguous()) + + out_c = self._run(contiguous) + out_s = self._run(strided) # pre-fix: matcher rejection raises here + + self.assertTrue( + torch.equal(out_s, out_c), + "strided-view kernel result differs from the contiguous reference", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py b/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py new file mode 100644 index 000000000..1d07d1b99 --- /dev/null +++ b/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py @@ -0,0 +1,319 @@ +# 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. +# ============================================================================== +"""Locked-full SWA tombstone-recovery under the unified pool (action handler). + +`RecoverSWAWithLockedFull` recovers a tombstoned SWA node whose full value is +LOCKED: the node cannot adopt the incoming request's ids wholesale, so the +static-pool recipe hands the node the INCOMING ids' swa pages, frees only their +FULL pages, and re-points the locked ids through `full_to_swa_index_mapping`. + +The unified composite has no mapping tensor — the swa sub-pool's v2p IS the +mapping — and its `set_full_to_swa_mapping` is an explicit no-op stub. The +pre-fix handler therefore raised AttributeError on `full_to_swa_index_mapping` +(and, had that line been removed, would have silently skipped the rebind while +line 1 freed swa pages the kept ids still referenced). The fix expresses the +same move as a page-ownership REBIND: bind the node's virtual pages to the +incoming pages' physical pages, tombstone the incoming ones, then free the +incoming ids through the composite — whose `swa_v2p_pages > 0` filter skips the +tombstoned swa side, releasing ONLY the full side. + +Why the recovery must succeed rather than decline (the v1 lesson, still true on +this branch): the TreeCore insert walk counts the node in `prefix_len` +regardless of component consumption, while the SWA match validator rejects a +`value is None` node — a declined recovery makes `insert` report a prefix the +follow-up `match_prefix` cannot honor, tripping +`new_prefix_len <= len(new_indices)` in `cache_unfinished_req`. + + python -m pytest test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py -v +""" + +import unittest + +import torch +from test_multi_ended_allocator import _FakeUnifiedSWAKVPool # sibling fixture + +from sglang.srt.mem_cache.multi_ended_allocator import UnifiedSWATokenToKVPoolAllocator +from sglang.srt.mem_cache.unified_cache.cache_action import RecoverSWAWithLockedFull +from sglang.srt.mem_cache.unified_cache.component_type import ComponentType +from sglang.srt.mem_cache.unified_cache.components.swa_component import SWAComponent +from sglang.srt.mem_cache.unified_memory_pool import MHASubPoolSpec, UnifiedKVPool +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=20, suite="base-a-test-cpu") + +_DEV = "cpu" +_SWA = ComponentType.SWA + + +def _build_swa_composite(n_full=64, n_swa=64): + full_spec = MHASubPoolSpec( + name="full", + layer_num=4, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="up", + ) + swa_spec = MHASubPoolSpec( + name="swa", + layer_num=2, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="down", + ) + 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, + ) + kvcache = _FakeUnifiedSWAKVPool(pool) + allocator = UnifiedSWATokenToKVPoolAllocator( + unified_buffer=pool, + kvcache=kvcache, + device=_DEV, + full_max_total_num_tokens=n_full, + swa_max_total_num_tokens=n_swa, + need_sort=False, + forward_stream=None, + ) + return allocator + + +class _StubTreeCore: + """Just what the handler touches: page_size + the device-value setter.""" + + def __init__(self, page_size=1): + self.page_size = page_size + self.set_calls = [] + + def set_component_device_value(self, node_id, component_type, value): + self.set_calls.append((node_id, component_type, value)) + + +class _Cache: + def __init__(self, allocator): + self.token_to_kv_pool_allocator = allocator + + +class _Probe(SWAComponent): + """SWAComponent wired to the real allocator and stub tree core.""" + + def __init__(self, allocator): + self.cache = _Cache(allocator) + self.tree_core = _StubTreeCore() + + +class _StaticAllocRecorder: + """Stands in for the STATIC SWATokenToKVPoolAllocator: has the mapping + tensor and a real set_full_to_swa_mapping. The handler must keep routing + static pools through the original recipe.""" + + def __init__(self, n=16): + self.full_to_swa_index_mapping = torch.arange(n, dtype=torch.int64) + self.mapping_calls = [] + self.clear_calls = [] + self.freed_full = [] + self.freed_via_inner = [] + self.full_attn_allocator = self + + def set_full_to_swa_mapping(self, full, swa): + # Honour the write like the real static allocator: the handler routes + # every mapping write THROUGH the API (never by indexing the tensor), + # so the fake must apply it for the mapping asserts to observe it. + self.mapping_calls.append((full, swa)) + self.full_to_swa_index_mapping[full.to(torch.int64)] = swa.to(torch.int64) + + def clear_full_to_swa_mapping(self, full): + self.clear_calls.append(full) + self.full_to_swa_index_mapping[full.to(torch.int64)] = 0 + + def free_full(self, indices): + self.freed_full.append(indices) + + def free(self, indices): + # The handler must not reach the inner allocator: that skips the + # free-group defer. + self.freed_via_inner.append(indices) + + def translate_loc_from_full_to_swa(self, full_indices): + return self.full_to_swa_index_mapping[full_indices.to(torch.int64)] + + +class _RecoverTestBase(unittest.TestCase): + def _probe(self): + allocator = _build_swa_composite() + self.assertIsInstance(allocator, UnifiedSWATokenToKVPoolAllocator) + return _Probe(allocator), allocator + + def _two_ranges(self, allocator, n=4): + kept = allocator.alloc(n) + incoming = allocator.alloc(n) + self.assertIsNotNone(kept) + self.assertIsNotNone(incoming) + return kept, incoming + + +class TestPagePairing(_RecoverTestBase): + def test_pairs_positionally_not_by_sorted_id(self): + """Allocation hands out virtual ids in no particular order; deduping + with `torch.unique` (which sorts) would bind the node's page k to an + unrelated incoming page — silent wrong-KV.""" + probe, _ = self._probe() + kept = torch.tensor([9, 7, 5], dtype=torch.int64) # descending + incoming = torch.tensor([2, 4, 6], dtype=torch.int64) # ascending + kept_pages, incoming_pages = probe._page_pairs(kept, incoming) + self.assertEqual(kept_pages.tolist(), [9, 7, 5]) + self.assertEqual(incoming_pages.tolist(), [2, 4, 6]) + + def test_length_mismatch_is_rejected(self): + probe, _ = self._probe() + with self.assertRaises(AssertionError): + probe._page_pairs( + torch.tensor([1, 2, 3], dtype=torch.int64), + torch.tensor([4, 5], dtype=torch.int64), + ) + + +class TestOwnershipTransfer(_RecoverTestBase): + def test_node_ids_end_up_owning_the_incoming_physical_pages(self): + probe, allocator = self._probe() + swa = allocator.swa_attn_allocator + kept, incoming = self._two_ranges(allocator) + donated = swa.virtual_to_physical[incoming.to(torch.int64)].clone() + + probe._transfer_swa_pages(allocator, kept, incoming) + + self.assertEqual( + swa.virtual_to_physical[kept.to(torch.int64)].tolist(), + donated.tolist(), + "the node's ids must now resolve to the donated physical pages", + ) + self.assertTrue( + bool((swa.virtual_to_physical[incoming.to(torch.int64)] == -1).all()), + "the incoming ids' swa side must be tombstoned", + ) + self.assertEqual( + swa.physical_to_virtual[donated].tolist(), + kept.to(torch.int64).tolist(), + "the inverse map must follow, or a later free credits the wrong id", + ) + + def test_sink_or_dead_donor_fails_loud(self): + """Handing the node the padding sink would serve zeros; refuse.""" + probe, allocator = self._probe() + kept, incoming = self._two_ranges(allocator) + allocator.free_swa(incoming) # donor no longer owns anything + with self.assertRaises(AssertionError): + probe._transfer_swa_pages(allocator, kept, incoming) + + +class TestRecoverActionHandler(_RecoverTestBase): + def test_recovery_sets_a_live_device_value_and_frees_only_the_full_side(self): + """End-to-end through apply_component_action — the pre-fix handler + raises AttributeError (`full_to_swa_index_mapping`) on this exact + call. Post-fix: the node gets a LIVE swa value, the HANDLER neither + allocates nor frees any swa page (ownership only moves), and the + incoming ids' FULL side returns to the pool.""" + probe, allocator = self._probe() + swa = allocator.swa_attn_allocator + kept, incoming = self._two_ranges(allocator) + allocator.free_swa(kept) # what eviction does when it tombstones + # Snapshot AFTER the setup traffic: the invariant under test is that + # the recovery handler itself moves ownership without moving capacity. + swa_live = swa.allocated_count() + full_avail = allocator.full_attn_allocator.available_size() + + probe.apply_component_action( + RecoverSWAWithLockedFull(node_id=7, kept_full=kept, incoming_full=incoming) + ) + + ((node_id, ct, value),) = probe.tree_core.set_calls + self.assertEqual((node_id, ct), (7, _SWA)) + self.assertEqual(len(value), len(kept)) + self.assertTrue( + bool((value > 0).all()), + "recovered value must address live swa pages, not the sink", + ) + self.assertEqual( + swa.allocated_count(), + swa_live, + "no swa page may be released or gained — ownership only moved", + ) + self.assertEqual( + allocator.full_attn_allocator.available_size(), + full_avail + len(incoming), + "the incoming ids' FULL side must come back", + ) + + def test_recovered_ids_translate_to_live_pages_not_the_sink(self): + """The tombstoned range translates to the clamped sink before the + recovery and to real pages after — recovering from the node's OWN + already-freed ids (instead of the donated ones) reintroduces the sink.""" + probe, allocator = self._probe() + kept, incoming = self._two_ranges(allocator) + allocator.free_swa(kept) + self.assertTrue( + bool((allocator.translate_loc_from_full_to_swa(kept) == 0).all()), + "precondition: a tombstoned range translates to the sink", + ) + probe.apply_component_action( + RecoverSWAWithLockedFull(node_id=1, kept_full=kept, incoming_full=incoming) + ) + self.assertTrue( + bool((allocator.translate_loc_from_full_to_swa(kept) > 0).all()), + "after recovery the node's ids must address live swa pages", + ) + + +class TestStaticPoolPathUnchanged(unittest.TestCase): + def test_static_allocator_keeps_the_mapping_recipe(self): + """A static SWA allocator (has the mapping tensor) must keep the + original recipe — the unified branch must not hijack it.""" + static = _StaticAllocRecorder() + probe = _Probe.__new__(_Probe) + probe.cache = _Cache(static) + probe.tree_core = _StubTreeCore() + + kept = torch.tensor([1, 2], dtype=torch.int64) + incoming = torch.tensor([5, 6], dtype=torch.int64) + probe.apply_component_action( + RecoverSWAWithLockedFull(node_id=3, kept_full=kept, incoming_full=incoming) + ) + + # Both mapping writes go through the allocator API -- the kept remap + # via set_full_to_swa_mapping, the incoming tombstone via + # clear_full_to_swa_mapping -- never by indexing + # `full_to_swa_index_mapping` (the tensor is absent on the unified + # composite by design). + self.assertEqual(len(static.mapping_calls), 1, "static recipe must run") + self.assertEqual(len(static.clear_calls), 1, "incoming must be tombstoned") + self.assertTrue( + bool( + (static.full_to_swa_index_mapping[incoming.to(torch.int64)] == 0).all() + ), + "incoming ids' mapping entries must be zeroed (static recipe)", + ) + # Through free_full, not the inner allocator: the latter skips the + # free-group defer. + self.assertEqual(len(static.freed_full), 1) + self.assertEqual(static.freed_via_inner, []) + ((node_id, ct, _),) = probe.tree_core.set_calls + self.assertEqual((node_id, ct), (3, _SWA)) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_swa_pool_v_head_dim.py b/test/registered/unit/mem_cache/test_swa_pool_v_head_dim.py new file mode 100644 index 000000000..a7667fe21 --- /dev/null +++ b/test/registered/unit/mem_cache/test_swa_pool_v_head_dim.py @@ -0,0 +1,124 @@ +# 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. +# ============================================================================== +"""`SWAKVPool.get_v_head_dim()` — the pool method a mambaish + hybrid-SWA +model reaches on boot. + +`TritonAttnBackend.__init__` picks its `v_head_dim` from one of three +branches, and the middle one asks the POOL: + + if sliding_window_size is not None and swa_v_head_dim != v_head_dim: + ... from model_config ... # asymmetric hybrid SWA + elif mambaish_config(model_config) is not None: + v_head_dim = token_to_kv_pool.get_v_head_dim() # <-- this one + else: + ... from get_value_buffer(start_layer) ... + +A model that is BOTH mambaish AND hybrid-SWA with MATCHING full/SWA value +head dims (Inkling-class) skips the first branch and lands in the second — +where its pool is an SWA-shaped pool, which had no `get_v_head_dim`. The +server died at backend construction with + + AttributeError: 'SWAKVPool' object has no attribute 'get_v_head_dim' + +on the STATIC pool and, identically, on `UnifiedSWAKVPool`. Neither the +mamba-hybrid pools (`HybridLinearKVPool` has the method) nor pure hybrid-SWA +models (not mambaish, so the branch is never taken) can reach it, which is +why it went unnoticed. + + python -m pytest test/registered/unit/mem_cache/test_swa_pool_v_head_dim.py -v +""" + +import inspect +import unittest + +import torch + +from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool +from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool +from sglang.srt.mem_cache.unified_memory_pool import UnifiedSWAKVPool +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=15, suite="base-a-test-cpu") + +_DEV = "cpu" +_FULL_V_HEAD_DIM = 8 +_SWA_V_HEAD_DIM = 8 # MATCHING — this is what routes Inkling into the branch + + +def _swa_pool(): + """A static SWAKVPool with the Inkling-class layer split: full and SWA + layers interleaved, layer 0 NOT a full-attention layer (which is exactly + why the backend asks the pool instead of indexing layer 0).""" + return SWAKVPool( + size=32, + size_swa=16, + page_size=1, + dtype=torch.float16, + head_num=2, + head_dim=_FULL_V_HEAD_DIM, + swa_attention_layer_ids=[0, 2], + full_attention_layer_ids=[1, 3], + device=_DEV, + enable_memory_saver=False, + ) + + +class TestSWAPoolVHeadDim(unittest.TestCase): + def test_static_pool_reports_the_full_side_value_head_dim(self): + """Red before the fix with AttributeError; the value must be the FULL + side's, since that is the geometry the caller means.""" + pool = _swa_pool() + self.assertEqual(pool.get_v_head_dim(), _FULL_V_HEAD_DIM) + + def test_answer_matches_the_full_pool_buffer_not_layer_zero(self): + """Layer 0 is an SWA layer here, so a naive `get_value_buffer(0)` + would read the SWA side. Pin that the method routes through the FULL + sub-pool at its own start_layer — the property that makes it correct + under pipeline parallelism too.""" + pool = _swa_pool() + want = pool.full_kv_pool.get_value_buffer(pool.full_kv_pool.start_layer).shape[ + -1 + ] + self.assertEqual(pool.get_v_head_dim(), want) + # And layer 0 really is the SWA side in this fixture. + _, is_swa = pool.layers_mapping[0] + self.assertTrue(is_swa, "fixture must keep layer 0 on the SWA side") + + def test_unified_swa_pool_inherits_it(self): + """`UnifiedSWAKVPool` subclasses `SWAKVPool`, so the unified tri-pool + path (mambaish + hybrid SWA in one buffer) is covered by the same + method — no second implementation to drift.""" + self.assertTrue(issubclass(UnifiedSWAKVPool, SWAKVPool)) + self.assertIs( + UnifiedSWAKVPool.get_v_head_dim, + SWAKVPool.get_v_head_dim, + "the unified pool must inherit the method, not shadow it", + ) + + def test_signature_matches_the_hybrid_linear_precedent(self): + """The backend calls this method on whichever pool it holds, so every + pool reachable from the mambaish branch must expose the SAME + zero-argument shape. `HybridLinearKVPool` is the precedent this one + mirrors; a future pool added to that branch has to match too.""" + for cls in (SWAKVPool, HybridLinearKVPool): + sig = inspect.signature(cls.get_v_head_dim) + self.assertEqual( + [p for p in sig.parameters if p != "self"], + [], + f"{cls.__name__}.get_v_head_dim must take no arguments", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py b/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py new file mode 100644 index 000000000..e9d2011c2 --- /dev/null +++ b/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py @@ -0,0 +1,306 @@ +# 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. +# ============================================================================== +"""The unified free path must not move anything across the PCIe bus. + +Two independent host syncs lived in `MultiEndedAllocator`'s free path: + + 1. Tombstone scatters written as ``t[idx] = -1``. The scalar RHS makes torch + materialise ``-1`` as a CPU tensor and copy it H2D, and a pageable H2D + copy BLOCKS the host until the stream drains. Invisible on decode-shaped + work; ~16 ms per free behind an 8192-token prefill. + + 2. `torch.unique` recovering distinct PAGE ids from freed TOKEN ids. Its + output shape is data-dependent, so it must D2H the count + (``_unique2 -> item -> _local_scalar_dense -> cudaStreamSynchronize``). + `PagedTokenToKVPoolAllocator` already solved this with `free_segment`: + a page's tokens sit consecutively in the kv row, so given `start_pos` the + page representatives are stride slices. The unified allocators simply + never implemented it and so were permanently on the syncing path. + +These tests mirror `test_paged_free_segment.py` -- the same sweep against the +`torch.unique` reference, the same free-group deferral -- because the unified +allocators now mirror that allocator's design rather than a parallel one. + + python -m pytest test/registered/unit/mem_cache/test_unified_free_no_host_sync.py -v +""" + +import ast +import inspect +import textwrap +import unittest +from unittest import mock + +import torch +from test_multi_ended_allocator import TestPagedMultiEndedAllocator as _PagedFixture + +from sglang.srt.mem_cache import multi_ended_allocator as mea +from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=20, suite="base-a-test-cpu") + +PAGE_SIZE = _PagedFixture.PAGE_SIZE + + +def _paged_allocator(lazy: bool): + """A real paged `MultiEndedAllocator` from the sibling fixture.""" + inst = _PagedFixture([m for m in dir(_PagedFixture) if m.startswith("test_")][0]) + _pool, full, _swa, _fkv, _skv = inst._build() + full.lazy_compaction = lazy + return full + + +# -------------------------------------------------------------------------- +# 1. tombstone scatters +# -------------------------------------------------------------------------- + +_TOMBSTONE_METHODS = [ + (mea.MultiEndedAllocator, "_free_lazy"), + (mea.MultiEndedAllocator, "free"), + (mea.MultiEndedAllocator, "_commit_move_batch"), +] +_TABLES = {"virtual_to_physical", "physical_to_virtual"} + + +def _scalar_index_assignments(fn): + """`self.[] = ` occurrences in fn's source. + + Slice assignments (``t[a:b] = -1``) are excluded: a slice is a view, so the + fill needs no index tensor. Tensor-valued scatters are excluded too -- only + the scalar RHS materialises a CPU value tensor. + """ + + def _is_scalar_literal(node): + # NOTE: `-1` parses as UnaryOp(USub, Constant(1)), NOT Constant. Testing + # only for Constant silently skips every negative literal -- i.e. every + # tombstone this scan exists to find. + if isinstance(node, ast.Constant): + return True + return isinstance(node, ast.UnaryOp) and isinstance(node.operand, ast.Constant) + + tree = ast.parse(textwrap.dedent(inspect.getsource(fn))) + bad = [] + for node in ast.walk(tree): + if not isinstance(node, ast.Assign) or not _is_scalar_literal(node.value): + continue + for tgt in node.targets: + if not isinstance(tgt, ast.Subscript): + continue + val = tgt.value + if not (isinstance(val, ast.Attribute) and val.attr in _TABLES): + continue + if isinstance(tgt.slice, ast.Slice): + continue + bad.append(ast.unparse(node)) + return bad + + +class TestTombstonesDoNotCrossTheBus(unittest.TestCase): + def test_no_scalar_index_assignment(self): + for cls, name in _TOMBSTONE_METHODS: + with self.subTest(method=f"{cls.__name__}.{name}"): + bad = _scalar_index_assignments(getattr(cls, name)) + self.assertEqual( + bad, + [], + msg=( + f"{cls.__name__}.{name} writes a tombstone with a scalar " + f"RHS: {bad}. That materialises -1 as a CPU tensor and " + f"copies it H2D, blocking the scheduler thread until the " + f"stream drains. Use `.index_fill_(0, idx, -1)`." + ), + ) + + def test_the_scan_detects_the_scalar_form_it_guards(self): + """Self-check. The scan is only as good as its AST matching, and it + silently missed every tombstone until `-1` was recognised as + UnaryOp(USub, Constant) rather than Constant. Pin that.""" + + def _offender(self): + self.virtual_to_physical[free_v_pages] = -1 # noqa: F821 + + self.assertEqual(len(_scalar_index_assignments(_offender)), 1) + + def test_free_paths_actually_use_index_fill(self): + """Positive form, so deleting the scatter entirely cannot pass.""" + for cls, name in _TOMBSTONE_METHODS: + with self.subTest(method=f"{cls.__name__}.{name}"): + self.assertIn("index_fill_", inspect.getsource(getattr(cls, name))) + + def test_index_fill_matches_scalar_assign_semantics(self): + """Behaviour-preserving, including the edge cases the free path hands + it: empty index, duplicate pages, full table.""" + for idx in ( + torch.tensor([], dtype=torch.int64), + torch.tensor([1, 3, 5], dtype=torch.int64), + torch.tensor([2, 2, 3], dtype=torch.int64), # duplicates + torch.arange(6, dtype=torch.int64), + ): + with self.subTest(n=int(idx.numel())): + a = torch.arange(6, dtype=torch.int64) + b = a.clone() + a[idx] = -1 + b.index_fill_(0, idx, -1) + self.assertTrue(torch.equal(a, b)) + + +# -------------------------------------------------------------------------- +# 2. free_segment: stride page extraction instead of torch.unique +# -------------------------------------------------------------------------- + + +class TestFreeSegment(unittest.TestCase): + """Mirrors `test_paged_free_segment.TestFreeSegment`.""" + + def test_matches_unique_over_alignments(self): + """Sweep (start, end) so segments cover aligned/unaligned head and + tail, a single partial page, and the full row.""" + for num_tokens in (1, PAGE_SIZE, PAGE_SIZE + 1, 3 * PAGE_SIZE - 1): + for start in range(0, num_tokens, max(1, num_tokens // 4)): + for end in (start + 1, num_tokens): + if end <= start: + continue + alloc = _paged_allocator(lazy=True) + row = alloc.alloc(3 * PAGE_SIZE) + seg = row[start:end] + expected = torch.unique(seg // PAGE_SIZE) + alloc.free_segment(seg, start_pos=start) + freed = torch.sort(alloc._free_phys_pages)[0] + with self.subTest(n=num_tokens, start=start, end=end): + # v2p is identity-ish here, so freed physical pages map + # 1:1 onto the expected virtual pages. + self.assertEqual(freed.numel(), expected.numel()) + + def test_never_calls_unique(self): + """The decisive check -- make `torch.unique` explode. A textual guard + can be fooled; this cannot.""" + for start in (0, 1, PAGE_SIZE - 1, PAGE_SIZE, PAGE_SIZE + 3): + alloc = _paged_allocator(lazy=True) + row = alloc.alloc(3 * PAGE_SIZE) + with self.subTest(start_pos=start): + with mock.patch.object( + torch, "unique", side_effect=AssertionError("sync path taken") + ): + alloc.free_segment(row[start : start + PAGE_SIZE], start_pos=start) + + def test_empty_segment_is_noop(self): + alloc = _paged_allocator(lazy=True) + before = alloc._free_phys_pages.numel() + alloc.free_segment(torch.empty(0, dtype=torch.int64), start_pos=0) + self.assertEqual(alloc._free_phys_pages.numel(), before) + + def test_page_size_one_takes_the_plain_path(self): + """token == page: nothing to dedup, so `free_segment` must not invent + a stride slice that would drop tokens.""" + alloc = _paged_allocator(lazy=True) + alloc.page_size = 1 + v = alloc.alloc(PAGE_SIZE) + n = v.numel() + alloc.free_segment(v, start_pos=0) + self.assertEqual(alloc._free_phys_pages.numel(), n) + + +class TestFreeGroupKeepsPositions(unittest.TestCase): + """Mirrors `test_paged_free_segment.test_group_defers_until_group_end`. + + Bug regression: buffering RAW tokens and `torch.cat`-ing them at + `free_group_end` destroys each segment's shape, so the merged tensor has no + recoverable page structure and falls back to `torch.unique`. Measured as 71 + of 77 `_free_lazy` calls still syncing on gpt-oss and Qwen3.5 ps=256 + (eval_429), all attributed to `free_group_end` via the decode path. The fix + buffers page REPRESENTATIVES, so the merge concatenates page ids. + """ + + def test_group_defers_until_group_end(self): + alloc = _paged_allocator(lazy=True) + row = alloc.alloc(2 * PAGE_SIZE) + before = alloc._free_phys_pages.numel() + alloc.free_group_begin() + alloc.free_segment(row, start_pos=0) + self.assertEqual( + alloc._free_phys_pages.numel(), before, "must defer inside the group" + ) + alloc.free_group_end() + self.assertEqual(alloc._free_phys_pages.numel(), before + 2) + + def test_group_end_does_not_sync(self): + """The property the whole fix exists for: a grouped segment free must + complete with `torch.unique` disabled.""" + alloc = _paged_allocator(lazy=True) + row = alloc.alloc(3 * PAGE_SIZE) + alloc.free_group_begin() + alloc.free_segment(row[:PAGE_SIZE], start_pos=0) + alloc.free_segment( + row[PAGE_SIZE + 3 : 2 * PAGE_SIZE + 3], start_pos=PAGE_SIZE + 3 + ) + with mock.patch.object( + torch, "unique", side_effect=AssertionError("sync path taken") + ): + alloc.free_group_end() + self.assertGreater(alloc._free_phys_pages.numel(), 0) + + def test_positionless_group_still_uses_the_unique_path(self): + """Plain `free()` inside a group has no position to keep, so it must + still go through the (syncing) dedup -- correctness over speed.""" + alloc = _paged_allocator(lazy=True) + row = alloc.alloc(2 * PAGE_SIZE) + alloc.free_group_begin() + alloc.free(row) + with self.assertRaises(AssertionError): + with mock.patch.object( + torch, "unique", side_effect=AssertionError("expected") + ): + alloc.free_group_end() + + +class TestEveryUnifiedAllocatorOverridesFreeSegment(unittest.TestCase): + """Completeness guard. The base `free_segment` DISCARDS `start_pos` and + calls plain `free`, so an allocator that inherits it sends every segment + free into the syncing dedup -- silently, with no error and no wrong answer, + just a stalled scheduler thread. That is exactly what happened: the SWA + composite was overridden and the Mamba composite was not, and 77 of 77 + `_free_lazy` calls on Qwen3.5 ps=256 still synced (eval_428). + """ + + def test_all_overridden(self): + for cls in ( + mea.MultiEndedAllocator, + mea.UnifiedMambaTokenToKVPoolAllocator, + mea.UnifiedSWATokenToKVPoolAllocator, + ): + with self.subTest(cls=cls.__name__): + self.assertIsNot( + cls.free_segment, + BaseTokenToKVPoolAllocator.free_segment, + msg=( + f"{cls.__name__} inherits the base `free_segment`, which " + f"discards `start_pos` -- every segment free will take the " + f"host-syncing dedup." + ), + ) + + def test_composites_buffer_reps_not_tokens_in_a_group(self): + """The group buffer must exist on every allocator that can receive a + segment free, or `free_segment` raises inside a group.""" + for cls in ( + mea.MultiEndedAllocator, + mea.UnifiedMambaTokenToKVPoolAllocator, + mea.UnifiedSWATokenToKVPoolAllocator, + ): + with self.subTest(cls=cls.__name__): + self.assertIn("free_page_reps_group", inspect.getsource(cls)) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index fd6951d03..ae9f0f889 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -7656,6 +7656,9 @@ class TestUnifiedRadixCacheActionRouting(CustomTestCase): # the incoming full's stale mapping is cleared, then its slot freed (full-only) alloc.clear_full_to_swa_mapping.assert_called_once_with(incoming_full) alloc.free_full.assert_called_once_with(incoming_full) + # Never by indexing the tensor: the unified composite has no + # `full_to_swa_index_mapping` to index into. + alloc.full_to_swa_index_mapping.__setitem__.assert_not_called() # not the inner allocator (skips the free-group defer) and not both halves alloc.full_attn_allocator.free.assert_not_called() alloc.free.assert_not_called()