diff --git a/python/sglang/kernels/ops/attention/fla/bench_gdn_replayssm_fold.py b/python/sglang/kernels/ops/attention/fla/bench_gdn_replayssm_fold.py new file mode 100644 index 000000000..2bc372643 --- /dev/null +++ b/python/sglang/kernels/ops/attention/fla/bench_gdn_replayssm_fold.py @@ -0,0 +1,127 @@ +"""Microbenchmark for ``commit_gdn_replayssm_fold_all_layers`` (defaults match +Qwen3.5-397B at TP4). The kernel is bound by the mandatory checkpoint +read+write (``bs * layers * HV * K * V * 4 B * 2``), so the reported GB/s +approximates achieved HBM bandwidth. + +Run: ``python -m sglang.kernels.ops.attention.fla.bench_gdn_replayssm_fold`` +""" + +from __future__ import annotations + +import argparse + +import torch +import triton + +from sglang.kernels.ops.attention.fla.gdn_replayssm_spec_fold import ( + commit_gdn_replayssm_fold_all_layers, +) + + +def _make_pool(num_layers, num_slots, HV, H, K, V, RL, device): + state = torch.randn( + num_layers, num_slots, HV, K, V, device=device, dtype=torch.float32 + ) + rawv = torch.randn( + num_layers, num_slots, HV, RL, V, device=device, dtype=torch.bfloat16 + ) + rawk = torch.randn( + num_layers, num_slots, H, RL, K, device=device, dtype=torch.bfloat16 + ) + g_ring = ( + -torch.rand(num_layers, num_slots, HV, RL, device=device, dtype=torch.float32) + * 0.5 + ) + beta = torch.rand(num_layers, num_slots, HV, RL, device=device, dtype=torch.float32) + return state, rawv, rawk, g_ring, beta + + +def _bench_bs(*, pool, bs, num_slots, num_layers, HV, H, K, V, RL, device, with_track): + state, rawv, rawk, g_ring, beta = pool + gen = torch.Generator(device=device).manual_seed(bs) + slots = torch.randperm(num_slots, device=device, generator=gen)[:bs].to(torch.int32) + accept_lens = torch.randint( + 1, RL + 1, (bs,), device=device, dtype=torch.int32, generator=gen + ) + if with_track: + track_indices = torch.randperm(num_slots, device=device, generator=gen)[:bs].to( + torch.int64 + ) + track_steps = torch.where( + torch.rand(bs, device=device, generator=gen) < 0.25, + accept_lens.to(torch.int64) - 1, + torch.full((bs,), -1, dtype=torch.int64, device=device), + ) + else: + track_indices = None + track_steps = None + + def run(): + commit_gdn_replayssm_fold_all_layers( + checkpoint_state=state, + rawv_cache=rawv, + rawk_cache=rawk, + g_cache=g_ring, + beta_cache=beta, + ssm_state_indices=slots, + accept_lens=accept_lens, + max_cache_len=RL, + num_k_heads=H, + mamba_track_indices=track_indices, + mamba_steps_to_track=track_steps, + null_block_id=-1, + ) + + ms = triton.testing.do_bench(run, warmup=25, rep=100) + traffic_gb = bs * num_layers * HV * K * V * 4 * 2 / 1e9 + return ms * 1000, traffic_gb / (ms / 1000) if ms > 0 else 0.0 + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--num-layers", type=int, default=45) + parser.add_argument("--num-slots", type=int, default=475) + parser.add_argument("--hv", type=int, default=16) + parser.add_argument("--h", type=int, default=4) + parser.add_argument("--k", type=int, default=128) + parser.add_argument("--v", type=int, default=128) + parser.add_argument("--ring-len", type=int, default=4) + parser.add_argument("--batch-sizes", type=int, nargs="+", default=[1, 8, 32, 128]) + args = parser.parse_args() + + device = "cuda" + pool = _make_pool( + args.num_layers, + args.num_slots, + args.hv, + args.h, + args.k, + args.v, + args.ring_len, + device, + ) + print( + f"config: layers={args.num_layers} slots={args.num_slots} HV={args.hv} " + f"H={args.h} K={args.k} V={args.v} ring_len={args.ring_len}" + ) + print(f"{'bs':>4} {'track':>6} {'us/launch':>10} {'achieved GB/s':>14}") + for bs in args.batch_sizes: + for with_track in (False, True): + us, gbs = _bench_bs( + pool=pool, + bs=bs, + num_slots=args.num_slots, + num_layers=args.num_layers, + HV=args.hv, + H=args.h, + K=args.k, + V=args.v, + RL=args.ring_len, + device=device, + with_track=with_track, + ) + print(f"{bs:>4} {str(with_track):>6} {us:>10.1f} {gbs:>14.0f}") + + +if __name__ == "__main__": + main() diff --git a/python/sglang/kernels/ops/attention/fla/fused_sigmoid_gating_recurrent.py b/python/sglang/kernels/ops/attention/fla/fused_sigmoid_gating_recurrent.py index 37a54e8b7..521b0edf4 100644 --- a/python/sglang/kernels/ops/attention/fla/fused_sigmoid_gating_recurrent.py +++ b/python/sglang/kernels/ops/attention/fla/fused_sigmoid_gating_recurrent.py @@ -51,6 +51,16 @@ def fused_sigmoid_gating_delta_rule_update_kernel( DISABLE_STATE_UPDATE: tl.constexpr = False, CACHE_INTERMEDIATE_STATES: tl.constexpr = False, HAS_EAGLE_TREE_CUSTOM_ATTN_MASK: tl.constexpr = False, + replayssm_rawv=None, + replayssm_rawk=None, + replayssm_g=None, + replayssm_beta=None, + stride_rawv_slot: tl.constexpr = 0, + stride_rawk_slot: tl.constexpr = 0, + stride_g_slot: tl.constexpr = 0, + stride_beta_slot: tl.constexpr = 0, + MAX_CACHE_LEN: tl.constexpr = 0, + CACHE_RING: tl.constexpr = False, ): """ Fused kernel that combines sigmoid gating computation with recurrent delta rule update. @@ -173,6 +183,47 @@ def fused_sigmoid_gating_delta_rule_update_kernel( # Compute beta = sigmoid(b) b_beta = 1.0 / (1.0 + tl.exp(-b_b)) + # Stored here, pre-l2norm k / pre-delta v, so the commit fold's replay + # is bit-identical to the update below; steps >= MAX_CACHE_LEN would + # smash the next slot's ring. + if CACHE_RING: + ring_slot = tl.load(h0_indices + i_n).to(tl.int64) + if ring_slot >= 0 and step_idx < MAX_CACHE_LEN: + tl.store( + replayssm_rawv + + ring_slot * stride_rawv_slot + + i_hv * MAX_CACHE_LEN * V + + step_idx * V + + o_v, + b_v.to(replayssm_rawv.dtype.element_ty), + mask=mask_v, + ) + if i_v == 0: + tl.store( + replayssm_rawk + + ring_slot * stride_rawk_slot + + i_h * MAX_CACHE_LEN * K + + step_idx * K + + o_k, + b_k.to(replayssm_rawk.dtype.element_ty), + mask=mask_k, + ) + tl.store( + replayssm_g + + ring_slot * stride_g_slot + + i_hv * MAX_CACHE_LEN + + step_idx, + b_g, + ) + if i_k == 0: + tl.store( + replayssm_beta + + ring_slot * stride_beta_slot + + i_hv * MAX_CACHE_LEN + + step_idx, + b_beta, + ) + # Apply L2 normalization if enabled if USE_QK_L2NORM_IN_KERNEL: b_q = b_q / (tl.sqrt(tl.sum(b_q * b_q) + 1e-6)) @@ -262,6 +313,11 @@ def fused_sigmoid_gating_delta_rule_update( int ] = None, # kept for API compat; stride is derived from ``intermediate_states_buffer.shape[1]`` retrieve_parent_token: Optional[torch.Tensor] = None, + cache_ring: bool = False, + replayssm_rawv: Optional[torch.Tensor] = None, + replayssm_rawk: Optional[torch.Tensor] = None, + replayssm_g: Optional[torch.Tensor] = None, + replayssm_beta: Optional[torch.Tensor] = None, ): """ Fused triton implementation of sigmoid gating delta rule update. @@ -319,6 +375,25 @@ def fused_sigmoid_gating_delta_rule_update( else 0 ) + if cache_ring: + assert not is_kda, "cache_ring supports GDN only (scalar gate layout)" + # stride(0) is used as the slot pitch, so a tensor still carrying the + # layer dim would scribble outside its slot. + assert ( + replayssm_rawv.dim() == 4 + and replayssm_rawk.dim() == 4 + and replayssm_g.dim() == 3 + and replayssm_beta.dim() == 3 + ), "cache_ring expects per-layer ring views" + max_cache_len = replayssm_rawv.shape[-2] + stride_rawv_slot = replayssm_rawv.stride(0) + stride_rawk_slot = replayssm_rawk.stride(0) + stride_g_slot = replayssm_g.stride(0) + stride_beta_slot = replayssm_beta.stride(0) + else: + max_cache_len = 0 + stride_rawv_slot = stride_rawk_slot = stride_g_slot = stride_beta_slot = 0 + fused_sigmoid_gating_delta_rule_update_kernel[grid]( A_log=A_log, a=a, @@ -361,6 +436,16 @@ def fused_sigmoid_gating_delta_rule_update( DISABLE_STATE_UPDATE=disable_state_update, CACHE_INTERMEDIATE_STATES=intermediate_states_buffer is not None, HAS_EAGLE_TREE_CUSTOM_ATTN_MASK=retrieve_parent_token is not None, + replayssm_rawv=replayssm_rawv, + replayssm_rawk=replayssm_rawk, + replayssm_g=replayssm_g, + replayssm_beta=replayssm_beta, + stride_rawv_slot=stride_rawv_slot, + stride_rawk_slot=stride_rawk_slot, + stride_g_slot=stride_g_slot, + stride_beta_slot=stride_beta_slot, + MAX_CACHE_LEN=max_cache_len, + CACHE_RING=cache_ring, num_warps=num_warps, num_stages=num_stages, ) diff --git a/python/sglang/kernels/ops/attention/fla/gdn_replayssm_spec_fold.py b/python/sglang/kernels/ops/attention/fla/gdn_replayssm_spec_fold.py new file mode 100644 index 000000000..2803ad6e8 --- /dev/null +++ b/python/sglang/kernels/ops/attention/fla/gdn_replayssm_spec_fold.py @@ -0,0 +1,261 @@ +# SPDX-License-Identifier: Apache-2.0 +"""GDN ReplaySSM fold-every-commit: replay the ring-written raw inputs of the +accepted draft prefix into the fp32 checkpoint on commit (replaces the +per-draft ``intermediate_ssm`` snapshots). + +The fold is a BITWISE CLONE of ``fused_sigmoid_gating_delta_rule_update_kernel``'s +GDN branch (same tile, division-form L2 norm, op order). Do NOT reorder into +tl.dot / reciprocal-multiply; keep num_warps=1 so the reduction trees match. +""" + +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def gdn_replayssm_exact_fold_kernel( + h0, # [num_slots, HV, K, V] fp32 checkpoint (folded in place) + rawv_cache, # [num_slots, HV, RL, V] raw v + rawk_cache, # [num_slots, H, RL, K] raw pre-norm k + g_cache, # [num_slots, HV, RL] fp32 scalar log-decay gate + beta_cache, # [num_slots, HV, RL] fp32 beta + ssm_state_indices, # [B] int physical slot per request + accept_lens, # [B] int committed prefix length per request (incl. bonus) + mamba_track_indices, # [B] int extra_buffer track slot (or NULL) per request + mamba_steps_to_track, # [B] int crossing step (or -1) per request + stride_state_slot: tl.constexpr, + stride_rawv_slot: tl.constexpr, + stride_rawk_slot: tl.constexpr, + stride_g_slot: tl.constexpr, + stride_beta_slot: tl.constexpr, + stride_state_layer: tl.constexpr, + stride_rawv_layer: tl.constexpr, + stride_rawk_layer: tl.constexpr, + stride_g_layer: tl.constexpr, + stride_beta_layer: tl.constexpr, + stride_indices: tl.constexpr, + stride_accept: tl.constexpr, + stride_track: tl.constexpr, + stride_steps: tl.constexpr, + H: tl.constexpr, + HV: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + MAX_CACHE_LEN: tl.constexpr, + USE_QK_L2NORM_IN_KERNEL: tl.constexpr, + NULL_BLOCK_ID: tl.constexpr, + HAS_TRACK: tl.constexpr, +): + i_v = tl.program_id(0) + i_n = tl.program_id(1) + i_hvl = tl.program_id(2) + # int64: layer stride * i_layer can overflow an int32 index product. + i_layer = (i_hvl // HV).to(tl.int64) + i_hv = i_hvl % HV + i_h = i_hv // (HV // H) + h0 = h0 + i_layer * stride_state_layer + rawv_cache = rawv_cache + i_layer * stride_rawv_layer + rawk_cache = rawk_cache + i_layer * stride_rawk_layer + g_cache = g_cache + i_layer * stride_g_layer + beta_cache = beta_cache + i_layer * stride_beta_layer + + state_idx = tl.load(ssm_state_indices + i_n * stride_indices).to(tl.int64) + if state_idx <= NULL_BLOCK_ID: + return + n_commit = tl.load(accept_lens + i_n * stride_accept).to(tl.int32) + if n_commit <= 0: + return + + if HAS_TRACK: + track_idx = tl.load(mamba_track_indices + i_n * stride_track).to(tl.int64) + track_step = tl.load(mamba_steps_to_track + i_n * stride_steps).to(tl.int32) + else: + track_idx = NULL_BLOCK_ID + track_step = -1 + + o_k = tl.arange(0, BK) + o_v = i_v * BV + tl.arange(0, BV) + mask_k = o_k < K + mask_v = o_v < V + mask_h = mask_k[:, None] & mask_v[None, :] + + # Checkpoint-bandwidth bound. Tuning tried and rejected: num_stages (no-op), + # num_warps > 1 (breaks the bitwise clone), evict_first (loses cold-L2). + p_h0 = ( + h0 + + state_idx * stride_state_slot + + i_hv * K * V + + o_v[None, :] * K + + o_k[:, None] + ) + b_h = tl.load(p_h0, mask=mask_h, other=0.0).to(tl.float32) + + for t in range(0, n_commit): + phys = t.to(tl.int64) + b_k = tl.load( + rawk_cache + + state_idx * stride_rawk_slot + + (i_h * MAX_CACHE_LEN + phys) * K + + o_k, + mask=mask_k, + other=0.0, + ).to(tl.float32) + b_v = tl.load( + rawv_cache + + state_idx * stride_rawv_slot + + (i_hv * MAX_CACHE_LEN + phys) * V + + o_v, + mask=mask_v, + other=0.0, + ).to(tl.float32) + b_g = tl.load( + g_cache + state_idx * stride_g_slot + i_hv * MAX_CACHE_LEN + phys + ).to(tl.float32) + b_beta = tl.load( + beta_cache + state_idx * stride_beta_slot + i_hv * MAX_CACHE_LEN + phys + ).to(tl.float32) + + if USE_QK_L2NORM_IN_KERNEL: + b_k = b_k / (tl.sqrt(tl.sum(b_k * b_k) + 1e-6)) + b_h *= tl.exp(b_g) + b_v -= tl.sum(b_h * b_k[:, None], 0) + b_v *= b_beta + b_h += b_k[:, None] * b_v[None, :] + + if HAS_TRACK: + if (t == track_step) and (track_idx > NULL_BLOCK_ID): + tl.store( + h0 + + track_idx * stride_state_slot + + i_hv * K * V + + o_v[None, :] * K + + o_k[:, None], + b_h.to(h0.dtype.element_ty), + mask=mask_h, + ) + + tl.store(p_h0, b_h.to(p_h0.dtype.element_ty), mask=mask_h) + + +def commit_gdn_replayssm_fold_all_layers( + checkpoint_state: torch.Tensor, # [num_layers, num_slots, HV, K, V], in place + rawv_cache: torch.Tensor, # [num_layers, num_slots, HV, RL, V] + rawk_cache: torch.Tensor, # [num_layers, num_slots, H, RL, K] + g_cache: torch.Tensor, # [num_layers, num_slots, HV, RL] fp32 + beta_cache: torch.Tensor, # [num_layers, num_slots, HV, RL] fp32 + ssm_state_indices: torch.Tensor, # [B] int (shared across layers) + accept_lens: torch.Tensor, # [B] int, incl. the bonus token + max_cache_len: int, + num_k_heads: int, + mamba_track_indices: torch.Tensor | None = None, + mamba_steps_to_track: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = True, + null_block_id: int = -1, +) -> None: + """Fold every layer's accepted window in one launch (layer packed into + grid axis 2); bit-identical to a per-layer loop.""" + num_layers, num_slots, HV = checkpoint_state.shape[:3] + K = rawk_cache.shape[-1] + V = rawv_cache.shape[-1] + B = ssm_state_indices.shape[0] + BK = triton.next_power_of_2(K) + BV = min(triton.next_power_of_2(V), 32) + grid = (triton.cdiv(V, BV), B, HV * num_layers) + has_track = mamba_track_indices is not None and mamba_steps_to_track is not None + if has_track: + track_idx_t = mamba_track_indices + steps_t = mamba_steps_to_track + stride_track = track_idx_t.stride(0) + stride_steps = steps_t.stride(0) + else: + # Unused when HAS_TRACK is False; pass a valid pointer. + track_idx_t = ssm_state_indices + steps_t = accept_lens + stride_track = 0 + stride_steps = 0 + gdn_replayssm_exact_fold_kernel[grid]( + checkpoint_state, + rawv_cache, + rawk_cache, + g_cache, + beta_cache, + ssm_state_indices, + accept_lens, + track_idx_t, + steps_t, + checkpoint_state.stride(1), + rawv_cache.stride(1), + rawk_cache.stride(1), + g_cache.stride(1), + beta_cache.stride(1), + checkpoint_state.stride(0), + rawv_cache.stride(0), + rawk_cache.stride(0), + g_cache.stride(0), + beta_cache.stride(0), + ssm_state_indices.stride(0), + accept_lens.stride(0), + stride_track, + stride_steps, + H=num_k_heads, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + MAX_CACHE_LEN=max_cache_len, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + NULL_BLOCK_ID=null_block_id, + HAS_TRACK=has_track, + num_warps=1, + num_stages=3, + ) + + +def commit_gdn_replayssm_fold_after_verify( + *, + spec_state, # MambaPool.SpeculativeState (all layers) + state_batch_indices: torch.Tensor, # [B] per-req mamba slot + accept_lens: torch.Tensor, # [B] int, incl. the bonus token + last_correct_step_indices: torch.Tensor, # [B] conv rollback target step + mamba_track_indices: torch.Tensor | None = None, + mamba_steps_to_track: torch.Tensor | None = None, + null_block_id: int = -1, +) -> None: + """Fold each layer's accepted window into ``temporal``, then do the usual + conv accept-rollback (+ the track-slot conv scatter under extra_buffer).""" + from sglang.kernels.ops.mamba.mamba_state_scatter_triton import ( + fused_conv_window_scatter_with_mask, + ) + + max_cache_len = spec_state.replayssm_rawv.shape[-2] + num_k_heads = spec_state.replayssm_rawk.shape[2] + commit_gdn_replayssm_fold_all_layers( + checkpoint_state=spec_state.temporal, + rawv_cache=spec_state.replayssm_rawv, + rawk_cache=spec_state.replayssm_rawk, + g_cache=spec_state.replayssm_g, + beta_cache=spec_state.replayssm_beta, + ssm_state_indices=state_batch_indices, + accept_lens=accept_lens, + max_cache_len=max_cache_len, + num_k_heads=num_k_heads, + mamba_track_indices=mamba_track_indices, + mamba_steps_to_track=mamba_steps_to_track, + null_block_id=null_block_id, + ) + for conv_states, interm_conv in zip( + spec_state.conv, spec_state.intermediate_conv_window + ): + fused_conv_window_scatter_with_mask( + conv_states, interm_conv, state_batch_indices, last_correct_step_indices + ) + if mamba_track_indices is not None and mamba_steps_to_track is not None: + fused_conv_window_scatter_with_mask( + conv_states, interm_conv, mamba_track_indices, mamba_steps_to_track + ) diff --git a/python/sglang/srt/configs/mamba_utils.py b/python/sglang/srt/configs/mamba_utils.py index 47ab2fcb4..11eefe174 100644 --- a/python/sglang/srt/configs/mamba_utils.py +++ b/python/sglang/srt/configs/mamba_utils.py @@ -124,6 +124,24 @@ class BaseLinearStateParams(ABC): + ssm_numel * self.dtype.temporal.itemsize ) * len(self.layers) + def replayssm_ring_bytes_per_req(self, record_len: int) -> int: + """Per-slot bytes of the ReplaySSM spec-verify fold window (all + layers). Not part of ``mamba_cache_per_req``, so the memory solver + must charge it separately. MUST mirror the ``MambaPool`` allocation: + raw v/k in the conv dtype + fp32 g and beta. GDN scalar-g layout + only.""" + assert not self.is_kda, "replayssm ring accounting supports GDN only" + hv, v_dim, k_dim = self.shape.temporal + h_k = self.shape.num_k_heads_per_tp + conv_b = self.dtype.conv.itemsize + fp32_b = 4 + per_layer = ( + hv * record_len * v_dim * conv_b + + h_k * record_len * k_dim * conv_b + + 2 * hv * record_len * fp32_b + ) + return per_layer * len(self.layers) + @property def is_kda(self) -> bool: """KDA per-K-channel gate vs GDN/Mamba2 per-head scalar gate. Selects diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 524f0bdb6..8bd0bb3b3 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -209,6 +209,9 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool): mamba_size: int = None, start_layer: int = None, speculative_eagle_topk: Optional[int] = None, + linear_replayssm_cache_len: int = 16, + mamba_envelope_layout: bool = False, + enable_gdn_replayssm_spec: bool = False, ): DecodeReqToTokenPool.__init__( self, @@ -253,6 +256,9 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool): enable_mamba_extra_buffer=self.enable_mamba_extra_buffer, speculative_num_draft_tokens=speculative_num_draft_tokens, speculative_eagle_topk=speculative_eagle_topk, + linear_replayssm_cache_len=linear_replayssm_cache_len, + mamba_envelope_layout=mamba_envelope_layout, + enable_gdn_replayssm_spec=enable_gdn_replayssm_spec, ) def clear(self): diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index dde215538..20af6d07f 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -590,23 +590,35 @@ class GDNAttnBackend(MambaAttnBackendBase): value = value.view(1, actual_seq_len, layer.num_v_heads, layer.head_v_dim) if is_target_verify: - # ReplaySSM spec-verify (Part B of #28511): when the per-slot ring is - # allocated (--enable-gdn-replayssm-spec, GDN + linear-chain topk<=1), - # reconstruct the verify output for the whole draft window from the - # frozen checkpoint (`temporal`) + the per-slot circular (d, k, g) ring - # instead of the recurrent verify that snapshots a full state per draft - # token. The cursors are advanced once per decode step by the worker - # (commit_gdn_replayssm_spec in spec_utils). GDN-only: KDA (per-K gate) - # routes through kda_backend and never reaches here; we additionally - # guard on `not replayssm_is_kda` for safety. Falls back to the - # recurrent verify when the ring is absent. + # ReplaySSM verify protocols: fold-every-commit (ring-write during + # verify, fold on commit), circular ring, or the snapshotting + # fallback when neither ring is allocated. mamba_pool = self.req_to_token_pool.mamba_pool + use_replayssm_fold = ( + mamba_cache_params.replayssm_rawv is not None + and getattr(mamba_pool, "replayssm_spec_fold", False) + and not getattr(mamba_pool, "replayssm_is_kda", False) + ) use_replayssm_spec = ( mamba_cache_params.replayssm_d is not None and getattr(mamba_pool, "replayssm_cache_base", None) is not None and not getattr(mamba_pool, "replayssm_is_kda", False) ) - if use_replayssm_spec: + if use_replayssm_fold: + core_attn_out = self._replayssm_fold_target_verify( + layer=layer, + query=query, + key=key, + value=value, + a=a, + b=b, + layer_cache=mamba_cache_params, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + retrieve_parent_token=retrieve_parent_token, + ) + elif use_replayssm_spec: core_attn_out = self._replayssm_target_verify( layer=layer, query=query, @@ -685,6 +697,55 @@ class GDNAttnBackend(MambaAttnBackendBase): return core_attn_out + def _replayssm_fold_target_verify( + self, + *, + layer: RadixLinearAttention, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + layer_cache: "MambaPool.SpeculativeState", + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + retrieve_parent_token: Optional[torch.Tensor], + ) -> torch.Tensor: + """Recurrent verify + fused ring-write; the commit fold replays the + accepted prefix into ``temporal``. Called directly, not via the kernel + dispatcher: the ring-write exists only in the Triton kernel.""" + from sglang.kernels.ops.attention.fla.fused_sigmoid_gating_recurrent import ( + fused_sigmoid_gating_delta_rule_update, + ) + + assert retrieve_parent_token is None, ( + "ReplaySSM fold-every-commit supports a linear draft chain only " + "(topk <= 1); EAGLE tree verify must use the recurrent verify." + ) + return fused_sigmoid_gating_delta_rule_update( + A_log=layer.A_log, + dt_bias=layer.dt_bias, + q=query, + k=key, + v=value, + a=a, + b=b, + initial_state_source=ssm_states, + initial_state_indices=cache_indices, + cu_seqlens=query_start_loc, + use_qk_l2norm_in_kernel=True, + softplus_beta=1.0, + softplus_threshold=20.0, + is_kda=False, + disable_state_update=True, + cache_ring=True, + replayssm_rawv=layer_cache.replayssm_rawv, + replayssm_rawk=layer_cache.replayssm_rawk, + replayssm_g=layer_cache.replayssm_g, + replayssm_beta=layer_cache.replayssm_beta, + ) + def _replayssm_target_verify( self, *, diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index bc4cf8154..1e55f464c 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -678,6 +678,12 @@ class KVCacheConfigurator: enable_overlap_schedule=not self.server_args.disable_overlap_schedule, mamba_size=self.server_args.max_mamba_cache_size, start_layer=self.layer_info.start_layer, + linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, + mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, + enable_gdn_replayssm_spec=( + self.server_args.enable_gdn_replayssm_spec + and self.hybrid_gdn_config is not None + ), ) return req_to_token_pool @@ -734,10 +740,6 @@ class KVCacheConfigurator: enable_linear_replayssm=self.server_args.enable_linear_replayssm, linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, - # ReplaySSM spec-verify is GDN-only: activate the pool machinery - # (rings + cursors + the intermediate_ssm gate) only for GDN-hybrid - # models, so any other mamba-ish model (Mamba2/Nemotron, lightning, - # ...) run with the flag set stays byte-identical to flag-off. enable_gdn_replayssm_spec=( self.server_args.enable_gdn_replayssm_spec and self.hybrid_gdn_config is not None @@ -1727,6 +1729,24 @@ class KVCacheConfigurator: assert config is not None has_spec_dec = not self.spec_algorithm.is_none() + # The ring is allocated per slot but is not part of mamba_cache_per_req; + # the solve must charge it too or num_slots is over-provisioned. + replayssm_active = ( + server_args.enable_gdn_replayssm_spec and self.hybrid_gdn_config is not None + ) + if replayssm_active: + record_len = ( + server_args.max_speculative_num_draft_tokens + if server_args.max_speculative_num_draft_tokens is not None + else server_args.linear_replayssm_cache_len + ) + replayssm_ring_per_req = ( + config.mamba2_cache_params.replayssm_ring_bytes_per_req( + record_len=record_len + ) + ) + else: + replayssm_ring_per_req = 0 if has_spec_dec: assert server_args.speculative_num_draft_tokens is not None assert server_args.max_running_requests is not None @@ -1739,7 +1759,7 @@ class KVCacheConfigurator: // self.ps.attn_dp_size, ) # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot) - if has_spec_dec: + if has_spec_dec and not replayssm_active: ratio = self._calculate_mamba_ratio() capped_reqs = min( server_args.max_running_requests // self.ps.attn_dp_size, @@ -1762,7 +1782,7 @@ class KVCacheConfigurator: // self.ps.attn_dp_size, ) # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot) - if has_spec_dec: + if has_spec_dec and not replayssm_active: intermediate_size = ( config.mamba2_cache_params.mamba_cache_per_req * (server_args.max_mamba_cache_size + 1) @@ -1784,7 +1804,7 @@ class KVCacheConfigurator: ) mamba_budget_bytes = mamba_budget * (1 << 30) - if has_spec_dec: + if has_spec_dec and not replayssm_active: ratio = self._calculate_mamba_ratio() D = server_args.speculative_num_draft_tokens # Joint solve: main_state + intermediate = mamba_budget @@ -1804,9 +1824,12 @@ class KVCacheConfigurator: intermediate_size = per_req * (capped_reqs + 1) * D total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) else: + per_slot = per_req + replayssm_ring_per_req server_args.override( "mamba_pool.memory_budget", - max_mamba_cache_size=int((mamba_budget_bytes - per_req) // per_req), + max_mamba_cache_size=int( + (mamba_budget_bytes - per_slot) // per_slot + ), ) # Validate: max_mamba_cache_size must be positive after memory allocation. @@ -1828,7 +1851,7 @@ class KVCacheConfigurator: # +1: the pool's padding slot mamba_state_memory = ( (server_args.max_mamba_cache_size + 1) - * config.mamba2_cache_params.mamba_cache_per_req + * (config.mamba2_cache_params.mamba_cache_per_req + replayssm_ring_per_req) / (1 << 30) ) return total_rest_memory - mamba_state_memory diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 66ef8e472..503dc3b6c 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -485,12 +485,13 @@ class MambaPool: self.debug_memory_pool = envs.SGLANG_DEBUG_MEMORY_POOL.get() self.enable_linear_replayssm = enable_linear_replayssm self.linear_replayssm_cache_len = linear_replayssm_cache_len - # ReplaySSM spec-verify (Part B of #28511) REUSES the linear_replayssm ring - # (replayssm_d/k/g + write_pos) and ADDS two per-slot cursors - # (replayssm_cache_base + replayssm_is_flush). Enabling the spec-verify path - # therefore implies the ring, so the d/k/g + write_pos allocation gates on - # `_replayssm_on` (either flag). GDN-only is enforced upstream + below. + # ReplaySSM: the decode ring (--enable-linear-replayssm) allocates the + # chunked (d, k) records + write_pos; the spec-verify flag + # (--enable-gdn-replayssm-spec) always uses fold-every-commit and + # allocates only the raw (v, k, g, beta) window -- no chunked records, + # no cursors. The shared g allocation gates on `_replayssm_on`. self.enable_gdn_replayssm_spec = enable_gdn_replayssm_spec + self.replayssm_spec_fold = bool(enable_gdn_replayssm_spec) _replayssm_on = enable_linear_replayssm or enable_gdn_replayssm_spec # for disagg with nvlink @@ -587,23 +588,32 @@ class MambaPool: # SSM dtype to halve the ring traffic. g stays fp32 everywhere # (exact-fold input). The two flags are mutually exclusive. ring_dtype = conv_dtype if enable_gdn_replayssm_spec else ssm_dtype - replayssm_d = torch.zeros( - size=(num_mamba_layers, num_slots, hv, L, v_dim), - dtype=ring_dtype, - device=device, - ) - replayssm_k = torch.zeros( - size=(num_mamba_layers, num_slots, h_k, L, k_dim), - dtype=ring_dtype, - device=device, - ) + # Fold-every-commit: one verify window, no chunked (d, k) records. + if self.replayssm_spec_fold: + record_len = ( + speculative_num_draft_tokens + if speculative_num_draft_tokens is not None + else L + ) + else: + record_len = L + replayssm_d = torch.zeros( + size=(num_mamba_layers, num_slots, hv, L, v_dim), + dtype=ring_dtype, + device=device, + ) + replayssm_k = torch.zeros( + size=(num_mamba_layers, num_slots, h_k, L, k_dim), + dtype=ring_dtype, + device=device, + ) # The log-decay gate ring (fp32): per-head SCALAR for the GDN - # gate -> [.., L]; per-K VECTOR for the KDA gate -> [.., L, K] - # (k_dim == temporal_state_shape[-1] for both). + # gate -> [.., record_len]; per-K VECTOR for the KDA gate -> + # [.., record_len, K] (k_dim == temporal_state_shape[-1] for both). g_shape = ( - (num_mamba_layers, num_slots, hv, L, k_dim) + (num_mamba_layers, num_slots, hv, record_len, k_dim) if cache_params.is_kda - else (num_mamba_layers, num_slots, hv, L) + else (num_mamba_layers, num_slots, hv, record_len) ) replayssm_g = torch.zeros( size=g_shape, @@ -617,32 +627,18 @@ class MambaPool: # (bit-identical to the recurrent baseline) instead of folding # the chunked `d` records open-loop. if enable_gdn_replayssm_spec: - # Backstop for the spec-verify ring invariants; this pool - # is sized with the final adaptive-aware draft maximum. - if L & (L - 1) != 0: - raise ValueError( - f"spec-verify ring length must be a power of two, got {L}" - ) - if ( - speculative_num_draft_tokens is not None - and L < 2 * speculative_num_draft_tokens - ): - raise ValueError( - f"spec-verify ring too small: {L} < " - f"2 * {speculative_num_draft_tokens} (early-flush margin)" - ) replayssm_rawv = torch.zeros( - size=(num_mamba_layers, num_slots, hv, L, v_dim), + size=(num_mamba_layers, num_slots, hv, record_len, v_dim), dtype=conv_dtype, device=device, ) replayssm_rawk = torch.zeros( - size=(num_mamba_layers, num_slots, h_k, L, k_dim), + size=(num_mamba_layers, num_slots, h_k, record_len, k_dim), dtype=conv_dtype, device=device, ) replayssm_beta = torch.zeros( - size=(num_mamba_layers, num_slots, hv, L), + size=(num_mamba_layers, num_slots, hv, record_len), dtype=torch.float32, device=device, ) @@ -790,10 +786,10 @@ class MambaPool: ) if _replayssm_on: logger.info( - f"GDN ReplaySSM ring buffers allocated (L=" - f"{linear_replayssm_cache_len}): " - f"d={get_tensor_size_bytes(replayssm_d) / GB:.3f}GB, " - f"k={get_tensor_size_bytes(replayssm_k) / GB:.3f}GB, " + f"GDN ReplaySSM ring buffers allocated " + f"(record_len={record_len}, fold={self.replayssm_spec_fold}): " + f"d={get_tensor_size_bytes(replayssm_d) / GB if replayssm_d is not None else 0.0:.3f}GB, " + f"k={get_tensor_size_bytes(replayssm_k) / GB if replayssm_k is not None else 0.0:.3f}GB, " f"g={get_tensor_size_bytes(replayssm_g) / GB:.3f}GB " + ( f"rawv={get_tensor_size_bytes(replayssm_rawv) / GB:.3f}GB, " @@ -813,23 +809,21 @@ class MambaPool: # the worker (spec-verify ring). Index 0..size; reset on slot (re)alloc. self.replayssm_write_pos = ( torch.zeros((size + 1,), dtype=torch.int32, device=device) - if _replayssm_on + if _replayssm_on and not self.replayssm_spec_fold else None ) # ReplaySSM spec-verify (Part B of #28511) extra per-slot cursors. The # circular ring's rolling origin (cache_base) + the per-slot flush flag # (is_flush). Block-keyed (indexed by the physical mamba slot), shared by # all GDN layers of one verify step; advanced by commit_gdn_replayssm_spec. - # Only allocated for the spec-verify ring (the decode ring does not use - # a circular buffer); None otherwise. self.replayssm_cache_base = ( torch.zeros((size + 1,), dtype=torch.int32, device=device) - if enable_gdn_replayssm_spec + if enable_gdn_replayssm_spec and not self.replayssm_spec_fold else None ) self.replayssm_is_flush = ( torch.zeros((size + 1,), dtype=torch.int8, device=device) - if enable_gdn_replayssm_spec + if enable_gdn_replayssm_spec and not self.replayssm_spec_fold else None ) mem_usage_bytes = self.mamba_cache.mem_usage_bytes() diff --git a/python/sglang/srt/mem_cache/unified_memory_pool.py b/python/sglang/srt/mem_cache/unified_memory_pool.py index 9e7cbf0df..4fcc7f4ee 100644 --- a/python/sglang/srt/mem_cache/unified_memory_pool.py +++ b/python/sglang/srt/mem_cache/unified_memory_pool.py @@ -542,6 +542,7 @@ class UnifiedMambaPool(MambaPool): self.replayssm_write_pos = None self.replayssm_is_kda = False self.enable_gdn_replayssm_spec = False + self.replayssm_spec_fold = False self.replayssm_cache_base = None self.replayssm_is_flush = None self.debug_memory_pool = False diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index b90f30c9d..2740aa872 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2542,14 +2542,13 @@ class ServerArgs: NS("exec.mamba"), ] = 16 # ReplaySSM spec-verify (Part B of RFC #28511): GDN linear-chain target-verify - # via a per-slot circular (d, k, g) ring + periodic flush instead of per-draft - # full-state snapshots. GDN only; linear-chain (topk <= 1) only. Reuses the - # `linear_replayssm` ring (replayssm_d/k/g + write_pos) and adds two per-slot - # cursors (cache_base, is_flush); the ring length reuses - # `linear_replayssm_cache_len`. + # via fold-every-commit instead of per-draft full-state snapshots -- the + # verify stores each draft step's raw inputs into a per-slot window and the + # commit replays the accepted prefix into the fp32 checkpoint. GDN only; + # linear-chain (topk <= 1) only. enable_gdn_replayssm_spec: A[ bool, - "Enable the ReplaySSM GDN spec-verify kernel (Part B of RFC #28511): a per-slot circular (d, k, g) ring + periodic flush replacing the recurrent verify's per-draft full-state snapshots. GDN only, linear-chain (--speculative-eagle-topk in {None, 1}) only. Reuses --linear-replayssm-cache-len for the ring length.", + "Enable the ReplaySSM GDN spec-verify (Part B of RFC #28511): fold-every-commit -- a per-slot raw-input window sized to the draft maximum replaces the recurrent verify's per-draft full-state snapshots. GDN only, linear-chain (--speculative-eagle-topk in {None, 1}) only.", NS("exec.mamba"), ] = False @@ -3537,7 +3536,6 @@ class ServerArgs: handle_speculative_decoding(self) # Needs the draft-token count derived just above. - self._validate_gdn_replayssm_spec_ring() # Validate the CuteDSL A2A token budget now that num_tokens_per_req is final. self._validate_cutedsl_a2a_token_budget() @@ -6061,32 +6059,29 @@ class ServerArgs: "EAGLE tree verify. Got " f"--speculative-eagle-topk={self.speculative_eagle_topk!r}." ) - if decode != "triton": + if decode not in ("triton", "flashinfer"): raise ValueError( - "--enable-gdn-replayssm-spec requires the Triton linear-attn " - "decode backend, got " + "--enable-gdn-replayssm-spec requires the triton or " + "flashinfer linear-attn decode backend, got " f"--linear-attn-decode-backend={decode!r}." ) - if self.enable_mamba_extra_buffer(): - # The spec-verify path does not yet implement the device-side - # force-flush needed to keep `temporal` consistent with the ring at - # radix mamba-track boundaries, so it is incompatible with - # extra_buffer (radix prefix caching). + # The auto->extra_buffer strategy resolution is still a declaration + # here, so read it through the resolved view. + view = self._resolved() + if ( + view.disable_radix_cache is False + and view.mamba_radix_cache_strategy == "extra_buffer_lazy" + ): raise ValueError( - "--enable-gdn-replayssm-spec is not yet compatible with mamba " - "extra_buffer (radix prefix caching); use --disable-radix-cache " - "or --mamba-radix-cache-strategy no_buffer." + "--enable-gdn-replayssm-spec is not validated with " + "--mamba-radix-cache-strategy extra_buffer_lazy yet; " + "use extra_buffer." ) - if self.disaggregation_mode != "null": + if self.disaggregation_mode == "prefill": raise ValueError( - "--enable-gdn-replayssm-spec is not supported under PD " - "disaggregation yet (follow-up). Got " - f"--disaggregation-mode={self.disaggregation_mode!r}." - ) - if self.linear_replayssm_cache_len < 1: - raise ValueError( - "--linear-replayssm-cache-len must be >= 1, got " - f"{self.linear_replayssm_cache_len}." + "--enable-gdn-replayssm-spec is not supported on a PD " + "prefill server: the ring is spec-verify-only scratch and " + "the prefill server never runs spec verify." ) if self.enable_linear_replayssm: raise ValueError( @@ -6095,57 +6090,23 @@ class ServerArgs: "with incompatible cursor protocols (per-decode-forward vs " "per-verify-commit advance)." ) - ring_len = self.linear_replayssm_cache_len - if ring_len & (ring_len - 1) != 0: - raise ValueError( - "--linear-replayssm-cache-len must be a power of two for the " - f"circular spec-verify ring, got {ring_len}." - ) - # ring_len >= 2 * max drafts is checked in - # _validate_gdn_replayssm_spec_ring() (draft tokens not derived yet). - # Closed-loop exact fold: the flush replays raw ring inputs through - # the recurrent update into the checkpoint, bit-identical to the - # recurrent baseline -- which keeps its state in fp32. A 16-bit - # checkpoint would re-quantize the exactly-folded state every flush - # and become the dominant residual error source, so require fp32. if self.mamba_ssm_dtype is None: logger.info( "--enable-gdn-replayssm-spec: setting --mamba-ssm-dtype " - "float32 (the closed-loop exact fold requires the fp32 SSM " - "checkpoint for recurrent-parity)." + "float32 (the closed-loop exact fold keeps the SSM checkpoint " + "bit-identical to the recurrent baseline)." ) self.mamba_ssm_dtype = "float32" elif self.mamba_ssm_dtype != "float32": - raise ValueError( - "--enable-gdn-replayssm-spec requires --mamba-ssm-dtype " - f"float32, got {self.mamba_ssm_dtype!r}. The closed-loop " - "exact fold keeps the committed state bit-identical to the " - "recurrent baseline, which is only meaningful against the " - "fp32 checkpoint; a 16-bit checkpoint would re-quantize it " - "every flush." + logger.warning( + "--enable-gdn-replayssm-spec with --mamba-ssm-dtype=%s: the " + "closed-loop fold re-quantizes the committed state each " + "commit/flush (fp32 keeps it bit-exact to the fp32 recurrent " + "baseline), so it may drift over long sequences. Validate " + "accuracy for your model.", + self.mamba_ssm_dtype, ) - def _validate_gdn_replayssm_spec_ring(self): - """Enforce ring_len >= 2 * max draft tokens for the spec-verify ring. - - Early-flush margin: write_pos + spec_len <= ring_len must hold on every - verify step (see _advance_gdn_spec_cursors_kernel). Runs after - handle_speculative_decoding() so the (adaptive-aware) max is final; - MambaPool re-checks at ring allocation as a backstop. - """ - if not self.enable_gdn_replayssm_spec: - return - max_drafts = self.max_speculative_num_draft_tokens - if max_drafts is None: - return - ring_len = self.linear_replayssm_cache_len - if ring_len < 2 * max_drafts: - raise ValueError( - "--linear-replayssm-cache-len must be >= 2 * the maximum " - "speculative draft-token count for the spec-verify ring " - f"(early-flush margin), got {ring_len} < {2 * max_drafts}." - ) - def _handle_legacy_cp_arguments(self): legacy_mode_to_strategy = { "in-seq-split": "zigzag", diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 7eb35b982..ce0e9e6ae 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -778,6 +778,51 @@ def prepare_mamba_track_for_verify(batch: ScheduleBatch) -> None: batch.mamba_track_seqlens = None +def _verify_commit_step_indices( + *, + batch: ScheduleBatch, + accept_index: torch.Tensor, + accept_lens: torch.Tensor, + draft_token_num: int, +) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Step indices for a post-verify state commit: per req, the tree step of + the last accepted node (reduces to accept_lens - 1 for topk == 1), and the + mamba-track interval-crossing step (-1 = no crossing; None when tracking + is off).""" + bs = accept_lens.shape[0] + accept_indices_offset = torch.arange( + 0, + bs * draft_token_num, + step=draft_token_num, + dtype=accept_lens.dtype, + device=accept_lens.device, + ) + req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device) + last_correct_step_indices = ( + accept_index[req_idx, (accept_lens - 1).to(torch.int64)] - accept_indices_offset + ) + if batch.mamba_track_indices is None: + return last_correct_step_indices, None + seq_lens_pre_verify = batch.seq_lens + seq_lens_post_verify = batch.seq_lens + accept_lens + mamba_track_interval = get_server_args().mamba_track_interval + to_track_mask = ( + seq_lens_pre_verify // mamba_track_interval + != seq_lens_post_verify // mamba_track_interval + ) + tracking_point = seq_lens_post_verify // mamba_track_interval * mamba_track_interval + to_track_ith = torch.clamp(tracking_point - seq_lens_pre_verify - 1, min=0).to( + torch.int64 + ) + candidate_track_steps = accept_index[req_idx, to_track_ith] - accept_indices_offset + mamba_steps_to_track = torch.where( + to_track_mask, + candidate_track_steps, + torch.full_like(candidate_track_steps, -1), + ) + return last_correct_step_indices, mamba_steps_to_track + + def commit_mamba_states_after_verify( target_worker: TpModelWorker, batch: ScheduleBatch, @@ -810,6 +855,40 @@ def commit_mamba_states_after_verify( # ring is allocated only then; KDA never allocates the cursors. req_pool = model_runner.req_to_token_pool mamba_pool = getattr(req_pool, "mamba_pool", None) + + # Fold-every-commit: replay the accepted prefix from the ring into + # `temporal`; the same fold stores the interval-crossing state to the + # track slot, so no SSM scatter or force-flush is needed here. + if ( + mamba_pool is not None + and getattr(mamba_pool, "replayssm_spec_fold", False) + and not getattr(mamba_pool, "replayssm_is_kda", False) + ): + if batch.forward_mode.is_idle() or accept_index.numel() == 0: + return + from sglang.kernels.ops.attention.fla.gdn_replayssm_spec_fold import ( + commit_gdn_replayssm_fold_after_verify, + ) + + spec_state = req_pool.get_speculative_mamba2_params_all_layers() + state_batch_indices = req_pool.get_mamba_indices(batch.req_pool_indices) + last_correct_step_indices, mamba_steps_to_track = _verify_commit_step_indices( + batch=batch, + accept_index=accept_index, + accept_lens=accept_lens, + draft_token_num=draft_token_num, + ) + commit_gdn_replayssm_fold_after_verify( + spec_state=spec_state, + state_batch_indices=state_batch_indices, + accept_lens=accept_lens, + last_correct_step_indices=last_correct_step_indices, + mamba_track_indices=batch.mamba_track_indices, + mamba_steps_to_track=mamba_steps_to_track, + null_block_id=-1, + ) + return + if ( mamba_pool is not None and getattr(mamba_pool, "replayssm_cache_base", None) is not None @@ -841,17 +920,11 @@ def commit_mamba_states_after_verify( ) # Roll back / commit the conv state to the last accepted draft step # (same logic as the recurrent commit, but conv-only). - accept_indices_offset = torch.arange( - 0, - bs * draft_token_num, - step=draft_token_num, - dtype=accept_lens.dtype, - device=accept_lens.device, - ) - req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device) - last_correct_step_indices = ( - accept_index[req_idx, (accept_lens - 1).to(torch.int64)] - - accept_indices_offset + last_correct_step_indices, _ = _verify_commit_step_indices( + batch=batch, + accept_index=accept_index, + accept_lens=accept_lens, + draft_token_num=draft_token_num, ) fused_conv_window_scatter_with_mask( spec_state.conv[0], @@ -871,47 +944,12 @@ def commit_mamba_states_after_verify( bs = accept_lens.shape[0] # `accept_lens` already includes the bonus token (drafts + 1 per req). if not batch.forward_mode.is_idle() and accept_index.numel() > 0: - accept_indices_offset = torch.arange( - 0, - bs * draft_token_num, - step=draft_token_num, - dtype=accept_lens.dtype, - device=accept_lens.device, + last_correct_step_indices, mamba_steps_to_track = _verify_commit_step_indices( + batch=batch, + accept_index=accept_index, + accept_lens=accept_lens, + draft_token_num=draft_token_num, ) - req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device) - # Per-req tree step of the last accepted node, i.e. the step whose - # mamba state to commit; reduces to accept_lens - 1 for topk == 1. - last_correct_step_indices = ( - accept_index[req_idx, (accept_lens - 1).to(torch.int64)] - - accept_indices_offset - ) - - if batch.mamba_track_indices is not None: - # If after verify, the request's seq_lens has crossed a mamba track interval, - # we need to update the mamba state for the request at the crossing point. - seq_lens_pre_verify = batch.seq_lens - seq_lens_post_verify = batch.seq_lens + accept_lens - mamba_track_interval = get_server_args().mamba_track_interval - to_track_mask = ( - seq_lens_pre_verify // mamba_track_interval - != seq_lens_post_verify // mamba_track_interval - ) - tracking_point = ( - seq_lens_post_verify // mamba_track_interval * mamba_track_interval - ) - to_track_ith = torch.clamp( - tracking_point - seq_lens_pre_verify - 1, min=0 - ).to(torch.int64) - candidate_track_steps = ( - accept_index[req_idx, to_track_ith] - accept_indices_offset - ) - mamba_steps_to_track = torch.where( - to_track_mask, - candidate_track_steps, - torch.full_like(candidate_track_steps, -1), - ) - else: - mamba_steps_to_track = None if hasattr(attn_backend, "update_mamba_state_after_mtp_verify"): attn_backend.update_mamba_state_after_mtp_verify( diff --git a/test/registered/attention/unittests/gdn/test_gdn_replayssm_spec_fold.py b/test/registered/attention/unittests/gdn/test_gdn_replayssm_spec_fold.py new file mode 100644 index 000000000..a75965039 --- /dev/null +++ b/test/registered/attention/unittests/gdn/test_gdn_replayssm_spec_fold.py @@ -0,0 +1,234 @@ +"""GDN ReplaySSM fold-every-commit: fused ring-write + commit fold. + +The design contract is BITWISE parity with the recurrent verify + per-draft +snapshot baseline, so every case asserts ``torch.equal``: ring-write leaves the +verify output unchanged; the folded checkpoint equals the accepted-step +snapshot (fp32 and bf16, each vs the same-dtype baseline); HAS_TRACK stores the +crossing-step state and skips -1 steps / null slots; a 256-iteration +verify->commit chain stays bitwise equal at every step (no accumulation channel +on the state path -- the long-decode drift failure mode). +""" + +import sys +import unittest +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[4])) + +from sglang.kernels.ops.attention.fla.fused_sigmoid_gating_recurrent import ( + fused_sigmoid_gating_delta_rule_update, +) +from sglang.kernels.ops.attention.fla.gdn_replayssm_spec_fold import ( + commit_gdn_replayssm_fold_all_layers, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large") + +B, T = 3, 4 +H, HV = 4, 8 +K = V = 64 +NUM_SLOTS = 8 +DEVICE = "cuda" + + +def _make_window(step_seed: int): + gen = torch.Generator(device=DEVICE).manual_seed(step_seed) + + def rand(*shape, dtype=torch.bfloat16): + return torch.randn(*shape, device=DEVICE, dtype=dtype, generator=gen) + + return { + "q": rand(1, B * T, H, K), + "k": rand(1, B * T, H, K), + "v": rand(1, B * T, HV, V), + "a": rand(B * T, HV), + "b": rand(B * T, HV), + } + + +def _run_verify(inputs, gating, state, slots, *, snapshots=None, rings=None): + kwargs = {} + if snapshots is not None: + kwargs.update( + intermediate_states_buffer=snapshots, + intermediate_state_indices=slots, + cache_steps=T, + ) + if rings is not None: + # Per-layer views, matching the backend's mamba2_layer_cache slices. + kwargs.update( + cache_ring=True, + replayssm_rawv=rings["rawv"][0], + replayssm_rawk=rings["rawk"][0], + replayssm_g=rings["g"][0], + replayssm_beta=rings["beta"][0], + ) + cu_seqlens = torch.arange(0, B * T + 1, step=T, dtype=torch.int32, device=DEVICE) + return fused_sigmoid_gating_delta_rule_update( + A_log=gating["A_log"], + dt_bias=gating["dt_bias"], + softplus_beta=1.0, + softplus_threshold=20.0, + q=inputs["q"], + k=inputs["k"], + v=inputs["v"], + b=inputs["b"], + a=inputs["a"], + initial_state_source=state, + initial_state_indices=slots, + cu_seqlens=cu_seqlens, + use_qk_l2norm_in_kernel=True, + is_kda=False, + disable_state_update=True, + **kwargs, + ) + + +def _make_rings(dtype=torch.bfloat16): + return { + "rawv": torch.zeros(1, NUM_SLOTS, HV, T, V, device=DEVICE, dtype=dtype), + "rawk": torch.zeros(1, NUM_SLOTS, H, T, K, device=DEVICE, dtype=dtype), + "g": torch.zeros(1, NUM_SLOTS, HV, T, device=DEVICE, dtype=torch.float32), + "beta": torch.zeros(1, NUM_SLOTS, HV, T, device=DEVICE, dtype=torch.float32), + } + + +def _fold(state, rings, slots, accept_lens, track_slots=None, track_steps=None): + commit_gdn_replayssm_fold_all_layers( + checkpoint_state=state, + rawv_cache=rings["rawv"], + rawk_cache=rings["rawk"], + g_cache=rings["g"], + beta_cache=rings["beta"], + ssm_state_indices=slots, + accept_lens=accept_lens, + max_cache_len=T, + num_k_heads=H, + mamba_track_indices=track_slots, + mamba_steps_to_track=track_steps, + null_block_id=-1, + ) + + +class TestGdnReplayssmSpecFold(CustomTestCase): + def setUp(self): + torch.manual_seed(0) + self.gating = { + "A_log": torch.randn(HV, device=DEVICE) * 0.1, + "dt_bias": torch.randn(HV, device=DEVICE) * 0.1, + } + self.slots = torch.tensor([5, 2, 7], dtype=torch.int32, device=DEVICE) + self.accept_lens = torch.tensor([3, 1, 4], dtype=torch.int32, device=DEVICE) + + def _state(self, dtype): + gen = torch.Generator(device=DEVICE).manual_seed(1) + return torch.randn( + NUM_SLOTS, HV, K, V, device=DEVICE, dtype=dtype, generator=gen + ) + + def test_ring_write_does_not_change_verify_output(self): + for dtype in (torch.float32, torch.bfloat16): + state = self._state(dtype) + inputs = _make_window(11) + out_plain = _run_verify(inputs, self.gating, state.clone(), self.slots) + out_ring = _run_verify( + inputs, self.gating, state.clone(), self.slots, rings=_make_rings() + ) + self.assertTrue(torch.equal(out_plain, out_ring), f"{dtype=}") + + def test_fold_matches_snapshot_baseline_bitwise(self): + for dtype in (torch.float32, torch.bfloat16): + state = self._state(dtype) + inputs = _make_window(22) + + snapshots = torch.zeros(NUM_SLOTS, T, HV, K, V, device=DEVICE, dtype=dtype) + _run_verify( + inputs, self.gating, state.clone(), self.slots, snapshots=snapshots + ) + + fold_state = state.clone().unsqueeze(0).contiguous() + rings = _make_rings() + _run_verify(inputs, self.gating, fold_state[0], self.slots, rings=rings) + _fold(fold_state, rings, self.slots, self.accept_lens) + + for s, n in zip(self.slots.tolist(), self.accept_lens.tolist()): + self.assertTrue( + torch.equal(snapshots[s, n - 1], fold_state[0, s]), + f"{dtype=} slot={s} accept_len={n}", + ) + untouched = set(range(NUM_SLOTS)) - set(self.slots.tolist()) + for s in untouched: + self.assertTrue(torch.equal(fold_state[0, s], state[s])) + + def test_track_store_and_null_slots(self): + dtype = torch.float32 + state = self._state(dtype) + inputs = _make_window(33) + + snapshots = torch.zeros(NUM_SLOTS, T, HV, K, V, device=DEVICE, dtype=dtype) + _run_verify(inputs, self.gating, state.clone(), self.slots, snapshots=snapshots) + + fold_state = state.clone().unsqueeze(0).contiguous() + rings = _make_rings() + _run_verify(inputs, self.gating, fold_state[0], self.slots, rings=rings) + + track_slots = torch.tensor([1, 0, 3], dtype=torch.int64, device=DEVICE) + track_steps = torch.tensor([1, -1, 2], dtype=torch.int64, device=DEVICE) + slots_with_null = self.slots.clone() + slots_with_null[1] = -1 + _fold( + fold_state, + rings, + slots_with_null, + self.accept_lens, + track_slots=track_slots, + track_steps=track_steps, + ) + + self.assertTrue(torch.equal(fold_state[0, 1], snapshots[5, 1])) + self.assertTrue(torch.equal(fold_state[0, 3], snapshots[7, 2])) + # Null slot: neither committed nor tracked (track step 1 is masked + # to -1 only for row 1's -1 step; row 1's slot itself was nulled). + self.assertTrue(torch.equal(fold_state[0, 2], state[2])) + self.assertTrue(torch.equal(fold_state[0, 0], state[0])) + + def test_long_chain_no_accumulation(self): + """256 chained verify->commit iterations stay bitwise equal to the + baseline chain at every iteration (fp32 + bf16): no error channel + can accumulate with sequence length.""" + num_iters = 256 + for dtype in (torch.float32, torch.bfloat16): + base_state = self._state(dtype) + fold_state = base_state.clone().unsqueeze(0).contiguous() + snapshots = torch.zeros(NUM_SLOTS, T, HV, K, V, device=DEVICE, dtype=dtype) + gen = torch.Generator().manual_seed(7) + for it in range(num_iters): + inputs = _make_window(1000 + it) + accept_lens = torch.randint(1, T + 1, (B,), generator=gen).to( + device=DEVICE, dtype=torch.int32 + ) + + out_base = _run_verify( + inputs, self.gating, base_state, self.slots, snapshots=snapshots + ) + for s, n in zip(self.slots.tolist(), accept_lens.tolist()): + base_state[s] = snapshots[s, n - 1] + + rings = _make_rings() + out_fold = _run_verify( + inputs, self.gating, fold_state[0], self.slots, rings=rings + ) + _fold(fold_state, rings, self.slots, accept_lens) + + self.assertTrue(torch.equal(out_base, out_fold), f"{dtype=} {it=}") + self.assertTrue( + torch.equal(base_state, fold_state[0]), f"{dtype=} {it=}" + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_replayssm_ring_accounting.py b/test/registered/unit/mem_cache/test_replayssm_ring_accounting.py new file mode 100644 index 000000000..6376126f5 --- /dev/null +++ b/test/registered/unit/mem_cache/test_replayssm_ring_accounting.py @@ -0,0 +1,67 @@ +"""ReplaySSM ring per-slot byte accounting (BaseLinearStateParams.replayssm_ring_bytes_per_req). + +The memory solver charges this on top of mamba_cache_per_req so num_slots is not +over-provisioned (the ring is allocated per slot but is NOT part of the state +cache cost). Pins the arithmetic against hand-computed byte counts for the +fold window (raw v / pre-norm k / g / beta). If the MambaPool allocation +changes shape, update both together. +""" + +import pytest +import torch + +from sglang.srt.configs.mamba_utils import ( + Mamba2CacheParams, + Mamba2StateDType, + Mamba2StateShape, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +# temporal = (hv=4, v_dim=8, k_dim=8), num_k_heads_per_tp = 4, record_len = 8, +# 2 layers. conv bf16 (2B), fp32 gate/beta (4B). Ring tensors (per slot, per +# layer): +# rawv hv*RL*v_dim, rawk h_k*RL*k_dim -> conv dtype +# g hv*RL -> fp32 +# beta hv*RL -> fp32 +DTYPE = Mamba2StateDType(conv=torch.bfloat16, temporal=torch.float32) +RL = 8 +LAYERS = [0, 1] + + +def _gdn_params(): + # Only shape.temporal and shape.num_k_heads_per_tp are read here; the rest + # are dummy (the accounting does not depend on them). + shape = Mamba2StateShape( + conv=[(4, 3)], + temporal=(4, 8, 8), + intermediate_size=0, + conv_dim=0, + ssm_state_size=0, + num_heads=0, + head_dim=0, + state_size=0, + conv_kernel=0, + num_k_heads_per_tp=4, + ) + return Mamba2CacheParams(shape=shape, dtype=DTYPE, layers=LAYERS) + + +class TestReplaySSMRingAccounting(CustomTestCase): + def test_gdn_fold(self): + # fold window: rawv 512 + rawk 512 + g(scalar, 4*8*4) 128 + beta 128 = 1280 + self.assertEqual( + _gdn_params().replayssm_ring_bytes_per_req(record_len=RL), + 1280 * len(LAYERS), + ) + + def test_zero_len_ring(self): + self.assertEqual(_gdn_params().replayssm_ring_bytes_per_req(record_len=0), 0) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/spec/test_ngram_mamba_verify_update.py b/test/registered/unit/spec/test_ngram_mamba_verify_update.py index 57db43e7f..e795f3f18 100644 --- a/test/registered/unit/spec/test_ngram_mamba_verify_update.py +++ b/test/registered/unit/spec/test_ngram_mamba_verify_update.py @@ -99,6 +99,9 @@ class TestNgramMambaVerifyUpdate(CustomTestCase): target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify = ( MagicMock() ) + mamba_pool = target_worker.model_runner.req_to_token_pool.mamba_pool + mamba_pool.replayssm_spec_fold = False + mamba_pool.replayssm_cache_base = None return target_worker def test_mamba_verify_update_called_with_correct_indices(self):