diff --git a/python/sglang/srt/arg_groups/fields/memory.py b/python/sglang/srt/arg_groups/fields/memory.py index c49ce09d3..76e414d2f 100644 --- a/python/sglang/srt/arg_groups/fields/memory.py +++ b/python/sglang/srt/arg_groups/fields/memory.py @@ -82,8 +82,10 @@ class Memory(msgspec.Struct): "Replace the statically-partitioned hybrid-model pools (full-attn KV + " "SWA/Mamba state) with one byte buffer split dynamically between " "sub-pools. Requires the Triton attention / linear-attn / Mamba " - "backends; not yet compatible with PD disaggregation or speculative " - "decoding.", + "backends. PD disaggregation is supported over mooncake at equal " + "attention TP with pp=1; not yet compatible with hierarchical / " + "host-tiered KV cache, prefill cuda-graph capture, or speculative " + "decoding other than DSPARK.", ] = False enable_session_radix_cache: A[ bool, diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 82ce1ff5c..c1cf16029 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1416,7 +1416,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): decode_req.req.kv.req_pool_idx, window_start:seq_len ] window_kv_indices_swa = ( - self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa( + self.token_to_kv_pool_allocator.translate_swa_indices_for_transfer( window_kv_indices_full ) ) diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 5c75146f7..9e468c4df 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -957,6 +957,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): prefill_data_indices=prefill_kv_indices, dst_data_indices=dst_kv_indices, executor=executor, + # The unified pool registers ONE region holding every layer's K and + # V inside each page envelope. The MHA branch would half-split that + # single region into K and V halves and compute num_kv_layers = 0, + # transferring nothing at all; the flat branch addresses the region + # as-is. MLA-unified already reaches the flat branch via + # is_mla_backend, so this only adds the MHA-unified peer. + force_flat=get_memory().enable_unified_memory, src_layer_ids=self.kv_args.kv_layer_ids, dst_layer_ids=dst_layer_ids, dst_device_data_indices=dst_device_kv_indices, @@ -1598,8 +1605,16 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): dst_data_indices=np.array(dst_indices_local, dtype=np.int32), executor=executor, state_type=st, - force_flat=st - in (StateType.QSA_PENDING, StateType.QSA_COMPRESSED), + # Two independent reasons to keep the flat layout. + # QSA's per-layer list must not be half-split into K/V; + # neither must a unified sub-pool's single region, which + # holds every layer's K and V per slot envelope -- the + # MHA branch would compute zero layers and ship nothing + # (same reason as in `send_kvcache`). + force_flat=( + st in (StateType.QSA_PENDING, StateType.QSA_COMPRESSED) + or get_memory().enable_unified_memory + ), src_layer_ids=src_state_layer_ids, dst_layer_ids=dst_state_layer_ids, ) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 255ae893d..285401064 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -1328,7 +1328,7 @@ class SchedulerDisaggregationPrefillMixin: req.kv.req_pool_idx, window_start:seq_len ] window_kv_indices_swa = ( - self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa( + self.token_to_kv_pool_allocator.translate_swa_indices_for_transfer( window_kv_indices_full ) ) diff --git a/python/sglang/srt/mem_cache/allocator/hisparse.py b/python/sglang/srt/mem_cache/allocator/hisparse.py index bef9b2a43..9dd096c3b 100644 --- a/python/sglang/srt/mem_cache/allocator/hisparse.py +++ b/python/sglang/srt/mem_cache/allocator/hisparse.py @@ -346,6 +346,17 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor): return self.logical_attn_allocator.translate_loc_from_full_to_swa(kv_indices) + def translate_swa_indices_for_transfer( + self, kv_indices: torch.Tensor + ) -> torch.Tensor: + # Delegated like the read-path translate above: this composite is not a + # SWA allocator itself, so it inherits neither the default nor an + # override, and the PD payload path calls this on whatever allocator + # the scheduler holds. + return self.logical_attn_allocator.translate_swa_indices_for_transfer( + kv_indices + ) + def full_available_size(self): return min( self.logical_attn_allocator.full_available_size(), diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index 29afd8822..cfdf732fb 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -203,6 +203,19 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): assert self._kvcache.full_to_swa_index_mapping is not None return self._kvcache.translate_loc_from_full_to_swa(kv_indices) + def translate_swa_indices_for_transfer( + self, kv_indices: torch.Tensor + ) -> torch.Tensor: + """Sliding-window token ids as the PD transfer engine addresses them. + + The sibling of `translate_kv_indices_for_transfer` for the SWA state + component. On a static pool the sliding-window buffers are indexed by + the same ids the kernels use, so the read-path translate IS the answer. + A virtual-id pool must override: the transfer addresses raw bytes and + needs PHYSICAL ids, not kernel-facing ones. + """ + return self.translate_loc_from_full_to_swa(kv_indices) + def alloc(self, need_size: int): assert self.page_size == 1 if need_size > self.full_attn_allocator.available_size(): diff --git a/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py b/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py index 775e0f778..b29a09824 100644 --- a/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py +++ b/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py @@ -17,7 +17,7 @@ sub-pools of one `UnifiedKVPool`, and the tri-pool variant that adds mamba state from __future__ import annotations import logging -from typing import List, Optional, Sequence +from typing import Callable, List, Optional, Sequence, Tuple import torch from torch.profiler import record_function @@ -34,6 +34,7 @@ from sglang.srt.mem_cache.allocator.unified_sub_pool import ( _flush_deferred_free_group, _full_tokens_before_mamba_recheck, _relieve_for_alloc, + install_move_gate, ) from sglang.srt.mem_cache.unified_memory_pool import UnifiedKVPool from sglang.srt.utils.common import get_num_new_pages @@ -320,6 +321,46 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): """Page-level physical->virtual table of the full sub-pool.""" return self.full_attn_allocator.physical_to_virtual + def translate_kv_indices_for_transfer( + self, kv_indices: torch.Tensor + ) -> torch.Tensor: + """Virtual TOKEN ids -> full-sub-pool PHYSICAL token ids for the PD + transfer engine. + + PHYSICAL, not kernel-facing: the transfer registers page ENVELOPES (see + `UnifiedMHATokenToKVPool.get_contiguous_buf_infos`). Without this + override the base identity would put VIRTUAL ids on the wire, which + address real bytes and so corrupt silently rather than fail. + """ + return self.full_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64)) + + def translate_swa_indices_for_transfer( + self, kv_indices: torch.Tensor + ) -> torch.Tensor: + """Virtual TOKEN ids -> swa-sub-pool PHYSICAL token ids. + + The SWA counterpart of the above. `translate_loc_from_full_to_swa` + cannot serve here: it returns KERNEL-FACING ids (the physical page + scaled by the sub-pool's per-page block count), which index the + per-layer views, whereas the SWA state component is registered as whole + page envelopes and addressed by physical page. + """ + return self.swa_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64)) + + def _move_gate_targets(self): + """Every member a compaction gate must cover. A subclass that adds an + end overrides THIS, and every gate widens with it.""" + return (self.full_attn_allocator, self.swa_attn_allocator) + + def set_disagg_move_gate(self, gate: Callable[[], bool]) -> None: + install_move_gate( + self._move_gate_targets(), + slot="disagg_move_gate", + gate=gate, + feature="PD disaggregation", + lazy_compaction=self.lazy_compaction, + ) + def translate_kv_loc_for_kernel( self, loc: torch.Tensor, @@ -374,6 +415,52 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): self.swa_attn_allocator.alloc_with_virtual(new_virtual_pages) return v_tokens + def _extend_in_virtual_space( + self, + prefix_lens: torch.Tensor, + prefix_lens_cpu: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_cpu: torch.Tensor, + last_loc: torch.Tensor, + extend_num_tokens: int, + ) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: + """Run the full side's paged extend and report which virtual PAGES it + newly took. Returns (virtual TOKEN ids, new virtual PAGE ids), or None + when the joint capacity check cannot fund the allocation. + + Both extend entries share this; they differ only in which of those pages + the sliding-window side then binds. + """ + num_new_pages = get_num_new_pages( + seq_lens=seq_lens_cpu, + page_size=self.page_size, + prefix_lens=prefix_lens_cpu, + ) + need_tokens = num_new_pages * self.page_size + if need_tokens > self.available_size(): + if not _relieve_for_alloc(self, need_tokens): + return None + + # Snapshot the virtual PAGES the kernel will consume; clone so swa keeps + # its view after the slice is consumed. + fa = self.full_attn_allocator + new_virtual_pages = fa.free_virtual_ids[:num_new_pages].clone() + + out_indices = fa.alloc_extend( + prefix_lens, + prefix_lens_cpu, + seq_lens, + seq_lens_cpu, + last_loc, + extend_num_tokens, + num_new_pages=num_new_pages, + ) + assert out_indices is not None, ( + "UnifiedSWA: full.alloc_extend returned None after joint pre-check " + "passed — internal-state inconsistency" + ) + return out_indices, new_virtual_pages + def alloc_extend( self, prefix_lens: torch.Tensor, @@ -386,37 +473,75 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): """Paged extend; returns virtual TOKEN ids. The same virtual page maps to full- and swa-physical, so swa binds exactly what the full kernel consumed.""" with record_function("UnifiedSWAAlloc.alloc_extend"): - num_new_pages = get_num_new_pages( - seq_lens=seq_lens_cpu, - page_size=self.page_size, - prefix_lens=prefix_lens_cpu, - ) - need_tokens = num_new_pages * self.page_size - if need_tokens > self.available_size(): - if not _relieve_for_alloc(self, need_tokens): - return None - - # Snapshot the virtual PAGES the kernel will consume; clone so swa keeps - # its view after the slice is consumed. - fa = self.full_attn_allocator - new_virtual_pages = fa.free_virtual_ids[:num_new_pages].clone() - - out_indices = fa.alloc_extend( + extended = self._extend_in_virtual_space( prefix_lens, prefix_lens_cpu, seq_lens, seq_lens_cpu, last_loc, extend_num_tokens, - num_new_pages=num_new_pages, - ) - assert out_indices is not None, ( - "UnifiedSWA.alloc_extend: full.alloc_extend returned None " - "after joint pre-check passed — internal-state inconsistency" ) + if extended is None: + return None + out_indices, new_virtual_pages = extended self.swa_attn_allocator.alloc_with_virtual(new_virtual_pages) return out_indices # virtual TOKEN ids + def alloc_extend_swa_tail( + self, + prefix_lens: torch.Tensor, + prefix_lens_cpu: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_cpu: torch.Tensor, + last_loc: torch.Tensor, + extend_num_tokens: int, + swa_tail_len: int, + ) -> Optional[torch.Tensor]: + """Decode-node prealloc: full KV for the whole sequence, sliding-window + KV for the live window tail only. + + The static composite allocates the two sides independently and records + a full->swa index mapping. That is not representable here: the two + sides SHARE one virtual id space (a virtual page names a full-physical + page and, if bound, a swa-physical one), which is why + `set_full_to_swa_mapping` is a no-op on this allocator and + `translate_loc_from_full_to_swa` derives the swa id from the virtual id + instead of a table. Running the static body would call `alloc_extend` + on the swa sub-allocator, which asserts it is not the id owner. + + The tail is expressed by binding swa for the TAIL's virtual pages only. + A new page left unbound has no swa-physical page, which reads as the + sink and is skipped by `free`'s `swa_v2p_page > 0` mask -- exactly the + out-of-window state the ratchet produces via `free_swa`. + + Admission is priced at the FULL side's page count, as plain + `alloc_extend` is: pessimistic when the tail is short, but it reuses + the composite's audited joint capacity path, and the bytes actually + held still follow the tail. + """ + assert len(prefix_lens_cpu) == 1 + assert 0 <= swa_tail_len <= extend_num_tokens + with record_function("UnifiedSWAAlloc.alloc_extend_swa_tail"): + extended = self._extend_in_virtual_space( + prefix_lens, + prefix_lens_cpu, + seq_lens, + seq_lens_cpu, + last_loc, + extend_num_tokens, + ) + if extended is None: + return None + out_indices, new_virtual_pages = extended + if swa_tail_len > 0 and new_virtual_pages.numel() > 0: + tail_pages = torch.unique(out_indices[-swa_tail_len:] // self.page_size) + # Only NEW pages need binding; a tail page carried in from the + # prefix is already bound on the swa side. + to_bind = new_virtual_pages[torch.isin(new_virtual_pages, tail_pages)] + if to_bind.numel() > 0: + self.swa_attn_allocator.alloc_with_virtual(to_bind) + return out_indices # virtual TOKEN ids + def alloc_decode( self, seq_lens: torch.Tensor, @@ -769,16 +894,9 @@ class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWATokenToKVPoolAllocator): binary search; the order matches the alloc path (full takes the high band). """ fa, sa = self.full_attn_allocator, self.swa_attn_allocator - e_f, e_s = fa.entry_bytes_per_page, sa.entry_bytes_per_page + e_f = fa.entry_bytes_per_page # full is grow-down: its chain gap IS the high band. b_high = fa._current_gap_bytes() - if sa._is_frontier_transparent(): - b_low = 0 - else: - b_low = max( - 0, - sa._byte_low_frontier() - sa._chain_high_frontier_below_bytes(), - ) h_f = len(fa._free_phys_pages) if fa.lazy_compaction else 0 h_s = sa._hole_pages() r_f = fa.num_pages - fa.min_page_index - fa._allocated_pages() @@ -819,6 +937,16 @@ class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWATokenToKVPoolAllocator): hi_n = mid - 1 return lo_n * self.page_size + def _move_gate_targets(self): + """All three members. The mamba end compacts independently and its slot + envelopes move as `StateType.MAMBA`, so leaving it out of a gate would + let a conv/SSM slot relocate under an in-flight transfer.""" + return ( + self.full_attn_allocator, + self.swa_attn_allocator, + self.mamba_allocator, + ) + def _flush_targets(self): """All three members, float FIRST: its zero-copy boundary absorption must land before the deficit math prices a relocation it already covered.""" diff --git a/python/sglang/srt/mem_cache/allocator/unified_mamba.py b/python/sglang/srt/mem_cache/allocator/unified_mamba.py index 4ecaf63b7..9e5bab9b6 100644 --- a/python/sglang/srt/mem_cache/allocator/unified_mamba.py +++ b/python/sglang/srt/mem_cache/allocator/unified_mamba.py @@ -31,6 +31,7 @@ from sglang.srt.mem_cache.allocator.unified_sub_pool import ( _flush_deferred_free_group, _full_tokens_before_mamba_recheck, _relieve_for_alloc, + install_move_gate, ) from sglang.srt.mem_cache.unified_memory_pool import UnifiedKVPool from sglang.srt.runtime_context import get_parallel @@ -306,15 +307,20 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): ) return self.full_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64)) + def _move_gate_targets(self): + """Every member a compaction gate must cover. The mamba end is gated + even where its state is not itself transferred: the gate is about the + MOVER, and the two ends compact as peers.""" + return (self.full_attn_allocator, self.mamba_allocator) + def set_disagg_move_gate(self, gate: Callable[[], bool]) -> None: - """Install the PD-disaggregation move gate on both sub-allocators.""" - assert self.lazy_compaction, ( - "PD disaggregation with the unified memory pool requires lazy " - "compaction (eager free-path compaction moves pages under " - "in-flight transfers)." + install_move_gate( + self._move_gate_targets(), + slot="disagg_move_gate", + gate=gate, + feature="PD disaggregation", + lazy_compaction=self.lazy_compaction, ) - self.full_attn_allocator.disagg_move_gate = gate - self.mamba_allocator.disagg_move_gate = gate def is_slot_allocated(self, slot: int) -> bool: return self.full_attn_allocator.is_slot_allocated(slot) diff --git a/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py b/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py index 7857f4196..81f040476 100644 --- a/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py +++ b/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py @@ -239,6 +239,31 @@ def _full_tokens_before_mamba_recheck( return -(-minimum_missing_bytes * dcp_size // full_allocator.entry_bytes) +def install_move_gate( + targets, + *, + slot: str, + gate: Callable[[], bool], + feature: str, + lazy_compaction: bool, +) -> None: + """Point every member of a composite at one compaction gate. + + A gate that reaches only some members is not a weaker gate, it is no gate: + the ungated end relocates its own pages under the same in-flight transfer. + So the member list is stated once per composite (`_move_gate_targets`) and + every gate installs over it, rather than each setter naming the members it + happens to remember. + """ + assert lazy_compaction, ( + f"{feature} with the unified memory pool requires lazy compaction " + "(eager free-path compaction moves pages under in-flight transfers)." + ) + assert slot in ("disagg_move_gate", "host_transfer_move_gate"), slot + for target in targets: + setattr(target, slot, gate) + + class MultiEndedAllocator(BaseTokenToKVPoolAllocator): """Allocator for one sub-pool over a `UnifiedKVPool`.""" diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index a0ae0d8ab..d6e448158 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -155,11 +155,21 @@ def resolve_decode_retraction_backup(*, tp_worker: BaseTpWorker) -> str: ) # Host-pool retraction transfers full and sliding-window components # only, so a model with recurrent state stays on cpu_tensor. - supports_host_pool = not uses_ssm_state( - tp_worker.model_runner.model_config - ) and ( - isinstance(kv_cache, MHATokenToKVPool) - or (isinstance(kv_cache, SWAKVPool) and full_tokens_per_layer > 0) + # + # The unified pool is excluded for the same reason hierarchical cache is + # (see `handle_unified_memory_pool`): the host-transfer path indexes the + # device buffers with the ids it is handed, and under the unified pool + # those are VIRTUAL. It also cannot be sized from `kv_cache.size`, which + # is a KERNEL-FACING row count (`num_pages * 2 * layer_num * page_size`) + # rather than a token capacity -- gpt-oss-20b reports 85M "tokens" and + # asks for 418 GB of host memory per component. + supports_host_pool = ( + not uses_ssm_state(tp_worker.model_runner.model_config) + and not memory.enable_unified_memory + and ( + isinstance(kv_cache, MHATokenToKVPool) + or (isinstance(kv_cache, SWAKVPool) and full_tokens_per_layer > 0) + ) ) schedule = get_schedule() priority_preemption = ( diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 8c0203e61..7fedaaf26 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -446,24 +446,12 @@ class KVCacheConfigurator: # from one byte buffer, then return. Gated to the target worker # (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4). if get_memory().enable_unified_memory and req_to_token_pool is None: - pd_enabled = get_disagg().disaggregation_mode != "null" is_dsv4 = is_deepseek_v4(self.model_config.hf_config) # Order matters: an Inkling-class model is BOTH mambaish and # hybrid-SWA, and the mamba pair would store every SWA layer's KV at # FULL lifetime -- its branch reads the HF config's # full_attention_layer_ids, which for Inkling is ALL layers. if self.mambaish_config is not None and self.is_hybrid_swa and not is_dsv4: - if pd_enabled: - # Same limitation as the 2-pool SWA branch below: the - # tri-pool carries an SWA sub-pool, and there is no - # whole-envelope transfer scheme for it. - raise ValueError( - "--enable-unified-memory with PD disaggregation does " - "not support hybrid-SWA models yet (no whole-envelope " - "transfer scheme for the SWA sub-pool); this model " - "routes to the mamba+SWA tri-pool, which has one. Drop " - "--enable-unified-memory or run without PD." - ) bundle = self._init_unified_mamba_swa_pools( max_num_reqs=sizes.max_running_requests, full_max_total_num_tokens=sizes.full_max_total_num_tokens, @@ -471,27 +459,12 @@ class KVCacheConfigurator: unified_total_bytes=sizes.unified_total_bytes, ) elif self.mambaish_config is not None: - if pd_enabled and not self.use_mla_backend: - raise ValueError( - "--enable-unified-memory with PD disaggregation " - "currently supports only MLA hybrid-Mamba models " - "(e.g. kimi-linear); this model uses the MHA full-" - "attention pool. Drop --enable-unified-memory or run " - "without PD disaggregation." - ) bundle = self._init_unified_mamba_pools( max_num_reqs=sizes.max_running_requests, max_total_num_tokens=sizes.max_total_num_tokens, unified_total_bytes=sizes.unified_total_bytes, ) elif self.is_hybrid_swa and not is_dsv4: - if pd_enabled: - raise ValueError( - "--enable-unified-memory with PD disaggregation does " - "not support hybrid-SWA models yet (no whole-envelope " - "transfer scheme for the SWA sub-pool). Drop " - "--enable-unified-memory or run without PD." - ) bundle = self._init_unified_swa_pools( max_num_reqs=sizes.max_running_requests, full_max_total_num_tokens=sizes.full_max_total_num_tokens, @@ -835,6 +808,13 @@ class KVCacheConfigurator: unified_total_bytes=(None if self.is_draft_worker else unified_total_bytes), # bs=1 feasibility floor input (context len is already passed). sliding_window_size=self.model_config.sliding_window_size, + # Decode nodes hand out request rows to PREALLOCATED transfers on + # top of the running set; the 2-pool mamba factory takes the same. + decode_pre_alloc_size=( + get_disagg().disaggregation_decode_extra_slots + if get_disagg().disaggregation_mode == "decode" + else 0 + ), ) def _init_unified_swa_pools( @@ -863,12 +843,28 @@ class KVCacheConfigurator: extra_max_context_len = 4 if get_spec().speculative_num_draft_tokens is not None: extra_max_context_len += get_spec().speculative_num_draft_tokens - req_to_token_pool = ReqToTokenPool( - size=max_num_reqs, - max_context_len=self.model_config.context_len + extra_max_context_len, - device=self.device, - enable_memory_saver=get_exec().features.enable_memory_saver, - ) + if get_disagg().disaggregation_mode == "decode": + # A decode node hands out request rows to PREALLOCATED transfers on + # top of its running set, so it needs the extra-slot pool (and the + # `pre_alloc_size` the scheduler's invariant checker reads). Mirrors + # `_build_req_to_token_pool`'s decode branch; the mamba composite + # already takes `decode_pre_alloc_size` the same way. + from sglang.srt.disaggregation.decode import DecodeReqToTokenPool + + req_to_token_pool = DecodeReqToTokenPool( + size=max_num_reqs, + max_context_len=self.model_config.context_len + extra_max_context_len, + device=self.device, + enable_memory_saver=get_exec().features.enable_memory_saver, + pre_alloc_size=get_disagg().disaggregation_decode_extra_slots, + ) + else: + req_to_token_pool = ReqToTokenPool( + size=max_num_reqs, + max_context_len=self.model_config.context_len + extra_max_context_len, + device=self.device, + enable_memory_saver=get_exec().features.enable_memory_saver, + ) head_num = self.model_config.get_num_kv_heads( get_parallel().attn_tp_size, get_parallel().attn_dcp_size diff --git a/python/sglang/srt/mem_cache/unified_memory_pool.py b/python/sglang/srt/mem_cache/unified_memory_pool.py index 8810e5fba..cdddc6231 100644 --- a/python/sglang/srt/mem_cache/unified_memory_pool.py +++ b/python/sglang/srt/mem_cache/unified_memory_pool.py @@ -624,10 +624,19 @@ class UnifiedMHATokenToKVPool(MHATokenToKVPool): 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." - ) + """PD-transfer registration: ONE entry, the raw buffer, addressed as + ``raw_ptr + physical_page_id * page_envelope_bytes``. + + Same whole-envelope contract as `UnifiedMLATokenToKVPool`: the transfer + item is one page across ALL layers and both K and V, because the + per-layer views overlap inside the envelope and index in kernel-facing + ids. A peer must therefore build an identical spec -- enforced on the + wire by `_validate_envelope_kv_layout`. + """ + # The address formula omits the anchor; a nonzero one would mis-address. + assert self._unified_buffer.anchor_bytes(self._sub_pool_name) == 0 + raw = self._unified_buffer._raw + return [raw.data_ptr()], [raw.numel()], [self._page_bytes] def get_cpu_copy(self, indices, mamba_indices=None): raise NotImplementedError( @@ -1882,6 +1891,7 @@ def init_unified_mamba_swa_pools( lazy_compaction: bool = False, unified_total_bytes: Optional[int] = None, sliding_window_size: Optional[int] = None, + decode_pre_alloc_size: int = 0, ) -> UnifiedPoolBundle: """Build the TRI-pool unified-memory-pool stack for models with full KV + SWA KV + mamba/conv state (Inkling-class: `mambaish_config` AND @@ -2007,6 +2017,7 @@ def init_unified_mamba_swa_pools( speculative_num_draft_tokens=speculative_num_draft_tokens, enable_overlap_schedule=not disable_overlap_schedule, start_layer=start_layer, + pre_alloc_size=decode_pre_alloc_size, ) allocator = UnifiedMambaSWATokenToKVPoolAllocator( unified_buffer=shared_pool, diff --git a/test/registered/disaggregation/test_disaggregation_unified_memory.py b/test/registered/disaggregation/test_disaggregation_unified_memory.py index 78bc04b37..382f625b6 100644 --- a/test/registered/disaggregation/test_disaggregation_unified_memory.py +++ b/test/registered/disaggregation/test_disaggregation_unified_memory.py @@ -9,6 +9,9 @@ from sglang.test.server_fixtures.disaggregation_fixture import ( register_cuda_ci(est_time=236, stage="base-b", runner_config="2-gpu-large") KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random" +# Smallest in-tree GDN hybrid: MHA full attention + gated-delta-net linear +# layers, i.e. the unified pool's MHA sub-pool rather than the MLA one. +QWEN_GDN_MODEL = "Qwen/Qwen3.5-0.8B" SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"} # --attention-backend and --enable-deterministic-inference are deliberately @@ -63,5 +66,21 @@ class TestUnifiedMemoryDisaggregationChunkedPrefill(TestUnifiedMemoryDisaggregat extra_decode_args = _chunked_args +class TestUnifiedMemoryDisaggregationMHA(TestUnifiedMemoryDisaggregation): + """The MHA full-attention sub-pool over the wire. + + Kimi-Linear above exercises the MLA sub-pool, whose whole-envelope + registration has always been the one PD supports. An MHA envelope is a + different shape -- `2 * layer_num` row-blocks per page instead of + `layer_num` -- and it reaches a different branch of + `_send_kvcache_generic`: without `force_flat` the MHA branch halves the + single registered region into K and V, computes `num_kv_layers = 0` and + transfers NOTHING, which shows up as garbage decode rather than an error. + Logprob parity against a non-PD unified reference is what catches that. + """ + + model = QWEN_GDN_MODEL + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/e2e/disaggregation/test_disaggregation_unified_memory_swa.py b/test/registered/e2e/disaggregation/test_disaggregation_unified_memory_swa.py new file mode 100644 index 000000000..7a5ced975 --- /dev/null +++ b/test/registered/e2e/disaggregation/test_disaggregation_unified_memory_swa.py @@ -0,0 +1,67 @@ +"""PD disaggregation for a hybrid-SWA model on the unified memory pool. + +A hybrid-SWA model ships TWO attention components: the full-attention KV on the +ordinary `kv_data_ptrs` channel and the sliding-window KV as `StateType.SWA`. +Under `--enable-unified-memory` both are whole page envelopes into the SAME raw +buffer, distinguished only by their per-page stride, and each is addressed by +its OWN sub-pool's physical page id -- the full and SWA sides run independent +compactions, so one virtual token names two unrelated physical pages. + +That makes three ways to be silently wrong rather than loud: + * shipping virtual ids (the base `translate_kv_indices_for_transfer` is the + identity, and virtual ids address real bytes); + * shipping the SWA side's KERNEL-FACING ids, which the read path uses, in + place of its physical ones; + * letting compaction relocate a page mid-transfer, which the SWA allocator + had no `set_disagg_move_gate` to prevent. + +Logprob parity against a non-PD unified reference catches all three; GSM8K on +gpt-oss is too noisy to (single-server unified and static both score 0.570 at +200 questions, and PD runs of each span 0.540-0.610). +""" + +import unittest + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.pd_parity_kit import PDLogprobParityMixin +from sglang.test.server_fixtures.disaggregation_fixture import ( + PDDisaggregationServerBase, +) +from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE + +register_cuda_ci(est_time=1200, stage="extra-a", runner_config="2-gpu-large") + +UNIFIED_SWA_ARGS = [ + "--skip-tokenizer-init", + "--random-seed", + "1", + "--enable-unified-memory", + # gpt-oss uses attention sinks, which flashinfer does not support; triton + # reads both sub-pools' per-layer views. + "--attention-backend", + "triton", + "--mem-fraction-static", + "0.7", + "--cuda-graph-backend-decode", + "disabled", + "--cuda-graph-backend-prefill", + "disabled", +] + + +class TestUnifiedMemoryDisaggregationSWA( + PDLogprobParityMixin, PDDisaggregationServerBase +): + """1 prefill + 1 decode, both unified, vs a non-PD unified reference.""" + + model = DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE + prefill_tp_size = 1 + decode_tp_size = 1 + decode_base_gpu_id = 1 + baseline_args = UNIFIED_SWA_ARGS + extra_prefill_args = UNIFIED_SWA_ARGS + extra_decode_args = UNIFIED_SWA_ARGS + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/e2e/disaggregation/test_disaggregation_unified_memory_tri.py b/test/registered/e2e/disaggregation/test_disaggregation_unified_memory_tri.py new file mode 100644 index 000000000..3e213c019 --- /dev/null +++ b/test/registered/e2e/disaggregation/test_disaggregation_unified_memory_tri.py @@ -0,0 +1,86 @@ +"""PD disaggregation for a TRI-pool model on the unified memory pool. + +Inkling is the only in-tree architecture that is both mambaish and hybrid-SWA, +so one unified buffer carries three components with three independent +compactions -- ``[conv state (up END) | swa (FLOAT) | full (down END)]`` -- and +PD must ship all three per request: full KV on the ``kv_data_ptrs`` channel, +sliding-window KV as ``StateType.SWA``, ShortConv state as ``StateType.MAMBA`` +(via the ``req_to_token_pool`` fallback, since the KV pool here is a +``UnifiedSWAKVPool`` rather than a ``HybridLinearKVPool``). + +Two failures this pins that the 2-pool cases cannot: + + * the FLOAT sub-pool moves for reasons neither END does, so a move gate that + reaches only full and swa still lets a conv slot relocate under an + in-flight state transfer; + * ``page_size > 1`` turns on the decode node's SWA-tail prealloc, whose + static body allocates the swa side independently -- an assertion failure + against this composite's single virtual id space, and, once that is + handled, the first path that can bind the WRONG swa pages. + +Logprob parity against a non-PD unified reference is the check: the tiny +``test`` revision is undertrained, so answer quality carries no signal, but a +dropped or misaddressed component moves logprobs immediately. +""" + +import os +import unittest + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.pd_parity_kit import PDLogprobParityMixin +from sglang.test.server_fixtures.disaggregation_fixture import ( + PDDisaggregationServerBase, +) + +register_cuda_ci(est_time=900, stage="extra-a", runner_config="2-gpu-large") + +_MODEL_PATH = os.environ.get("INKLING_TEST_MODEL_PATH", "thinkingmachines/Inkling") +_MODEL_REVISION = os.environ.get("INKLING_TEST_MODEL_REVISION", "test") + +# The unified radix tree is what merges the three components into one tree. +SERVER_ENV = {"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"} + +UNIFIED_TRI_ARGS = [ + "--skip-tokenizer-init", + "--random-seed", + "1", + "--enable-unified-memory", + # Unified requires the Triton strided page-major read/write paths. + "--attention-backend", + "triton", + "--page-size", + "128", + "--mamba-radix-cache-strategy", + "extra_buffer", + "--swa-full-tokens-ratio", + "0.1", + "--mamba-full-memory-ratio", + "0.1", + "--mem-fraction-static", + "0.5", + # Inkling defaults to a FULL prefill graph, which unified rejects at boot. + "--cuda-graph-backend-prefill", + "disabled", + "--revision", + _MODEL_REVISION, +] + + +class TestUnifiedMemoryDisaggregationTriPool( + PDLogprobParityMixin, PDDisaggregationServerBase +): + """1 prefill + 1 decode, both unified, vs a non-PD unified reference.""" + + model = _MODEL_PATH + extra_prefill_env = SERVER_ENV + extra_decode_env = SERVER_ENV + prefill_tp_size = 1 + decode_tp_size = 1 + decode_base_gpu_id = 1 + baseline_args = UNIFIED_TRI_ARGS + extra_prefill_args = UNIFIED_TRI_ARGS + extra_decode_args = UNIFIED_TRI_ARGS + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/disaggregation/test_unified_memory_move_gate.py b/test/registered/unit/disaggregation/test_unified_memory_move_gate.py index 362e8b988..e702a2595 100644 --- a/test/registered/unit/disaggregation/test_unified_memory_move_gate.py +++ b/test/registered/unit/disaggregation/test_unified_memory_move_gate.py @@ -186,5 +186,210 @@ class TestMoveGateRejectsNonPdNode(CustomTestCase): unified_memory_disagg_move_gate(scheduler) +class TestUnifiedAllocatorsPublishTheTransferContract(CustomTestCase): + """Every unified composite allocator must OVERRIDE the two PD hooks. + + `BaseTokenToKVPoolAllocator.translate_kv_indices_for_transfer` is the + IDENTITY, and `set_disagg_move_gate` exists only where a composite defines + it. Inheriting either is silent, not loud: identity puts VIRTUAL ids on the + wire (they address real bytes, so the peer gets plausible garbage), and a + missing gate lets lazy compaction relocate pages under in-flight RDMA. + An AST-level check because instantiating these composites needs a GPU. + """ + + # Composites that own the full-side virtual ids and so must define the + # transfer translate themselves. + _COMPOSITES = ( + "UnifiedMambaTokenToKVPoolAllocator", + "UnifiedSWATokenToKVPoolAllocator", + ) + # Every composite must define the gate setter, including the tri-pool, + # which inherits the SWA translates (same full side) but has a THIRD + # member the 2-pool setter does not reach. + _GATE_COMPOSITES = _COMPOSITES + ("UnifiedMambaSWATokenToKVPoolAllocator",) + + @staticmethod + def _own_methods(cls_name: str) -> Set[str]: + """Names this class defines ITSELF, inheritance excluded. + + Resolved off the class object rather than by parsing a named module: + these composites have already been moved once (out of + `multi_ended_allocator` into `allocator/unified_*`), and a hardcoded + module path turns that kind of move into a test failure that says + nothing about the contract. `__dict__` needs no GPU -- it is the class + body, not an instance. + """ + from sglang.srt.mem_cache.allocator import ( + unified_hybrid_swa, + unified_mamba, + ) + + for mod in (unified_mamba, unified_hybrid_swa): + cls = getattr(mod, cls_name, None) + if cls is not None: + return set(vars(cls)) + raise AssertionError(f"class {cls_name} not found in the unified allocators") + + def test_transfer_translate_is_not_inherited_identity(self): + for name in self._COMPOSITES: + with self.subTest(composite=name): + self.assertIn( + "translate_kv_indices_for_transfer", + self._own_methods(name), + f"{name} inherits the identity transfer translate; PD would " + "ship VIRTUAL ids and corrupt KV without any error", + ) + + # Every sub-allocator attribute a composite can hold. The stub carries all + # of them regardless of composite, so the assertion is on what installation + # REACHES rather than on what the stub was given. + _MEMBER_ATTRS = ("full_attn_allocator", "swa_attn_allocator", "mamba_allocator") + + # The members each composite's gate must reach. The tri-pool row is the one + # that matters: it inherits the setter, so an enumeration written inside + # that setter would silently leave the third member ungated. + _EXPECTED_COVERAGE = { + "UnifiedMambaTokenToKVPoolAllocator": { + "full_attn_allocator", + "mamba_allocator", + }, + "UnifiedSWATokenToKVPoolAllocator": { + "full_attn_allocator", + "swa_attn_allocator", + }, + "UnifiedMambaSWATokenToKVPoolAllocator": { + "full_attn_allocator", + "swa_attn_allocator", + "mamba_allocator", + }, + } + + def _members_reached(self, cls_name: str, slot: str) -> Set[str]: + """Install one gate on a stub composite and report which members got it. + + `object.__new__` skips `__init__` (which needs a GPU); the setter reads + only `lazy_compaction` and the member attributes. + """ + from sglang.srt.mem_cache.allocator import unified_hybrid_swa, unified_mamba + + cls = getattr(unified_mamba, cls_name, None) or getattr( + unified_hybrid_swa, cls_name + ) + alloc = object.__new__(cls) + alloc.lazy_compaction = True + for attr in self._MEMBER_ATTRS: + member = type("_Member", (), {})() + member.disagg_move_gate = None + member.host_transfer_move_gate = None + setattr(alloc, attr, member) + + def gate() -> bool: + return True + + alloc.set_disagg_move_gate(gate) + return { + attr + for attr in self._MEMBER_ATTRS + if getattr(getattr(alloc, attr), slot) is gate + } + + def test_the_gate_reaches_every_member(self): + """A gate that reaches only some members is not a weaker gate, it is no + gate: the ungated end relocates its own pages under the very transfer + the gate was installed for. + """ + for name, expected in self._EXPECTED_COVERAGE.items(): + with self.subTest(composite=name): + self.assertEqual( + self._members_reached(name, "disagg_move_gate"), + expected, + f"{name}.disagg_move_gate does not cover every member", + ) + + def test_gate_setters_do_not_enumerate_members_themselves(self): + """The structural half of the rule above: a setter that names its + members is one a new member silently escapes. Installation must go + through the shared helper, which drives off `_move_gate_targets`. + """ + import inspect + + from sglang.srt.mem_cache.allocator import unified_hybrid_swa, unified_mamba + + for name in self._EXPECTED_COVERAGE: + cls = getattr(unified_mamba, name, None) or getattr( + unified_hybrid_swa, name + ) + if "set_disagg_move_gate" not in vars(cls): + continue # inherited, and the inherited one is checked above + with self.subTest(composite=name): + body = inspect.getsource(cls.set_disagg_move_gate) + self.assertIn("install_move_gate", body) + self.assertNotIn("_move_gate = ", body) + + def test_swa_composite_translates_the_swa_side_separately(self): + """The SWA sub-pool runs its OWN compaction, so a full-side physical id + does not name the SWA page holding the same virtual token. The read-path + `translate_loc_from_full_to_swa` cannot stand in either: it returns + kernel-facing ids, and the transfer addresses raw page envelopes.""" + self.assertIn( + "translate_swa_indices_for_transfer", + self._own_methods("UnifiedSWATokenToKVPoolAllocator"), + ) + + +class TestEverySwaAllocatorAnswersTheTransferTranslate(CustomTestCase): + """Any allocator with a full->SWA read translate needs the transfer sibling. + + `_swa_payload` on both PD sides calls + `translate_swa_indices_for_transfer` on whatever allocator the scheduler + holds. Most get it by inheriting `SWATokenToKVPoolAllocator`, but a + composite that merely DELEGATES the read translate (the DSV4 HiSparse + allocator derives from `BaseTokenToKVPoolAllocator`) inherits neither the + default nor an override, and PD aborts with an AttributeError the moment a + sliding-window payload is built. + + Derived from the live class tree rather than a hand-kept list: a list would + pass forever the day someone adds the next delegating composite. + """ + + @staticmethod + def _allocator_classes(): + import importlib + import inspect + import pkgutil + + import sglang.srt.mem_cache.allocator as pkg + from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator + + found = {} + for mod_info in pkgutil.iter_modules(pkg.__path__): + try: + mod = importlib.import_module( + f"sglang.srt.mem_cache.allocator.{mod_info.name}" + ) + except Exception: + continue # optional backends need hardware this runner may lack + for _, cls in inspect.getmembers(mod, inspect.isclass): + if issubclass(cls, BaseTokenToKVPoolAllocator): + found[cls.__name__] = cls + return found + + def test_read_translate_implies_transfer_translate(self): + classes = self._allocator_classes() + # Guard the guard: an import failure that empties this set would make + # the assertion below vacuous. + self.assertIn("SWATokenToKVPoolAllocator", classes) + for name, cls in sorted(classes.items()): + if not hasattr(cls, "translate_loc_from_full_to_swa"): + continue + with self.subTest(allocator=name): + self.assertTrue( + hasattr(cls, "translate_swa_indices_for_transfer"), + f"{name} translates full->SWA for reads but cannot answer " + "translate_swa_indices_for_transfer; PD's _swa_payload " + "calls it on whatever allocator the scheduler holds", + ) + + if __name__ == "__main__": unittest.main() 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 396168eb0..8ecfb9676 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,9 +18,11 @@ import unittest import torch from sglang.srt.mem_cache.layout.page_major import ( + build_mha_views, build_mla_views, build_page_major_mamba_views, mamba_entry_bytes, + mha_entry_bytes, mla_entry_bytes, ) from sglang.test.ci.ci_register import register_cpu_ci @@ -74,6 +76,118 @@ class TestMLAEnvelopeTransferAddressing(CustomTestCase): self.assertTrue(torch.equal(got, val), (page, layer, off)) +class TestMHAEnvelopeTransferAddressing(CustomTestCase): + """The MHA counterpart of the MLA case above. + + An MHA page envelope holds ``2 * layer_num`` row-blocks (layer l's K at + block 2l, its V at 2l+1). PD ships that whole envelope as one item, so a + row written through ANY per-layer view must land inside its own page's + ``page_envelope_bytes`` block -- otherwise the transfer would carry a + page's K but another page's V and every kernel would still read fine + locally. + """ + + def test_page_envelope_matches_per_layer_views(self): + layer_num, page_size, head_num, head_dim, num_pages = 3, 4, 2, 8, 6 + store_dtype = torch.bfloat16 + entry_bytes = mha_entry_bytes( + layer_num=layer_num, + head_num=head_num, + head_dim=head_dim, + v_head_dim=head_dim, + itemsize=store_dtype.itemsize, + ) + page_bytes = page_size * entry_bytes + row_bytes = head_num * head_dim * store_dtype.itemsize + self.assertEqual(page_bytes, page_size * 2 * layer_num * row_bytes) + + # One page envelope of tail pad, as UnifiedKVPool allocates for MHA. + raw = torch.zeros((num_pages + 1) * page_bytes, dtype=torch.uint8) + k_views, v_views = build_mha_views( + raw, + layer_num=layer_num, + head_num=head_num, + head_dim=head_dim, + v_head_dim=head_dim, + store_dtype=store_dtype, + page_size=page_size, + num_pages=num_pages, + anchor_bytes=0, + ) + + blocks = 2 * layer_num + for page in range(num_pages): + for layer in range(layer_num): + for is_v, views in ((0, k_views), (1, v_views)): + for pos in range(page_size): + row = page * blocks * page_size + pos + views[layer][row].fill_(1) + (nz,) = torch.nonzero(raw, as_tuple=True) + lo, hi = int(nz.min()), int(nz.max()) + self.assertGreaterEqual( + lo, + page * page_bytes, + f"page={page} layer={layer} v={is_v} pos={pos} " + "wrote below its page envelope", + ) + self.assertLess( + hi, + (page + 1) * page_bytes, + f"page={page} layer={layer} v={is_v} pos={pos} " + "wrote past its page envelope", + ) + views[layer][row].zero_() + + def test_envelope_move_is_a_whole_page_copy(self): + """Relocating a page envelope must move every layer's K and V with it; + this is what `UnifiedMHATokenToKVPool.move_kv_cache` relies on and what + makes a physical page id a valid PD transfer index after compaction.""" + layer_num, page_size, head_num, head_dim, num_pages = 2, 2, 1, 4, 4 + store_dtype = torch.bfloat16 + entry_bytes = mha_entry_bytes( + layer_num=layer_num, + head_num=head_num, + head_dim=head_dim, + v_head_dim=head_dim, + itemsize=store_dtype.itemsize, + ) + page_bytes = page_size * entry_bytes + raw = torch.zeros((num_pages + 1) * page_bytes, dtype=torch.uint8) + k_views, v_views = build_mha_views( + raw, + layer_num=layer_num, + head_num=head_num, + head_dim=head_dim, + v_head_dim=head_dim, + store_dtype=store_dtype, + page_size=page_size, + num_pages=num_pages, + anchor_bytes=0, + ) + blocks = 2 * layer_num + # Distinct content in source page 1, every layer, K and V. + for layer in range(layer_num): + for pos in range(page_size): + row = 1 * blocks * page_size + pos + k_views[layer][row].fill_(layer + 1) + v_views[layer][row].fill_(-(layer + 1)) + + env = raw[: num_pages * page_bytes].view(num_pages, page_bytes) + env[3] = env[1] + + for layer in range(layer_num): + for pos in range(page_size): + row = 3 * blocks * page_size + pos + self.assertTrue( + torch.all(k_views[layer][row] == layer + 1), + f"K layer {layer} did not ride the envelope move", + ) + self.assertTrue( + torch.all(v_views[layer][row] == -(layer + 1)), + f"V layer {layer} did not ride the envelope move", + ) + + class TestMambaEnvelopeTransferAddressing(CustomTestCase): def test_slot_envelope_is_self_contained(self): """A slot's conv+temporal state for all layers must live exactly in diff --git a/test/registered/unit/mem_cache/test_unified_mha_views.py b/test/registered/unit/mem_cache/test_unified_mha_views.py index b97de3d34..d306b5d6f 100644 --- a/test/registered/unit/mem_cache/test_unified_mha_views.py +++ b/test/registered/unit/mem_cache/test_unified_mha_views.py @@ -345,11 +345,12 @@ class TestUnifiedMHATokenToKVPool(unittest.TestCase): ) def test_transfer_entry_points_fail_loud(self): - """PD / CPU-copy entry points assume per-layer buffers indexed by TOKEN - id and would silently mis-index the row space, so each must raise.""" + """The entry points that assume per-layer buffers indexed by TOKEN id + would silently mis-index against the row space (or hit a missing-attr + AttributeError), so each must raise. `get_contiguous_buf_infos` is NOT + among them: PD addresses this pool as whole page envelopes, pinned by + `test_pd_registration_is_one_whole_envelope` below.""" _, 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): @@ -357,6 +358,24 @@ class TestUnifiedMHATokenToKVPool(unittest.TestCase): with self.assertRaises(NotImplementedError): pool.set_kv_buffer_prefix_valid() + def test_pd_registration_is_one_whole_envelope(self): + """PD registers ONE region -- the whole raw buffer -- with the page + envelope as the item, so the transfer engine addresses it as + `raw_ptr + physical_page * page_envelope_bytes`. Per-layer regions + would be wrong here: the per-layer views overlap inside the envelope + and index in kernel-facing ids, not token ids.""" + kv, pool = _make_pool_and_kv(1) + ptrs, lens, item_lens = pool.get_contiguous_buf_infos() + self.assertEqual(len(ptrs), 1) + self.assertEqual(len(lens), 1) + self.assertEqual(len(item_lens), 1) + self.assertEqual(ptrs[0], kv._raw.data_ptr()) + self.assertEqual(lens[0], kv._raw.numel()) + self.assertEqual(item_lens[0], pool._page_bytes) + # The whole addressable page range must fit the registered region, or + # the last page's write would run off the end of the RDMA mapping. + self.assertLessEqual(pool._num_pages * item_lens[0], lens[0]) + def test_hnd_env_cannot_hijack_layout(self): """SGLANG_USE_HND_KVCACHE must not flip this pool's layout: HND indexes 4-D while the per-layer views are 3-D, so the pinned label has to win."""