From c55e40617634fe3bf21bd6b900c84efb7a3578d8 Mon Sep 17 00:00:00 2001 From: leavelet Date: Thu, 11 Jun 2026 20:13:04 +0000 Subject: [PATCH] Add symm-heap current-page exchange for CP shared-KV compose (Step B) Replaces the last remaining NCCL collective in the bs>1 compose layer loop (the compact-current all-reduce) with a DeepEP-style counting barrier plus a CUDA-IPC gather from peers' symmetric dense buffers, behind SGLANG_CP_SHARED_KV_COMPOSE_SYMM (+ COMPOSE_ARENA, both default off). - CpComposeArena.register_symm: fixes capacity at the pool-derived bound (logical pages x dense page unit, overridable via SGLANG_CP_SHARED_KV_SYMM_HEAP_MB), allocates slab + flags in CUDA-IPC memory, exchanges handles once over the CP group; deterministic bump carve means peer_base + my_offset addresses any peer's dense buffer with no per-layer handshakes. Registration happens only in the token-KV compose (uniform first-use point); growth after registration raises. - Current pages are single-writer at page granularity under the page-aligned in-seq split, so the exchange is the existing gather_cuda_ipc_peer_pages with src==dst page ids and writer (compute owner) ranks; writers are built per batch by build_batch_current_page_writer_ranks and gated by maybe_build_current_page_writer_ranks (page_aligned metadata required). The barrier runs even with zero remote pages (counts must match). - _agreed_tai_ipc_peer_ptrs: the per-rank IPC capability probe is now agreed across the CP group (one-time MIN all-reduce per pool tensor) so ranks can never split between gather and collective paths and deadlock on mismatched NCCL shapes. - ComposePlan cache re-anchored ON the slot_remap object (forward-batch lifetime) instead of a module-level tensor-identity key, which could go stale when a freed tensor's address is reused by the next batch. Validation (g0034 syh-dev-new): tai-kernel cp_symm_barrier correctness (200 adversarial iterations, rotating 10ms producer delays, phase-safety, flags drained at quiescence) and perf (7.7us max-rank latency) both pass; 8-rank GPU test extends to the symm path - byte-identical to compose_v2 across 4 layers (parity halves exercised, slab registered); mem_cache suite 464 passed with the only failures being a documented pre-existing sys.modules stub pollution pair, reproduced identically with SGLANG_CP_SHARED_KV_COMPOSE_V2=0. Co-Authored-By: Claude Fable 5 --- python/sglang/srt/environ.py | 7 + .../attention/nsa/cp_shared_kv_compose.py | 201 +++++++++++-- .../attention/nsa/cp_shared_kv_runtime.py | 272 ++++++++++++++++-- .../srt/layers/attention/nsa/nsa_indexer.py | 8 + .../srt/layers/attention/nsa_backend.py | 32 +++ .../test_cp_shared_kv_compose_v2_8rank.py | 50 +++- .../mem_cache/test_cp_shared_kv_compose.py | 116 +++++++- .../mem_cache/test_cp_shared_kv_runtime.py | 8 - 8 files changed, 626 insertions(+), 68 deletions(-) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 4368fff7b..1b33ff455 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -229,6 +229,13 @@ class Envs: # layer-parity halves). Default off in Step A; Step B registers the arena # as symmetric memory and flips this on. SGLANG_CP_SHARED_KV_COMPOSE_ARENA = EnvBool(False) + # Step B: exchange current pages via barrier + symm-heap IPC gather (zero + # NCCL in the layer loop). Requires COMPOSE_ARENA=1, page-aligned in-seq + # split, and tai-kernel cp_symm_barrier. Default off until e2e parity. + SGLANG_CP_SHARED_KV_COMPOSE_SYMM = EnvBool(False) + # Total symm slab size override in MB (both halves). 0 = pool-derived + # hard bound (logical pages x dense page unit, see design doc). + SGLANG_CP_SHARED_KV_SYMM_HEAP_MB = EnvInt(0) # NSA MQA logits are materialized as fp32 [q, k] buffers inside DeepGEMM. # Lower values split query rows more aggressively to cap peak temporary memory. SGLANG_NSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2) diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_compose.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_compose.py index 3671aa796..19efb0512 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_compose.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_compose.py @@ -51,6 +51,14 @@ def cp_shared_kv_compose_arena_enabled() -> bool: return bool(envs.SGLANG_CP_SHARED_KV_COMPOSE_ARENA.get()) +def cp_shared_kv_compose_symm_enabled() -> bool: + """Step B: exchange current pages via symm-heap IPC gather (zero NCCL). + + Requires the arena (the symm slab IS the arena) — checked at use site. + """ + return bool(envs.SGLANG_CP_SHARED_KV_COMPOSE_SYMM.get()) + + # -------------------------------------------------------------------------- # Compose plan: layer-invariant descriptor tensors, built once per batch. # -------------------------------------------------------------------------- @@ -65,6 +73,11 @@ class ComposePlan: (current pages, invalid pages) carries -1 and is zero-filled by the IPC slot-dense kernel. ``current_dense_pages`` are dense page ids (slot + 1, dummy-page convention) of all current slots in merged-span order. + + When the caller provides per-current-page writer (compute-owner) ranks, + ``remote_current_*`` hold the gather list for the symm exchange: only the + pages written by OTHER ranks, src page == dst page (peers' dense buffers + are at identical offsets). """ prefix_owner_ranks: torch.Tensor @@ -73,6 +86,8 @@ class ComposePlan: total_slots: int num_prefix_slots: int num_current_pages: int + remote_current_writer_ranks: torch.Tensor | None = None + remote_current_dense_pages: torch.Tensor | None = None def _spans_key(spans: list[tuple[int, int]] | None) -> tuple[tuple[int, int], ...]: @@ -81,47 +96,44 @@ def _spans_key(spans: list[tuple[int, int]] | None) -> tuple[tuple[int, int], .. return tuple((int(s), int(e)) for s, e in spans) -def _tensor_identity(t: torch.Tensor) -> tuple[int, int, tuple[int, ...], str, str]: - return ( - int(t.untyped_storage().data_ptr()), - int(t.data_ptr()), - tuple(int(d) for d in t.shape), - str(t.dtype), - str(t.device), - ) - - -_PLAN_CACHE: dict[tuple, ComposePlan] = {} -_PLAN_CACHE_MAX = 32 - - def get_or_build_compose_plan( *, - slot_logical_pages: torch.Tensor, + slot_remap, layout: "CpSharedKVLayout", physical_page_capacity: int | None, prefix_spans: list[tuple[int, int]], current_spans: list[tuple[int, int]], kind: str, + current_page_writer_ranks: list[int] | None = None, ) -> ComposePlan: """Build (or fetch the per-batch cached) compose plan. - The plan is identical for every layer of a batch: the cache key is the - identity of ``slot_logical_pages`` (per-batch tensor, see - ``get_or_build_shared_token_kv_slot_remap``) plus the span lists and - layout parameters. + The plan is identical for every layer of a batch, so it is anchored ON + the slot_remap object (which itself lives on the forward batch — same + lifetime idiom as ``get_or_build_shared_token_kv_slot_remap``). A + module-level cache keyed by tensor identity would go stale when a freed + tensor's address is reused by the next batch. + + ``current_page_writer_ranks`` (compute owners, batch current-slot order, + page-aligned in-seq split) enables the symm current-page gather list. """ + plans = getattr(slot_remap, "_compose_plans", None) + if plans is None: + plans = {} + object.__setattr__(slot_remap, "_compose_plans", plans) key = ( - _tensor_identity(slot_logical_pages), int(layout.cp_size), int(layout.cp_rank), int(physical_page_capacity or -1), _spans_key(prefix_spans), _spans_key(current_spans), kind, + tuple(current_page_writer_ranks) + if current_page_writer_ranks is not None + else None, ) - plan = _PLAN_CACHE.get(key) + plan = plans.get(key) if plan is not None: return plan @@ -130,7 +142,7 @@ def get_or_build_compose_plan( build_cp_shared_kv_ipc_page_descriptors, ) - flat_pages = slot_logical_pages.reshape(-1) + flat_pages = slot_remap.slot_logical_pages.reshape(-1) total_slots = int(flat_pages.numel()) device = flat_pages.device @@ -159,6 +171,23 @@ def get_or_build_compose_plan( else: current_dense_pages = torch.empty(0, device=device, dtype=torch.long) + remote_writer_ranks = None + remote_dense_pages = None + if current_page_writer_ranks is not None: + num_current = int(current_dense_pages.numel()) + if len(current_page_writer_ranks) != num_current: + raise ValueError( + "[CP_SHARED_KV_FAIL_FAST][compose_plan] current page writer " + f"count mismatch: writers={len(current_page_writer_ranks)} " + f"current_pages={num_current} (page-aligned split required)" + ) + writers = torch.tensor( + current_page_writer_ranks, dtype=torch.long, device=device + ) + remote = writers != int(layout.cp_rank) + remote_writer_ranks = writers[remote].contiguous() + remote_dense_pages = current_dense_pages[remote].contiguous() + plan = ComposePlan( prefix_owner_ranks=prefix_owner_ranks, prefix_src_pages=prefix_src_pages, @@ -166,19 +195,13 @@ def get_or_build_compose_plan( total_slots=total_slots, num_prefix_slots=num_prefix_slots, num_current_pages=int(current_dense_pages.numel()), + remote_current_writer_ranks=remote_writer_ranks, + remote_current_dense_pages=remote_dense_pages, ) - if len(_PLAN_CACHE) >= _PLAN_CACHE_MAX: - # Plans are per-batch; a bounded FIFO keeps stale batches from piling - # up without adding bookkeeping on the (hot) cache-hit path. - _PLAN_CACHE.pop(next(iter(_PLAN_CACHE))) - _PLAN_CACHE[key] = plan + plans[key] = plan return plan -def reset_compose_plan_cache() -> None: - _PLAN_CACHE.clear() - - # -------------------------------------------------------------------------- # Compose arena: tier-S carve discipline (symm-registration ready). # -------------------------------------------------------------------------- @@ -208,6 +231,11 @@ class CpComposeArena: self._parity_layer: list[int | None] = [None, None] self._high_water = 0 self._registered = False # Step B: set after IPC registration + # Step B symm state (populated by register_symm): + self.peer_slab_bases: torch.Tensor | None = None # CPU int64 [cp] + self.flag_ptrs: torch.Tensor | None = None # CUDA int64 [cp] + self._flags: torch.Tensor | None = None + self.cp_rank: int = -1 @staticmethod def _align(nbytes: int) -> int: @@ -253,6 +281,117 @@ class CpComposeArena: def high_water_bytes(self) -> int: return self._high_water + @property + def symm_ready(self) -> bool: + return self._registered + + def slab_offset_of(self, buffer: torch.Tensor) -> int: + assert self._slab is not None + offset = buffer.data_ptr() - self._slab.data_ptr() + if not (0 <= offset < 2 * self._half_bytes): + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][compose_arena] buffer is not from " + f"this arena slab: offset={offset}" + ) + return int(offset) + + def register_symm( + self, + *, + cp_group, + cp_rank: int, + cp_size: int, + half_bytes: int, + ) -> None: + """Fix capacity, allocate the slab + flags in CUDA-IPC memory, and + exchange handles ONCE across the CP group (collective — every rank + must call this at the same point; the first symm compose of a batch + is such a point).""" + + from tai_kernel.nsa_prefill.ipc import ( + allocate_cuda_ipc_buffer, + get_cuda_ipc_mem_handle, + open_cuda_ipc_mem_handles, + ) + + if self._registered: + return + if self._slab is not None: + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][compose_arena] register_symm must " + "run before any non-symm slab allocation (enable " + "SGLANG_CP_SHARED_KV_COMPOSE_SYMM from startup, not mid-run)" + ) + half_bytes = self._align(half_bytes) + self._slab = allocate_cuda_ipc_buffer(2 * half_bytes, device=self.device) + self._half_bytes = half_bytes + flags = allocate_cuda_ipc_buffer(max(cp_size * 4, 8), device=self.device) + flags.zero_() + self._flags = flags + + def _exchange(buffer: torch.Tensor) -> torch.Tensor: + handle = get_cuda_ipc_mem_handle(buffer) + gathered = torch.empty( + (cp_size, int(handle.numel())), + dtype=torch.uint8, + device=self.device, + ) + cp_group.all_gather_into_tensor( + gathered, handle.to(device=self.device) + ) + return open_cuda_ipc_mem_handles( + gathered.cpu(), buffer, self_rank=cp_rank + ) + + self.peer_slab_bases = _exchange(self._slab) + flag_peer_ptrs = _exchange(flags) + torch.cuda.synchronize(self.device) + self.flag_ptrs = flag_peer_ptrs.to(self.device) + self.cp_rank = int(cp_rank) + self._registered = True + logger.info( + "[CP-Compose-Arena] symm slab registered: half=%.1f MiB " + "total=%.1f MiB cp_rank=%s cp_size=%s", + half_bytes / (1 << 20), + 2 * half_bytes / (1 << 20), + cp_rank, + cp_size, + ) + + def peer_dense_ptrs_for(self, buffer: torch.Tensor) -> torch.Tensor: + """CPU int64 [cp]: each peer's pointer to ITS copy of ``buffer``. + + Valid because every rank carves identical offsets (deterministic + bump over rank-invariant shapes).""" + + assert self.peer_slab_bases is not None + return self.peer_slab_bases + self.slab_offset_of(buffer) + + +def compute_symm_half_bytes( + *, + kv_pool_tokens: int, + page_size: int, + cp_size: int, + kv_page_nbytes: int, + index_page_nbytes: int, + slack_pages: int = 512, +) -> int: + """Pool-derived hard bound for one arena half (see design doc §1). + + Every logical page a batch can materialize is backed by the shared pool: + logical capacity = per-rank pool pages × cp_size. One half must hold the + dense KV and dense index buffers of one layer plus rounding slack. + """ + + explicit_mb = int(envs.SGLANG_CP_SHARED_KV_SYMM_HEAP_MB.get()) + if explicit_mb > 0: + return explicit_mb << 19 # MB of TOTAL slab -> half bytes + + logical_pages = (int(kv_pool_tokens) // int(page_size)) * int(cp_size) + pages = logical_pages + int(slack_pages) + return pages * (int(kv_page_nbytes) + int(index_page_nbytes)) + _ARENAS: dict[torch.device, CpComposeArena] = {} diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py index 15b78d3bd..c5cc11a59 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -11,7 +11,11 @@ import torch from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( acquire_dense_buffer, + compute_symm_half_bytes, + cp_shared_kv_compose_arena_enabled, + cp_shared_kv_compose_symm_enabled, cp_shared_kv_compose_v2_enabled, + get_compose_arena, get_or_build_compose_plan, ) from sglang.srt.layers.attention.nsa.utils import ( @@ -2274,6 +2278,55 @@ def _get_or_open_tai_ipc_peer_ptrs( return None +_TAI_IPC_CAPABILITY_AGREEMENT: dict[tuple[object, ...], bool] = {} + + +def _agreed_tai_ipc_peer_ptrs( + tensor: torch.Tensor, + layout: CpSharedKVLayout, +) -> tuple[Any, torch.Tensor] | None: + """Probe peer-IPC capability and AGREE on it across the CP group. + + The probe alone is per-rank; if it diverged (one rank's open failing), + ranks would issue different collectives downstream and deadlock with a + shape mismatch. A one-time MIN-agreement per pool tensor pins every + rank to the same path. Called at uniform points only (the compose path + runs identically on all ranks). + """ + + state = _get_or_open_tai_ipc_peer_ptrs(tensor, layout) + if layout.cp_size <= 1: + return state + if not torch.distributed.is_initialized(): + # Single-process context (unit tests / tools): no group to agree + # with. Production CP always has torch.distributed initialized. + return state + cp_group = get_attention_cp_group() + key = _tai_ipc_peer_ptr_cache_key(tensor, layout, cp_group) + agreed = _TAI_IPC_CAPABILITY_AGREEMENT.get(key) + if agreed is None: + flag = torch.tensor( + [1 if state is not None else 0], + dtype=torch.int32, + device=tensor.device, + ) + torch.distributed.all_reduce( + flag, + op=torch.distributed.ReduceOp.MIN, + group=cp_group.device_group, + ) + agreed = bool(int(flag.item()) == 1) + _TAI_IPC_CAPABILITY_AGREEMENT[key] = agreed + if state is not None and not agreed: + logger.warning( + "[CP_SHARED_KV_FALLBACK][compose_v2] peer IPC works on this " + "rank but not on every CP rank; the whole group uses the " + "collective fallback. cp_rank=%s", + layout.cp_rank, + ) + return state if agreed else None + + def _page_nbytes_from_page_tensor(tensor: torch.Tensor) -> int: if tensor.ndim == 0: raise ValueError("CP shared KV IPC page tensor must have at least one dim") @@ -4332,6 +4385,136 @@ def materialize_local_token_kv_page_slots_into( dense_range.copy_(torch.where(owned_view, gathered, zero)) +def build_batch_current_page_writer_ranks( + *, + prefix_lens_cpu, + extend_lens_cpu, + page_size: int, + cp_size: int, +) -> list[int]: + """Per-current-page compute owner (writer) in batch current-slot order. + + Valid ONLY for the page-aligned in-seq split, where each current page is + written by exactly one rank. Order matches + ``build_batch_current_slot_spans`` (request-major ascending slots). + """ + + from sglang.srt.mem_cache.cp_shared_kv_compute_owner import ( + build_in_seq_page_compute_owners, + ) + + writers: list[int] = [] + for prefix_len, extend_len in zip(prefix_lens_cpu, extend_lens_cpu): + writers.extend( + int(owner) + for owner in build_in_seq_page_compute_owners( + extend_len=int(extend_len), + extend_prefix_len=int(prefix_len), + page_size=page_size, + cp_size=cp_size, + ) + ) + return writers + + +def maybe_build_current_page_writer_ranks( + *, + forward_batch, + prefix_lens_cpu, + extend_lens_cpu, + page_size: int, + layout: CpSharedKVLayout, +) -> list[int] | None: + """Caller-side gate for the symm current-page exchange. + + Returns writer ranks only when the symm path is enabled AND the batch + uses the page-aligned in-seq split (single writer per current page); + otherwise None, which keeps compose on the collective current exchange. + """ + + if not ( + cp_shared_kv_compose_symm_enabled() and cp_shared_kv_compose_arena_enabled() + ): + return None + metadata = getattr(forward_batch, "nsa_cp_metadata", None) + if metadata is None or not getattr(metadata, "page_aligned", False): + return None + if prefix_lens_cpu is None or extend_lens_cpu is None: + return None + return build_batch_current_page_writer_ranks( + prefix_lens_cpu=prefix_lens_cpu, + extend_lens_cpu=extend_lens_cpu, + page_size=page_size, + cp_size=int(layout.cp_size), + ) + + +def _symm_arena_ready_or_register( + *, + layout: CpSharedKVLayout, + kv_cache: torch.Tensor, + page_size: int, +) -> bool: + """Register the symm slab on first use (collective; uniform call point). + + Only the token-KV compose registers (it knows the exact KV page bytes); + the index compose uses symm once registration has happened. + """ + + arena = get_compose_arena(kv_cache.device) + if arena.symm_ready: + return True + cp_group = get_attention_cp_group() + # NSA index page bytes are model constants (index_head_dim 128, + # quant_block 128 -> head + 4 scale bytes per token). + index_page_nbytes = page_size * (128 + 4) + arena.register_symm( + cp_group=cp_group, + cp_rank=int(layout.cp_rank), + cp_size=int(layout.cp_size), + half_bytes=compute_symm_half_bytes( + kv_pool_tokens=int(kv_cache.shape[0]), + page_size=page_size, + cp_size=int(layout.cp_size), + kv_page_nbytes=_token_kv_page_nbytes(kv_cache, page_size), + index_page_nbytes=index_page_nbytes, + ), + ) + return arena.symm_ready + + +def _symm_exchange_current_pages( + dense_buffer: torch.Tensor, + plan, + layout: CpSharedKVLayout, + *, + page_nbytes: int, +) -> None: + """Step B current-page exchange: barrier + IPC gather from peers' symm + dense buffers (identical offsets on every rank). The barrier runs even + with zero remote pages — barrier COUNTS must match across ranks.""" + + from tai_kernel.nsa_prefill.ipc import ( + cp_symm_barrier, + gather_cuda_ipc_peer_pages, + ) + + arena = get_compose_arena(dense_buffer.device) + cp_symm_barrier(arena.flag_ptrs, self_rank=int(layout.cp_rank)) + if ( + plan.remote_current_dense_pages is not None + and plan.remote_current_dense_pages.numel() > 0 + ): + gather_cuda_ipc_peer_pages( + arena.peer_dense_ptrs_for(dense_buffer), + dense_buffer, + plan.remote_current_writer_ranks, + plan.remote_current_dense_pages, + plan.remote_current_dense_pages, + page_nbytes=page_nbytes, + ) + + def _resolve_partial_current_spans( *, current_slot_spans: list[tuple[int, int]] | None, @@ -4404,35 +4587,56 @@ def _compose_token_kv_partial_current_v2( layer_id: int | None, nvtx_source: str, timing_start: object, + current_page_writer_ranks: list[int] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: - """Step A compose for the MLA token-KV dense buffer. + """Step A/B compose for the MLA token-KV dense buffer. Fast path: one IPC slot-dense gather covers ALL prefix spans (slots - outside them are -1 sentinels, zero-filled by the kernel) + one collective - over the compact current pages. Fallback (no peer IPC): local materialize - of all prefix spans (no per-span collectives) + ONE sum-all-reduce over - the whole dense buffer, exact because every row is writer-exclusive at - reduce time. + outside them are -1 sentinels, zero-filled by the kernel); current pages + are then shared either via barrier + symm-heap IPC gather (Step B, zero + NCCL) or via one collective over the compact current pages. Fallback + (no peer IPC): local materialize of all prefix spans (no per-span + collectives) + ONE sum-all-reduce over the whole dense buffer, exact + because every row is writer-exclusive at reduce time. """ + use_symm = ( + cp_shared_kv_compose_symm_enabled() + and cp_shared_kv_compose_arena_enabled() + and current_page_writer_ranks is not None + and layer_id is not None + ) plan = get_or_build_compose_plan( - slot_logical_pages=slot_remap.slot_logical_pages, + slot_remap=slot_remap, layout=layout, physical_page_capacity=kv_cache.shape[0] // page_size, prefix_spans=prefix_spans, current_spans=current_spans, kind="token_kv", + current_page_writer_ranks=( + current_page_writer_ranks if use_symm else None + ), ) dense_rows = int(slot_remap.dense_num_pages) * page_size - # Path selection is a capability decision, not error handling: the probe - # (cached per pool tensor) tells us whether peer-IPC is available. After - # a successful probe, a failing gather is a bug — let it raise. + # Path selection is a capability decision, not error handling: the + # group-agreed probe (cached per pool tensor) tells us whether peer-IPC + # is available everywhere. After a successful probe, a failing gather + # is a bug — let it raise. ipc_state = ( - _get_or_open_tai_ipc_peer_ptrs(kv_cache, layout) + _agreed_tai_ipc_peer_ptrs(kv_cache, layout) if plan.num_prefix_slots > 0 else None ) + if use_symm: + # Symm requires the prefix IPC capability (same transport) and the + # registered slab; registration is collective and happens here, at + # the uniform first-use point. + use_symm = ipc_state is not None and _symm_arena_ready_or_register( + layout=layout, + kv_cache=kv_cache, + page_size=page_size, + ) gathered = ipc_state is not None if gathered: kernels, peer_ptrs = ipc_state @@ -4487,7 +4691,14 @@ def _compose_token_kv_partial_current_v2( ) if gathered: - if plan.num_current_pages > 0: + if use_symm: + _symm_exchange_current_pages( + mixed_kv_cache, + plan, + layout, + page_nbytes=_token_kv_page_nbytes(kv_cache, page_size), + ) + elif plan.num_current_pages > 0: _reduce_current_pages_compact( mixed_kv_cache, int(slot_remap.dense_num_pages), @@ -4551,6 +4762,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots( logical_locs_row_ids: torch.Tensor | None = None, layer_id: int | None = None, nvtx_source: str = "mla.partial_current_sync", + current_page_writer_ranks: list[int] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Synchronously compose prefix materialization with current KV rows. @@ -4643,6 +4855,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots( layer_id=layer_id, nvtx_source=nvtx_source, timing_start=timing_start, + current_page_writer_ranks=current_page_writer_ranks, ) dense_kv_cache = kv_cache.new_zeros( @@ -4780,25 +4993,43 @@ def _compose_index_partial_current_v2( layer_id: int | None, nvtx_source: str, timing_start: object, + current_page_writer_ranks: list[int] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: - """Step A compose for the indexer page buffer (see token-KV variant).""" + """Step A/B compose for the indexer page buffer (see token-KV variant). + The index compose never registers the symm slab itself (the token-KV + compose does, knowing the exact KV page bytes); it uses symm only once + registration already happened — deterministic across ranks because the + layer order is identical everywhere. + """ + + use_symm = ( + cp_shared_kv_compose_symm_enabled() + and cp_shared_kv_compose_arena_enabled() + and current_page_writer_ranks is not None + and layer_id is not None + and get_compose_arena(page_buffer.device).symm_ready + ) plan = get_or_build_compose_plan( - slot_logical_pages=slot_remap.slot_logical_pages, + slot_remap=slot_remap, layout=layout, physical_page_capacity=page_buffer.shape[0], prefix_spans=prefix_spans, current_spans=current_spans, kind="index", + current_page_writer_ranks=( + current_page_writer_ranks if use_symm else None + ), ) dense_num_pages = int(slot_remap.dense_num_pages) ipc_state = ( - _get_or_open_tai_ipc_peer_ptrs(page_buffer, layout) + _agreed_tai_ipc_peer_ptrs(page_buffer, layout) if plan.num_prefix_slots > 0 else None ) gathered = ipc_state is not None + use_symm = use_symm and gathered if gathered: kernels, peer_ptrs = ipc_state dense_page_buffer = acquire_dense_buffer( @@ -4841,7 +5072,14 @@ def _compose_index_partial_current_v2( ) if gathered: - if plan.num_current_pages > 0: + if use_symm: + _symm_exchange_current_pages( + dense_page_buffer, + plan, + layout, + page_nbytes=_page_nbytes_from_page_tensor(page_buffer), + ) + elif plan.num_current_pages > 0: _reduce_current_pages_compact( dense_page_buffer, dense_num_pages, @@ -4902,6 +5140,7 @@ def materialize_prefix_and_reuse_current_index_page_slots( current_slot_spans: list[tuple[int, int]] | None = None, layer_id: int | None = None, nvtx_source: str = "index.partial_current_sync", + current_page_writer_ranks: list[int] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Synchronously compose prefix index materialization with current index rows.""" timing_start = cp_shared_kv_bs_gt1_timing_start() @@ -4981,6 +5220,7 @@ def materialize_prefix_and_reuse_current_index_page_slots( layer_id=layer_id, nvtx_source=nvtx_source, timing_start=timing_start, + current_page_writer_ranks=current_page_writer_ranks, ) dense_page_buffer = page_buffer.new_zeros( diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 5ea47fa06..accd22966 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -31,6 +31,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( log_cp_draft_shared_kv_debug, materialize_prefix_and_reuse_current_index_page_slots, materialize_shared_paged_buffer, + maybe_build_current_page_writer_ranks, should_reuse_current_extend_kv, tensor_debug_checksum, tensor_debug_summary, @@ -689,6 +690,13 @@ class Indexer(MultiPlatformOp): prefix_slot_spans=prefix_slot_spans, current_slot_spans=current_slot_spans, layer_id=layer_id, + current_page_writer_ranks=maybe_build_current_page_writer_ranks( + forward_batch=forward_batch, + prefix_lens_cpu=prefix_lens_cpu, + extend_lens_cpu=extend_lens_cpu, + page_size=page_size, + layout=layout, + ), ) ) if ( diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index a336a0323..433d6d41d 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -36,6 +36,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( mask_batch_logical_pages_to_valid_lengths, materialize_prefix_and_reuse_current_kv_page_slots, materialize_shared_token_kv_buffer, + maybe_build_current_page_writer_ranks, pack_current_mla_kv_for_reuse, tensor_debug_checksum, tensor_debug_summary, @@ -2118,6 +2119,19 @@ class NativeSparseAttnBackend( current_slot_spans=current_slot_spans, layer_id=layer.layer_id, nvtx_source="mla.current_only_page_slots", + current_page_writer_ranks=( + maybe_build_current_page_writer_ranks( + forward_batch=forward_batch, + prefix_lens_cpu=getattr( + forward_batch, "extend_prefix_lens_cpu", None + ), + extend_lens_cpu=getattr( + forward_batch, "extend_seq_lens_cpu", None + ), + page_size=page_size, + layout=forward_batch.cp_shared_kv_layout, + ) + ), ) ) else: @@ -2256,6 +2270,15 @@ class NativeSparseAttnBackend( prefix_slot_spans=prefix_slot_spans, current_slot_spans=current_slot_spans, layer_id=layer.layer_id, + current_page_writer_ranks=( + maybe_build_current_page_writer_ranks( + forward_batch=forward_batch, + prefix_lens_cpu=prefix_lens_cpu, + extend_lens_cpu=extend_lens_cpu, + page_size=page_size, + layout=forward_batch.cp_shared_kv_layout, + ) + ), ) ) if ( @@ -2645,6 +2668,15 @@ class NativeSparseAttnBackend( current_req_id=ragged_current_req_id, layer_id=layer.layer_id, nvtx_source="mla.ragged_partial_current_sync", + current_page_writer_ranks=( + maybe_build_current_page_writer_ranks( + forward_batch=forward_batch, + prefix_lens_cpu=prefix_lens_cpu, + extend_lens_cpu=extend_lens_cpu, + page_size=page_size, + layout=forward_batch.cp_shared_kv_layout, + ) + ), ) ) if envs.SGLANG_NSA_DEQUANT_ONLY_TOPK.get(): diff --git a/test/manual/test_cp_shared_kv_compose_v2_8rank.py b/test/manual/test_cp_shared_kv_compose_v2_8rank.py index 17682dd0d..3e32269bf 100644 --- a/test/manual/test_cp_shared_kv_compose_v2_8rank.py +++ b/test/manual/test_cp_shared_kv_compose_v2_8rank.py @@ -21,9 +21,6 @@ import torch.distributed as dist from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime -from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( - reset_compose_plan_cache, -) from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout @@ -132,7 +129,12 @@ def _build_scenario(rank: int, cp_size: int, device: torch.device): phys_page * page_size : (phys_page + 1) * page_size ] = payload.view(page_size, 1, kv_dim).to(device) - # Current rows owned by this rank (writer = page owner in this scenario). + # Current rows owned by this rank (writer = page owner in this scenario, + # so the symm writer list is the storage owner per current page). + current_page_writer_ranks = [] + for req_pages, prefix_pages in zip(logical_rows, prefix_pages_by_req): + for page in req_pages[prefix_pages:]: + current_page_writer_ranks.append((page - 1) % cp_size) cur_locs = [ loc for req_id, loc in current_locs_all if (loc // page_size - 1) % cp_size == rank @@ -191,10 +193,11 @@ def _build_scenario(rank: int, cp_size: int, device: torch.device): page_size=page_size, prefix_slot_spans=prefix_slot_spans, current_slot_spans=current_slot_spans, + current_page_writer_ranks=current_page_writer_ranks, ) -def _compose(s, layer_id: int): +def _compose(s, layer_id: int, *, writers: list[int] | None = None): return runtime.materialize_prefix_and_reuse_current_kv_page_slots( kv_cache=s["kv_cache"], logical_locs=s["logical_locs"], @@ -209,6 +212,7 @@ def _compose(s, layer_id: int): prefix_slot_spans=s["prefix_slot_spans"], current_slot_spans=s["current_slot_spans"], layer_id=layer_id, + current_page_writer_ranks=writers, ) @@ -230,7 +234,6 @@ def main() -> None: torch.cuda.synchronize() dist.barrier() - reset_compose_plan_cache() with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override(True): v2_kv, v2_locs = _compose(s, layer_id=0) torch.cuda.synchronize() @@ -263,6 +266,41 @@ def main() -> None: f"(dense rows={int(ref_kv.shape[0])}, world={world})", flush=True, ) + + # ---- Step B: symm exchange (arena + barrier + peer gather, zero NCCL + # in the current-page phase). Multiple layers exercise parity halves. ---- + with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override( + True + ), envs.SGLANG_CP_SHARED_KV_COMPOSE_ARENA.override( + True + ), envs.SGLANG_CP_SHARED_KV_COMPOSE_SYMM.override(True): + for layer_id in range(4): + symm_kv, symm_locs = _compose( + s, layer_id=layer_id, writers=s["current_page_writer_ranks"] + ) + torch.cuda.synchronize() + assert torch.equal(ref_locs, symm_locs), ( + f"rank{rank} layer{layer_id}: symm locs mismatch" + ) + if not torch.equal(ref_kv, symm_kv): + diff = (ref_kv != symm_kv).any(dim=-1).any(dim=-1) + bad = torch.nonzero(diff).reshape(-1)[:8].cpu().tolist() + raise AssertionError( + f"rank{rank} layer{layer_id}: symm dense kv mismatch at " + f"rows {bad} (of {int(diff.sum())})" + ) + dist.barrier() + from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( + get_compose_arena, + ) + + assert get_compose_arena(device).symm_ready, "symm slab was not registered" + if rank == 0: + print( + "PASS: symm compose byte-identical to v2 across 4 layers " + "(arena registered, barrier + peer gather engaged)", + flush=True, + ) dist.destroy_process_group() diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_compose.py b/test/registered/unit/mem_cache/test_cp_shared_kv_compose.py index 3cb6c77f6..95f1c3e9a 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_compose.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_compose.py @@ -8,11 +8,12 @@ import unittest import torch from sglang.srt.environ import envs +from types import SimpleNamespace + from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( CpComposeArena, acquire_dense_buffer, get_or_build_compose_plan, - reset_compose_plan_cache, ) from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout @@ -22,9 +23,6 @@ register_cpu_ci(est_time=1, suite="stage-a-test-cpu") class TestComposePlan(unittest.TestCase): - def setUp(self): - reset_compose_plan_cache() - def _layout(self, cp_size=8, cp_rank=0, page_size=64): return CpSharedKVLayout( page_size=page_size, cp_size=cp_size, cp_rank=cp_rank @@ -39,7 +37,7 @@ class TestComposePlan(unittest.TestCase): prefix_spans = [(0, 3), (5, 7)] current_spans = [(3, 5), (7, 10)] plan = get_or_build_compose_plan( - slot_logical_pages=slot_logical_pages, + slot_remap=SimpleNamespace(slot_logical_pages=slot_logical_pages), layout=layout, physical_page_capacity=None, prefix_spans=prefix_spans, @@ -77,7 +75,7 @@ class TestComposePlan(unittest.TestCase): layout = self._layout() slot_logical_pages = torch.tensor([1, 0, 3], dtype=torch.int64) plan = get_or_build_compose_plan( - slot_logical_pages=slot_logical_pages, + slot_remap=SimpleNamespace(slot_logical_pages=slot_logical_pages), layout=layout, physical_page_capacity=None, prefix_spans=[(0, 3)], @@ -91,7 +89,7 @@ class TestComposePlan(unittest.TestCase): layout = self._layout() slot_logical_pages = torch.tensor([1, 2, 3, 4], dtype=torch.int64) kwargs = dict( - slot_logical_pages=slot_logical_pages, + slot_remap=SimpleNamespace(slot_logical_pages=slot_logical_pages), layout=layout, physical_page_capacity=None, prefix_spans=[(0, 2)], @@ -109,6 +107,110 @@ class TestComposePlan(unittest.TestCase): self.assertIsNot(plan_a, plan_d) +class TestComposePlanSymm(unittest.TestCase): + def test_remote_current_lists_exclude_self_written_pages(self): + layout = CpSharedKVLayout(page_size=64, cp_size=8, cp_rank=3) + slot_logical_pages = torch.tensor([1, 2, 3, 4, 5], dtype=torch.int64) + plan = get_or_build_compose_plan( + slot_remap=SimpleNamespace(slot_logical_pages=slot_logical_pages), + layout=layout, + physical_page_capacity=None, + prefix_spans=[(0, 2)], + current_spans=[(2, 5)], + kind="token_kv", + current_page_writer_ranks=[3, 0, 3], + ) + # current dense pages are slots 2,3,4 -> dense 3,4,5; writers 3,0,3; + # self rank 3 -> only the page written by rank 0 is remote. + self.assertEqual(plan.current_dense_pages.tolist(), [3, 4, 5]) + self.assertEqual(plan.remote_current_writer_ranks.tolist(), [0]) + self.assertEqual(plan.remote_current_dense_pages.tolist(), [4]) + + def test_writer_count_mismatch_fails_fast(self): + layout = CpSharedKVLayout(page_size=64, cp_size=8, cp_rank=0) + with self.assertRaises(ValueError): + get_or_build_compose_plan( + slot_remap=SimpleNamespace( + slot_logical_pages=torch.tensor([1, 2, 3], dtype=torch.int64) + ), + layout=layout, + physical_page_capacity=None, + prefix_spans=[(0, 1)], + current_spans=[(1, 3)], + kind="token_kv", + current_page_writer_ranks=[0], # 2 current pages + ) + + def test_writer_ranks_helper_matches_compute_owner_per_request(self): + from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + build_batch_current_page_writer_ranks, + ) + from sglang.srt.mem_cache.cp_shared_kv_compute_owner import ( + build_in_seq_page_compute_owners, + ) + + prefix_lens = [640, 320] + extend_lens = [95, 130] + writers = build_batch_current_page_writer_ranks( + prefix_lens_cpu=prefix_lens, + extend_lens_cpu=extend_lens, + page_size=64, + cp_size=8, + ) + expected = [] + for p, e in zip(prefix_lens, extend_lens): + expected.extend( + int(o) + for o in build_in_seq_page_compute_owners( + extend_len=e, + extend_prefix_len=p, + page_size=64, + cp_size=8, + ) + ) + self.assertEqual(writers, expected) + + def test_maybe_build_gates_on_env_and_page_alignment(self): + from types import SimpleNamespace + + from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + maybe_build_current_page_writer_ranks, + ) + + layout = CpSharedKVLayout(page_size=64, cp_size=8, cp_rank=0) + fb_aligned = SimpleNamespace( + nsa_cp_metadata=SimpleNamespace(page_aligned=True) + ) + fb_unaligned = SimpleNamespace( + nsa_cp_metadata=SimpleNamespace(page_aligned=False) + ) + kwargs = dict( + prefix_lens_cpu=[640], + extend_lens_cpu=[95], + page_size=64, + layout=layout, + ) + with envs.SGLANG_CP_SHARED_KV_COMPOSE_SYMM.override( + True + ), envs.SGLANG_CP_SHARED_KV_COMPOSE_ARENA.override(True): + self.assertIsNotNone( + maybe_build_current_page_writer_ranks( + forward_batch=fb_aligned, **kwargs + ) + ) + self.assertIsNone( + maybe_build_current_page_writer_ranks( + forward_batch=fb_unaligned, **kwargs + ) + ) + with envs.SGLANG_CP_SHARED_KV_COMPOSE_SYMM.override(False): + self.assertIsNone( + maybe_build_current_page_writer_ranks( + forward_batch=fb_aligned, **kwargs + ) + ) + + class TestComposeArena(unittest.TestCase): def test_parity_halves_reset_per_layer_and_offsets_are_deterministic(self): arena = CpComposeArena(torch.device("cpu")) diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py index f545d4618..260828567 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py @@ -2166,9 +2166,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): per-span range reduce, never a whole-buffer reduce).""" from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime - from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( - reset_compose_plan_cache, - ) from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout page_size = 4 @@ -2185,7 +2182,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): layout=layout, page_size=page_size, ) - reset_compose_plan_cache() captured = {} test_case = self @@ -2695,9 +2691,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): """compose_v2 contract for the index buffer (see token-KV twin).""" from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime - from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( - reset_compose_plan_cache, - ) from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout page_size = 4 @@ -2717,7 +2710,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): dtype=torch.uint8, ) current_scale = torch.tensor([[1.25], [2.5]], dtype=torch.float32) - reset_compose_plan_cache() captured = {} test_case = self