Harden symm compose per review: prefetch gate, round parity, hot-path CPU

Addresses the opus review of the Step A/B compose stack:

- FAIL FAST on symm + prefetcher coexistence: the MLA/index prefetchers
  issue their own per-span collectives and bypass the symm exchange, so
  letting them coexist would make COMPOSE_SYMM a silent no-op on every
  prefetch hit (measured "parity" that never ran). maybe_build_current_
  page_writer_ranks now raises when a prefetcher is attached or
  SGLANG_CP_SHARED_KV_ENABLE_MLA_PREFETCH is set.
- Arena parity now derives from a monotonic round epoch instead of raw
  layer_id: EAGLE draft layers reuse decoder layer ids, so id-parity would
  stop alternating halves (draft L0 -> next forward target L0 lands on the
  same half) and repeated-id rounds would keep appending into one half.
  begin_round(layer_id, kind) starts a new round (other half, offsets
  reset) when the id changes OR a kind repeats; regression test included.
- Per-batch writer-ranks cache on the forward batch and a presence flag in
  the plan key: previously the ~bs x current-pages writer list was rebuilt
  AND tuple-hashed on every layer per call site, pure launch-path waste.
- Rank-uniformity invariant documented at _symm_exchange_current_pages
  (any per-rank gate must go through a group agreement first) and an
  actionable arena-overflow message naming SGLANG_CP_SHARED_KV_SYMM_HEAP_MB.

Deferred follow-ups (documented in the design doc): factor the two near-
clone v2 helpers, single barrier per F-layer (~308us/batch), persistent
device-side peer-ptr tables for the gather.

182 targeted tests + 8-rank GPU byte-exactness (v2 and symm) re-validated
on g0034 syh-dev-new.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-06-11 20:26:27 +00:00
parent c55e406176
commit b356773d2f
3 changed files with 102 additions and 29 deletions

View File

@@ -129,9 +129,10 @@ def get_or_build_compose_plan(
_spans_key(prefix_spans),
_spans_key(current_spans),
kind,
tuple(current_page_writer_ranks)
if current_page_writer_ranks is not None
else None,
# Writer ranks are a deterministic function of the batch lens (cached
# per forward batch), so a presence flag suffices — hashing the ~bs x
# pages int list every layer would be pure cache-hit overhead.
current_page_writer_ranks is not None,
)
plan = plans.get(key)
if plan is not None:
@@ -228,9 +229,16 @@ class CpComposeArena:
self._slab: torch.Tensor | None = None
self._half_bytes = 0
self._offsets = [0, 0]
self._parity_layer: list[int | None] = [None, None]
self._high_water = 0
self._registered = False # Step B: set after IPC registration
# Round/epoch state: a "round" is one layer-instance's compose calls
# (index then kv on F-layers). Parity derives from a monotonic epoch,
# NOT from layer_id — EAGLE draft layers reuse decoder layer ids, so
# raw-id parity would stop alternating halves (and same-id rounds
# would keep appending into one half across forwards).
self._epoch = 0
self._round_layer_id: int | None = None
self._round_kinds: set[str] = set()
# 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]
@@ -241,12 +249,24 @@ class CpComposeArena:
def _align(nbytes: int) -> int:
return (nbytes + 255) & ~255
def begin_layer(self, parity: int, layer_token: int) -> None:
"""Reset the parity half when a new layer starts using it."""
parity &= 1
if self._parity_layer[parity] != layer_token:
self._parity_layer[parity] = layer_token
self._offsets[parity] = 0
def begin_round(self, layer_id: int, kind: str) -> int:
"""Advance to a new round (and the other parity half) when needed.
A new round starts when the layer id changes OR when the same buffer
kind repeats for one layer id (e.g. draft layer 0 followed by the
next forward's target layer 0). Within a round, all kinds carve the
same half at deterministic offsets. The call sequence is identical
on every rank, so epochs (and therefore offsets) stay symmetric.
Returns the parity of the current round.
"""
if layer_id != self._round_layer_id or kind in self._round_kinds:
self._epoch += 1
self._round_layer_id = layer_id
self._round_kinds = set()
self._offsets[self._epoch & 1] = 0
self._round_kinds.add(kind)
return self._epoch & 1
def _ensure_capacity(self, half_bytes: int) -> None:
if self._slab is not None and half_bytes <= self._half_bytes:
@@ -255,7 +275,9 @@ class CpComposeArena:
raise RuntimeError(
"[CP_SHARED_KV_FAIL_FAST][compose_arena] arena growth requested "
f"after symm registration: need_half={half_bytes} "
f"have_half={self._half_bytes}"
f"have_half={self._half_bytes}. Increase "
"SGLANG_CP_SHARED_KV_SYMM_HEAP_MB (total slab MB) or check "
"the pool-derived sizing in compute_symm_half_bytes."
)
new_half = max(half_bytes, self._half_bytes * 2, 64 << 20)
logger.info(
@@ -422,8 +444,7 @@ def acquire_dense_buffer(
if not cp_shared_kv_compose_arena_enabled() or layer_id is None:
return torch.empty(shape, dtype=dtype, device=device)
arena = get_compose_arena(device)
parity = int(layer_id) & 1
arena.begin_layer(parity, int(layer_id))
parity = arena.begin_round(int(layer_id), kind)
numel = 1
for dim in shape:
numel *= int(dim)

View File

@@ -4430,23 +4430,50 @@ def maybe_build_current_page_writer_ranks(
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.
The list is batch-invariant (a pure function of the batch lens), so it is
cached on the forward batch — rebuilding ~bs x current-pages ints per
layer per call site would sit on the launch-critical path for nothing.
"""
if not (
cp_shared_kv_compose_symm_enabled() and cp_shared_kv_compose_arena_enabled()
):
return None
# The prefetchers issue their own per-span collectives and never the
# barrier/symm exchange; letting them coexist would make SYMM a silent
# no-op on every prefetch hit. Fail fast instead of measuring parity.
if (
getattr(forward_batch, "cp_shared_kv_mla_prefetcher", None) is not None
or getattr(forward_batch, "cp_shared_kv_index_prefetcher", None) is not None
or envs.SGLANG_CP_SHARED_KV_ENABLE_MLA_PREFETCH.get()
):
raise RuntimeError(
"[CP_SHARED_KV_FAIL_FAST][compose_symm] "
"SGLANG_CP_SHARED_KV_COMPOSE_SYMM requires the CP shared-KV "
"prefetchers to be disabled (the prefetch path bypasses the "
"symm exchange and would keep issuing per-span all-reduces). "
"Unset SGLANG_CP_SHARED_KV_ENABLE_MLA_PREFETCH / index prefetch "
"or disable COMPOSE_SYMM."
)
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(
cache_key = (int(page_size), int(layout.cp_size))
cached_key = getattr(forward_batch, "cp_shared_kv_current_writer_key", None)
if cached_key == cache_key:
return forward_batch.cp_shared_kv_current_writer_ranks
writers = 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),
)
forward_batch.cp_shared_kv_current_writer_key = cache_key
forward_batch.cp_shared_kv_current_writer_ranks = writers
return writers
def _symm_arena_ready_or_register(
@@ -4492,7 +4519,14 @@ def _symm_exchange_current_pages(
) -> 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."""
with zero remote pages — barrier COUNTS must match across ranks.
INVARIANT: every gate on the path to this call (compose env flags,
page_aligned metadata, symm_ready, the agreed IPC capability, prefetch
absence) MUST be rank-uniform. A per-rank divergence here desyncs the
barrier counting and hangs the CP group — never add a per-rank condition
without routing it through a group agreement first
(see _agreed_tai_ipc_peer_ptrs)."""
from tai_kernel.nsa_prefill.ipc import (
cp_symm_barrier,