diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index 45b1038e0..6bf623a56 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -288,6 +288,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac check_environments() initialize_linear_attn_config(runner.server_args) + hybrid_backend_cls = HybridLinearAttnBackend if runner.hybrid_gdn_config is not None: if is_blackwell(): assert ( @@ -303,7 +304,46 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac logger.info(f"Using hybrid linear attention backend for hybrid GDN models.") linear_attn_backend = GDNAttnBackend(runner) elif runner.mamba2_config is not None: - linear_attn_backend = Mamba2AttnBackend(runner) + from sglang.srt.configs.lfm2 import Lfm2Config + from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig + from sglang.srt.configs.lfm2_vl import Lfm2VlConfig + from sglang.srt.configs.zaya import ZayaConfig + + # Short-conv hybrids (ZAYA1 CCA, LFM2 short conv) share a conv-state + # sidecar that owns the per-request state plumbing and is invoked by + # the model via conv_state_metadata (never as a full-vs-linear + # alternative). Other mamba2 models keep the full Mamba2 SSM backend. + short_conv_cfgs = ( + ZayaConfig, + Lfm2Config, + Lfm2MoeConfig, + Lfm2VlConfig, + ) + if isinstance(runner.mamba2_config, short_conv_cfgs): + if is_npu(): + # The model conv layers call + # get_attn_backend().conv_state_metadata() unconditionally, + # but the Ascend hybrid/mamba backend has no such method. + # Fail here (before model execution) with a clear message + # rather than an AttributeError deep in the first conv layer. + raise NotImplementedError( + "Short-conv hybrid models (ZAYA1 CCA, LFM2 / LFM2-MoE) " + "are not yet supported on NPU: the conv-state sidecar " + "(ShortConvAttnBackend.conv_state_metadata) has no Ascend " + "implementation. Add an Ascend conv-state backend before " + "serving these models on NPU." + ) + from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( + ShortConvHybridAttnBackend, + ) + from sglang.srt.layers.attention.linear.short_conv_backend import ( + ShortConvAttnBackend, + ) + + linear_attn_backend = ShortConvAttnBackend(runner) + hybrid_backend_cls = ShortConvHybridAttnBackend + else: + linear_attn_backend = Mamba2AttnBackend(runner) elif runner.kimi_linear_config is not None: linear_attn_backend = KDAAttnBackend(runner) elif runner.hybrid_lightning_config is not None: @@ -325,7 +365,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac full_attn_layers = [0] else: full_attn_layers = cfg.full_attention_layer_ids - return HybridLinearAttnBackend( + return hybrid_backend_cls( full_attn_backend, linear_attn_backend, full_attn_layers ) diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 93df15c6d..4dd1f99cb 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -1096,3 +1096,31 @@ class HybridLinearAttnBackend(AttentionBackend): mamba_track_indices, mamba_steps_to_track, ) + + +class ShortConvHybridAttnBackend(HybridLinearAttnBackend): + """HybridLinearAttnBackend variant for short-conv hybrid models (ZAYA1 CCA, + LFM2 short conv). + + The linear sidecar is a :class:`ShortConvAttnBackend + ` + that owns the per-request conv-state plumbing. The model's conv module + reaches it via :meth:`conv_state_metadata` (``get_attn_backend()`` returns + this wrapper) and runs its own conv kernel against the returned handle, so + the model definition holds no pool access. The sidecar is never reached + through the full-vs-linear ``forward_decode`` / ``forward_extend`` dispatch. + """ + + def __init__( + self, + full_attn_backend: AttentionBackend, + short_conv_backend: MambaAttnBackendBase, + full_attn_layers: list, + ): + # Register short_conv_backend as the linear sidecar so it rides in + # attn_backend_list and inherits the metadata / cuda-graph fan-out. + super().__init__(full_attn_backend, short_conv_backend, full_attn_layers) + self.short_conv_backend = short_conv_backend + + def conv_state_metadata(self, layer_id: int, forward_batch: ForwardBatch): + return self.short_conv_backend.conv_state_metadata(layer_id, forward_batch) diff --git a/python/sglang/srt/layers/attention/linear/short_conv_backend.py b/python/sglang/srt/layers/attention/linear/short_conv_backend.py new file mode 100644 index 000000000..fba352f3b --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/short_conv_backend.py @@ -0,0 +1,219 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Short-convolution attention backend. + +Several hybrid models interleave a *causal short conv with per-request conv +state* (stored in the centralized ``MambaPool``) with softmax attention layers: + +* **LFM2** (:class:`Lfm2ShortConv `) -- + a depthwise gated short conv (``causal_conv1d_fn`` / ``causal_conv1d_update``) + as a standalone token mixer on its own conv layers. +* **ZAYA1** (:class:`CCA `) -- a two-stage grouped + conv plus a one-token ``prev_hs`` lag, preprocessing q/k for the layer's + softmax attention. + +These share the *state plumbing* -- resolving the per-request slot indices, the +``has_initial_state`` prefix mask, the ``query_start_loc`` cu-seqlens, and the +cuda-graph static index buffers, all once per forward step -- but NOT the conv +kernel itself. ``ShortConvAttnBackend`` owns only the plumbing and hands it out +via :meth:`conv_state_metadata` as a :class:`ShortConvMetadata`; each model runs +its own conv kernel against that handle, so the model definition holds no pool +access. + +The backend is a *sidecar*: it is invoked directly by the model (through +:class:`ShortConvHybridAttnBackend +`), +never through the full-vs-linear ``forward_decode`` / ``forward_extend`` +dispatch. Metadata + cuda-graph capture/replay come from +:class:`MambaAttnBackendBase`. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, List, NamedTuple, Optional + +import torch + +from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( + MambaAttnBackendBase, +) +from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +if TYPE_CHECKING: + from sglang.srt.model_executor.model_runner import ModelRunner + + +class ShortConvMetadata(NamedTuple): + """Per-(layer, step) conv-state handle handed to a model's conv kernel. + + ``layer_cache`` exposes the per-layer pool views (``conv[0]`` = conv state, + ``conv[1]`` = an optional second state such as ZAYA1's ``prev_hs``, + ``temporal`` = SSM state, unused by pure short convs). The device tensors are + cuda-graph-static on the decode/replay path; the ``*_cpu`` host mirrors are + built once per step only for models whose extend path runs a host loop + (e.g. ZAYA1 v1) and are ``None`` on decode. + """ + + layer_cache: Any + cache_indices: torch.Tensor + # cu-seqlens for the varlen prefill conv (device, int32). None on decode. + query_start_loc: Optional[torch.Tensor] = None + # Per-request "resumes a cached prefix" mask (device bool). None on decode. + has_initial_state: Optional[torch.Tensor] = None + # Host mirror of cache_indices for extend host loops. None on decode. + slot_ids_cpu: Optional[List[int]] = None + # Host mirror of has_initial_state for extend host loops. None on decode. + has_prefix_cpu: Optional[List[bool]] = None + + +class ShortConvAttnBackend(MambaAttnBackendBase): + """Owns the short-conv per-request state plumbing (see module docstring).""" + + # State IO is index-driven; no host seq-lens plumbing required from the + # runner. (The extend path reads ``extend_*_cpu`` off the batch, which is + # always populated for extend regardless of this flag.) + needs_cpu_seq_lens: bool = False + + def __init__(self, model_runner: ModelRunner): + super().__init__(model_runner) + mamba_cache = self.req_to_token_pool.mamba_pool.mamba_cache + # conv[0] == conv_state: [n_layers, n_slots, conv_dim, conv_kernel - 1] + self.conv_states_shape = mamba_cache.conv[0].shape + + # Per-step state, resolved ONCE per step in init_forward_metadata / + # init_forward_metadata_out_graph (never per conv layer). The extend host + # mirrors drive the extend loop; ``_cache_indices`` is the int64 slot + # index view shared by all conv layers within the step. + self._has_initial_state: Optional[torch.Tensor] = None + self._slot_ids_cpu: Optional[List[int]] = None + self._has_prefix_cpu: Optional[List[bool]] = None + self._cache_indices: Optional[torch.Tensor] = None + self._cache_indices_buf: Optional[torch.Tensor] = None + + def _reset_step_state(self): + self._has_initial_state = None + self._slot_ids_cpu = None + self._has_prefix_cpu = None + + def _alloc_cache_indices_buf(self, max_bs: int): + # Persistent int64 index buffer, refilled in place per step so the + # captured (cuda or cpu) graph reads a stable address. + self._cache_indices_buf = torch.empty( + max_bs, dtype=torch.int64, device=self.device + ) + + def _refresh_cache_indices(self): + # Resolve the int64 slot-index view ONCE per step, shared by every conv + # layer. When a graph index buffer is allocated and large enough, refill + # it IN PLACE and hand out a view -- the captured graph then reads a + # stable address that this (pre-replay) hook keeps current, so it is + # cuda- and cpu-graph safe. Otherwise (eager, or bs beyond the buffer) + # a fresh cast is fine. + md = self.forward_metadata + idx = md.mamba_cache_indices if md is not None else None + buf = self._cache_indices_buf + if idx is None: + self._cache_indices = None + elif buf is not None and idx.shape[0] <= buf.shape[0]: + n = idx.shape[0] + buf[:n].copy_(idx) + self._cache_indices = buf[:n] + else: + self._cache_indices = idx.to(torch.long) + + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): + super().init_cuda_graph_state(max_bs, max_num_tokens) + self._alloc_cache_indices_buf(max_bs) + + def init_cpu_graph_state(self, max_bs: int, max_num_tokens: int): + super().init_cpu_graph_state(max_bs, max_num_tokens) + self._alloc_cache_indices_buf(max_bs) + + def init_forward_metadata(self, forward_batch: ForwardBatch): + # Eager path (also the CPU-graph replay path). Builds + # self.forward_metadata and runs the deferred mamba clear/COW ops. + super().init_forward_metadata(forward_batch) + self._reset_step_state() + self._refresh_cache_indices() + mode = forward_batch.forward_mode + if ( + mode.is_extend() + and not mode.is_target_verify() + and not mode.is_draft_extend_v2() + ): + self._has_initial_state = forward_batch.extend_prefix_lens > 0 + if self._cache_indices is not None: + self._slot_ids_cpu = self._cache_indices.tolist() + self._has_prefix_cpu = [ + int(p) > 0 for p in forward_batch.extend_prefix_lens_cpu + ] + + def init_forward_metadata_out_graph( + self, forward_batch: ForwardBatch, in_capture: bool = False + ): + # Decode cuda-graph capture + replay path -- no extend prefix state. + super().init_forward_metadata_out_graph(forward_batch, in_capture) + self._reset_step_state() + self._refresh_cache_indices() + + def init_forward_metadata_capture_cpu_graph(self, *args, **kwargs): + # Decode CPU-graph capture path. The base fills forward_metadata but not + # the int64 view; without this the conv layers would capture a ``None`` + # index (crash / corrupt state). Replay goes through init_forward_metadata + # and refills the SAME buffer, so the captured cpu graph reads a stable + # address kept current at replay. + super().init_forward_metadata_capture_cpu_graph(*args, **kwargs) + self._reset_step_state() + self._refresh_cache_indices() + + def conv_state_metadata( + self, layer_id: int, forward_batch: ForwardBatch + ) -> ShortConvMetadata: + """Return the conv-state handle for ``layer_id`` at the current step. + + The per-step fields are already resolved on ``self.forward_metadata`` / + ``self._*`` (in ``init_forward_metadata`` / ``_out_graph``); + ``forward_batch`` is accepted for interface parity with the unit-test + mock and is not otherwise required here. + """ + layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) + md = self.forward_metadata + + # Slot indices are cached ONCE per step in init_forward_metadata / + # init_forward_metadata_out_graph (int64). Hand back the cached view -- no + # per-layer recompute. Decode is cuda-graph-safe because that view is a + # persistent buffer refilled in place before each replay. + return ShortConvMetadata( + layer_cache=layer_cache, + cache_indices=self._cache_indices, + query_start_loc=md.query_start_loc, + has_initial_state=self._has_initial_state, + slot_ids_cpu=self._slot_ids_cpu, + has_prefix_cpu=self._has_prefix_cpu, + ) + + # The short-conv layers are invoked via conv_state_metadata + the model's own + # conv kernel, never through the HybridLinearAttnBackend full-vs-linear + # dispatch. Mirror Mamba2AttnBackend and guard the routed entrypoints. + def forward_decode(self, *args, **kwargs): + raise NotImplementedError( + "ShortConvAttnBackend is invoked via conv_state_metadata; " + "it does not run through forward_decode." + ) + + def forward_extend(self, *args, **kwargs): + raise NotImplementedError( + "ShortConvAttnBackend is invoked via conv_state_metadata; " + "it does not run through forward_extend." + ) diff --git a/python/sglang/srt/layers/attention/mamba/causal_conv1d.py b/python/sglang/srt/layers/attention/mamba/causal_conv1d.py index 5d2f0aa3a..ba94f3105 100644 --- a/python/sglang/srt/layers/attention/mamba/causal_conv1d.py +++ b/python/sglang/srt/layers/attention/mamba/causal_conv1d.py @@ -96,6 +96,8 @@ def causal_conv1d_fn( x = x.contiguous() bias = bias.contiguous() if bias is not None else None + if cache_indices is not None and cache_indices.dtype != torch.int32: + cache_indices = cache_indices.to(torch.int32) causal_conv1d_fwd( x, weight, @@ -162,6 +164,8 @@ def causal_conv1d_update( unsqueeze = x.dim() == 2 if unsqueeze: x = x.unsqueeze(-1) + if conv_state_indices is not None and conv_state_indices.dtype != torch.int32: + conv_state_indices = conv_state_indices.to(torch.int32) causal_conv1d_update_kernel( x, conv_state, diff --git a/python/sglang/srt/models/lfm2.py b/python/sglang/srt/models/lfm2.py index 3a93ddf1d..d494c6a59 100644 --- a/python/sglang/srt/models/lfm2.py +++ b/python/sglang/srt/models/lfm2.py @@ -40,7 +40,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.model_executor.forward_context import get_req_to_token_pool +from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_loader.weight_utils import ( default_weight_loader, sharded_weight_loader, @@ -265,10 +265,11 @@ class Lfm2ShortConv(nn.Module): if forward_batch.forward_mode.is_idle(): return hidden_states - layer_cache = get_req_to_token_pool().mamba2_layer_cache(self.layer_idx) - conv_state = layer_cache.conv[0] - req_pool_indices = forward_batch.req_pool_indices - mamba_indices = get_req_to_token_pool().get_mamba_indices(req_pool_indices) + # The backend owns the per-request conv-state plumbing (slot indices, + # prefix mask, cu-seqlens, cuda-graph buffers); this layer just runs its + # depthwise conv against the returned handle. + meta = get_attn_backend().conv_state_metadata(self.layer_idx, forward_batch) + conv_state = meta.layer_cache.conv[0] # Project and split into gates: B (pre-conv), C (post-conv), x (input) proj, _ = self.in_proj(hidden_states) @@ -283,40 +284,18 @@ class Lfm2ShortConv(nn.Module): self.conv_weight, self.conv_bias, activation=None, - conv_state_indices=mamba_indices.to(torch.int32), + conv_state_indices=meta.cache_indices, ) else: # Prefill: multiple tokens, use varlen kernel - T = hidden_states.shape[0] Bx_t = Bx.transpose(0, 1).contiguous() - - # Build query_start_loc: [0, cumsum(seq_lens)...] - extend_start_loc = forward_batch.extend_start_loc - if extend_start_loc is not None and len(extend_start_loc) > 1: - query_start_loc = torch.cat( - [ - extend_start_loc, - torch.tensor( - [T], dtype=torch.int32, device=hidden_states.device - ), - ] - ) - cache_indices = mamba_indices.to(torch.int32) - has_initial_state = forward_batch.extend_prefix_lens > 0 - else: - query_start_loc = torch.tensor( - [0, T], dtype=torch.int32, device=hidden_states.device - ) - cache_indices = mamba_indices[:1].to(torch.int32) - has_initial_state = forward_batch.extend_prefix_lens[:1] > 0 - conv_out = causal_conv1d_fn( Bx_t, self.conv_weight, self.conv_bias, - query_start_loc=query_start_loc, - cache_indices=cache_indices, - has_initial_state=has_initial_state, + query_start_loc=meta.query_start_loc, + cache_indices=meta.cache_indices, + has_initial_state=meta.has_initial_state, conv_states=conv_state, activation=None, ).transpose(0, 1) diff --git a/python/sglang/srt/models/lfm2_moe.py b/python/sglang/srt/models/lfm2_moe.py index b3c033a89..7e41f19eb 100644 --- a/python/sglang/srt/models/lfm2_moe.py +++ b/python/sglang/srt/models/lfm2_moe.py @@ -42,7 +42,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.model_executor.forward_context import get_req_to_token_pool +from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_loader.weight_utils import ( default_weight_loader, sharded_weight_loader, @@ -328,10 +328,11 @@ class Lfm2MoeShortConv(nn.Module): if forward_batch.forward_mode.is_idle(): return hidden_states - layer_cache = get_req_to_token_pool().mamba2_layer_cache(self.layer_idx) - conv_state = layer_cache.conv[0] - req_pool_indices = forward_batch.req_pool_indices - mamba_indices = get_req_to_token_pool().get_mamba_indices(req_pool_indices) + # The backend owns the per-request conv-state plumbing (slot indices, + # prefix mask, cu-seqlens, cuda-graph buffers); this layer just runs its + # depthwise conv against the returned handle. + meta = get_attn_backend().conv_state_metadata(self.layer_idx, forward_batch) + conv_state = meta.layer_cache.conv[0] proj, _ = self.in_proj(hidden_states) B_gate, C_gate, x = proj.chunk(3, dim=-1) @@ -344,36 +345,17 @@ class Lfm2MoeShortConv(nn.Module): self.conv_weight, self.conv_bias, activation=None, - conv_state_indices=mamba_indices.to(torch.int32), + conv_state_indices=meta.cache_indices, ) else: - T = hidden_states.shape[0] Bx_t = Bx.transpose(0, 1).contiguous() - - # Build query_start_loc for variable-length sequences - # causal_conv1d_fn expects [start0, start1, ..., startN, T] - extend_start_loc = forward_batch.extend_start_loc - if extend_start_loc is not None and len(extend_start_loc) > 1: - # Multiple sequences: append T to extend_start_loc - # Allocate and fill to avoid torch.cat overhead - query_start_loc = extend_start_loc.new_empty(len(extend_start_loc) + 1) - query_start_loc[:-1] = extend_start_loc - query_start_loc[-1] = T - cache_indices = mamba_indices.to(torch.int32) - has_initial_state = forward_batch.extend_prefix_lens > 0 - else: - # Single sequence: [0, T] - query_start_loc = hidden_states.new_tensor([0, T], dtype=torch.int32) - cache_indices = mamba_indices[:1].to(torch.int32) - has_initial_state = forward_batch.extend_prefix_lens[:1] > 0 - conv_out = causal_conv1d_fn( Bx_t, self.conv_weight, self.conv_bias, - query_start_loc=query_start_loc, - cache_indices=cache_indices, - has_initial_state=has_initial_state, + query_start_loc=meta.query_start_loc, + cache_indices=meta.cache_indices, + has_initial_state=meta.has_initial_state, conv_states=conv_state, activation=None, ).transpose(0, 1) diff --git a/python/sglang/srt/models/zaya.py b/python/sglang/srt/models/zaya.py index 9fb63238c..1f63f2019 100644 --- a/python/sglang/srt/models/zaya.py +++ b/python/sglang/srt/models/zaya.py @@ -31,9 +31,13 @@ for the full design notes): averaging across MoE layers) and MOD (mixture-of-depths skip expert). - Per-layer :class:`ResidualScaling` keeps the residual stream in fp32 with affine scale/bias both on the residual and on the post-mixer hidden states. -- Per-request CCA state (``conv_state`` + ``prev_hs``) is managed by - SGLang's centralized ``MambaPool`` inside ``HybridReqToTokenPool``, - accessed via ``get_req_to_token_pool().mamba2_layer_cache()``. +- Per-request CCA state (``conv_state`` + ``prev_hs``) lives in SGLang's + centralized ``MambaPool`` inside ``HybridReqToTokenPool``. The per-request + state plumbing (slot indices, prefix mask, cuda-graph buffers) is owned by + ``ShortConvAttnBackend`` and reached via + ``get_attn_backend().conv_state_metadata()``, so the model holds no pool + access; CCA runs its own conv (:func:`cca_extend` / :func:`cca_decode`) + against the returned handle. """ from __future__ import annotations @@ -41,7 +45,7 @@ from __future__ import annotations import logging import re from collections.abc import Iterable -from typing import Optional +from typing import List, Optional, Tuple import torch import torch.nn.functional as F @@ -70,7 +74,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors -from sglang.srt.model_executor.forward_context import get_req_to_token_pool +from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs @@ -78,14 +82,6 @@ from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs logger = logging.getLogger(__name__) -# Attribute names used to memoize the per-request MambaPool slot indices on the -# ForwardBatch. The req -> slot mapping is identical for every CCA layer in a -# step, so caching it here makes the lookup (and its GPU->CPU sync) run once per -# forward step instead of once per attention layer. -_MAMBA_INDICES_ATTR = "_zaya_mamba_indices" -_MAMBA_INDICES_CPU_ATTR = "_zaya_mamba_indices_cpu" - - # --------------------------------------------------------------------------- # Residual scaling # --------------------------------------------------------------------------- @@ -141,6 +137,187 @@ def _apply_norm_with_fp32_residual( return norm(residual.to(target_dtype)) +# --------------------------------------------------------------------------- +# CCA conv-state kernels (v1 torch) +# +# ZAYA1-specific conv step: the CCA conv is a causal two-stage conv over +# ``qk = [W_q hs || W_k hs]`` plus a one-token ``prev_hs`` lag for val_proj2. +# The per-request conv state lives in the centralized MambaPool; the backend +# (ShortConvAttnBackend) hands out the slot indices + prefix flags and CCA runs +# these functions against them. ``conv_qk`` is the module's two-stage conv; +# both functions mutate ``conv_state`` / ``prev_hs_state`` in place and return +# ``(qk_out, v2_input)`` -- the conv output ``[T, in_out_ch]`` and the (shifted) +# ``val_proj2`` input ``[T, hidden_size]``. +# --------------------------------------------------------------------------- + + +def cca_extend( + qk: torch.Tensor, + hidden_states: torch.Tensor, + conv_qk: nn.Module, + conv_state: torch.Tensor, + prev_hs_state: torch.Tensor, + slot_ids: List[int], + has_prefix: List[bool], + extend_seq_lens_cpu: List[int], + total_padding: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Prefill / extend conv-state step (v1, pure torch). + + Walks each request in the batch, applies ``conv_qk`` with the request's own + initial state (zeros on a fresh first chunk, the cached ``conv_state`` slot + otherwise), writes the updated ``conv_state`` / ``prev_hs_state`` back, and + returns the concatenated ``(qk_out, v2_input)`` in the original token layout. + + ``slot_ids`` is the host mirror of the per-request MambaPool slot indices and + ``has_prefix[i]`` is ``True`` when request ``i`` resumes a cached prefix. + + The Triton swap (:func:`cca_conv1d_fn`) removes this per-request loop. + """ + dtype = hidden_states.dtype + if total_padding is None: + total_padding = conv_state.shape[-1] + in_out_ch = qk.shape[-1] + hidden_size = hidden_states.shape[-1] + + qk_out = torch.empty_like(qk) + v2_input = torch.empty_like(hidden_states) + + # Fresh-prefill fast path: when no request has a cached prefix the per-request + # convs can be coalesced into a single packed convolution. Each request's + # segment is laid out as ``[total_padding zeros, S_i tokens]``. + all_fresh = bool(extend_seq_lens_cpu) and not any(has_prefix) + + if all_fresh: + seq_lens = [int(s) for s in extend_seq_lens_cpu] + pad = total_padding + offsets_in = [0] + for s in seq_lens: + offsets_in.append(offsets_in[-1] + s + pad) + packed = qk.new_zeros((1, in_out_ch, offsets_in[-1])) + start = 0 + for i, s in enumerate(seq_lens): + end = start + s + packed[0, :, offsets_in[i] + pad : offsets_in[i + 1]] = qk[ + start:end + ].transpose(0, 1) + start = end + + packed_out = conv_qk(packed) # [1, C, offsets_in[-1] - pad] + + start = 0 + for i, s in enumerate(seq_lens): + end = start + s + a_i = offsets_in[i] + qk_out[start:end] = packed_out[0, :, a_i : a_i + s].transpose(0, 1) + new_state = packed[0, :, a_i + s : a_i + s + pad] + conv_state[slot_ids[i]] = new_state.to(conv_state.dtype) + + hs_cur = hidden_states[start:end] + first = hidden_states.new_zeros((1, hidden_size)) + v2_input[start:end] = torch.cat([first, hs_cur[:-1]], dim=0) + prev_hs_state[slot_ids[i]] = ( + hs_cur[-1].unsqueeze(-1).to(prev_hs_state.dtype) + ) + start = end + else: + start = 0 + for i, seq_len in enumerate(extend_seq_lens_cpu): + end = start + int(seq_len) + slot = slot_ids[i] + prefix = bool(has_prefix[i]) + + qk_cur = qk[start:end].transpose(0, 1).unsqueeze(0) # [1, C, S_cur] + if prefix: + left_pad = conv_state[slot].unsqueeze(0).to(dtype) + else: + left_pad = qk_cur.new_zeros((1, in_out_ch, total_padding)) + padded = torch.cat([left_pad, qk_cur], dim=-1) + + out = conv_qk(padded) # [1, C, S_cur] + qk_out[start:end] = out.squeeze(0).transpose(0, 1) + + new_state = padded[..., -total_padding:] + conv_state[slot] = new_state.squeeze(0).to(conv_state.dtype) + + hs_cur = hidden_states[start:end] + if prefix: + first = prev_hs_state[slot].squeeze(-1).to(dtype).unsqueeze(0) + else: + first = hidden_states.new_zeros((1, hidden_size)) + v2_input[start:end] = torch.cat([first, hs_cur[:-1]], dim=0) + + prev_hs_state[slot] = hs_cur[-1].unsqueeze(-1).to(prev_hs_state.dtype) + start = end + + return qk_out, v2_input + + +def cca_decode( + qk: torch.Tensor, + hidden_states: torch.Tensor, + conv_qk: nn.Module, + conv_state: torch.Tensor, + prev_hs_state: torch.Tensor, + mamba_indices: torch.Tensor, + total_padding: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Single-token decode conv-state step (v1, pure torch). + + Gathers each request's cached ``conv_state`` / ``prev_hs_state`` via + ``index_select``, runs ``conv_qk`` on the ``[T, C, total_padding + 1]`` + window, and scatters the updated state back with ``index_copy_``. All ops are + on-device (``mamba_indices`` is a device ``long`` tensor), so this stays + CUDA-graph capturable. Returns ``(qk_out, prev_hs)`` where ``prev_hs`` is the + previous hidden state feeding ``val_proj2``. + + The Triton swap is :func:`cca_conv1d_update`. + """ + dtype = hidden_states.dtype + if total_padding is None: + total_padding = conv_state.shape[-1] + + left_pad = conv_state.index_select(0, mamba_indices).to(dtype) + cur = qk.unsqueeze(-1) # [T, C, 1] + padded = torch.cat([left_pad, cur], dim=-1) # [T, C, total_padding + 1] + out = conv_qk(padded) # [T, C, 1] + qk_out = out.squeeze(-1) # [T, C] + + new_state = padded[..., -total_padding:] + conv_state.index_copy_(0, mamba_indices, new_state.to(conv_state.dtype)) + + # Read the previous hidden state (val_proj2 input) BEFORE overwriting the + # slot with the current token. + prev_hs = prev_hs_state.index_select(0, mamba_indices).squeeze(-1).to(dtype) + prev_hs_state.index_copy_( + 0, mamba_indices, hidden_states.unsqueeze(-1).to(prev_hs_state.dtype) + ) + return qk_out, prev_hs + + +# Fused kernel seam (TODO) -- perf swap for the v1 torch paths above. These +# mirror the ``causal_conv1d_fn`` / ``causal_conv1d_update`` contract but for +# CCA's two-stage *grouped* conv (conv_qk[0] depthwise + conv_qk[1] grouped +# per-head), which the stock depthwise ``causal_conv1d`` cannot express. Once +# implemented they replace the per-request loop in ``cca_extend`` and the +# separate gather/conv/scatter launches in ``cca_decode`` with a single +# index-driven kernel. Same ``(qk_out, v2_input)`` return contract. + + +def cca_conv1d_fn(*args, **kwargs): + raise NotImplementedError( + "Fused CCA prefill conv-with-state kernel not implemented yet; " + "the model uses cca_extend (v1 torch) in the meantime." + ) + + +def cca_conv1d_update(*args, **kwargs): + raise NotImplementedError( + "Fused CCA decode conv-with-state kernel not implemented yet; " + "the model uses cca_decode (v1 torch) in the meantime." + ) + + # --------------------------------------------------------------------------- # CCA: Compressed Convolutional Attention QKV projection # --------------------------------------------------------------------------- @@ -401,58 +578,6 @@ class CCA(nn.Module): # ----- helpers --------------------------------------------------------- - @staticmethod - def _get_mamba_indices(forward_batch: ForwardBatch) -> torch.Tensor: - """Per-request MambaPool slot indices as an int64 device tensor. - - The req -> slot mapping depends only on ``forward_batch.req_pool_indices``, - which is constant for every CCA layer within one forward step. Computing - it inside each of the ~60 attention layers would issue one redundant - gather per layer, so it is computed once and memoized on the ForwardBatch - (whose lifetime is exactly one forward step). The lookup is pure on-device - work, so this stays compatible with CUDA graph capture on the decode path. - """ - cached = getattr(forward_batch, _MAMBA_INDICES_ATTR, None) - if cached is None: - cached = ( - get_req_to_token_pool() - .get_mamba_indices(forward_batch.req_pool_indices) - .to(torch.long) - ) - setattr(forward_batch, _MAMBA_INDICES_ATTR, cached) - return cached - - @staticmethod - def _get_mamba_indices_cpu( - forward_batch: ForwardBatch, mamba_indices: torch.Tensor - ) -> list[int]: - """Host mirror of :meth:`_get_mamba_indices`, memoized per forward step. - - Only the extend/prefill path needs the indices on the host to drive its - per-request Python loop; the decode path indexes the pool entirely - on-device. Memoizing turns the previous one-``.tolist()``-sync-per-layer - behavior into a single GPU->CPU sync per forward step. This helper is - never reached on the decode path that CUDA graphs capture. - """ - cached = getattr(forward_batch, _MAMBA_INDICES_CPU_ATTR, None) - if cached is None: - cached = mamba_indices.tolist() - setattr(forward_batch, _MAMBA_INDICES_CPU_ATTR, cached) - return cached - - def _get_pool_state(self, forward_batch: ForwardBatch): - """Retrieve per-request CCA state from the centralized MambaPool. - - ``conv_state`` / ``prev_hs_state`` are layer-local pool views, but the - ``mamba_indices`` req -> slot mapping is shared across layers and so is - memoized on the ForwardBatch (see :meth:`_get_mamba_indices`). - """ - layer_cache = get_req_to_token_pool().mamba2_layer_cache(self.layer_id) - conv_state = layer_cache.conv[0] - prev_hs_state = layer_cache.conv[1] - mamba_indices = self._get_mamba_indices(forward_batch) - return conv_state, prev_hs_state, mamba_indices - def _normalize_qk( self, query: torch.Tensor, key: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor]: @@ -583,195 +708,6 @@ class CCA(nn.Module): value = self._slice_v_per_rank(value_full) return query, key, value - def _forward_extend( - self, - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Prefill / extend path. - - Walks every request in the batch, applies the conv with each request's - own initial state (zero on first chunk, cached otherwise), writes the - updated state and ``prev_hs`` back into the centralized MambaPool, and - returns the concatenated q/k/v in the original token layout. - """ - dtype = hidden_states.dtype - T = hidden_states.shape[0] - - q_raw, _ = self.linear_q(hidden_states) # [T, latent_q] - k_raw, _ = self.linear_k(hidden_states) - qk = torch.cat([q_raw, k_raw], dim=-1) # [T, in_out_ch] - - query_pre = q_raw.view(T, self.num_q_heads, self.head_dim) - key_base = k_raw.view(T, self.num_k_heads, self.head_dim) - - qk_out = torch.empty_like(qk) - v2_input = torch.empty_like(hidden_states) - - conv_state, prev_hs_state, mamba_indices = self._get_pool_state(forward_batch) - # Host view of the slot indices to drive the per-request loop below. - # Memoized on the ForwardBatch, so the GPU->CPU sync runs once per forward - # step rather than once per attention layer (~60 syncs/step otherwise). - mamba_idx_cpu = self._get_mamba_indices_cpu(forward_batch, mamba_indices) - - extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu - extend_prefix_lens_cpu = forward_batch.extend_prefix_lens_cpu - - # Fresh-prefill fast path: when no request has a cached prefix the - # per-request convs (one launch each, ×60 attention layers ×B requests) - # can be coalesced into a single packed convolution. The conv chain is - # two ``kernel_size=2`` convs (effective receptive field = 3), so each - # request's S valid outputs are produced from the packed positions - # ``[a_i, a_i + S_i - 1]`` where the input segment for request i is - # ``[pad, pad, x_0, ..., x_{S-1}]`` of length ``S_i + total_padding``. - all_fresh = bool(extend_seq_lens_cpu) and not any( - int(p) > 0 for p in extend_prefix_lens_cpu - ) - - if all_fresh: - seq_lens = [int(s) for s in extend_seq_lens_cpu] - pad = self.total_padding - # Build packed buffer: per request -> [pad zeros, S_i tokens]. - offsets_in = [0] - for s in seq_lens: - offsets_in.append(offsets_in[-1] + s + pad) - packed = qk.new_zeros((1, self.in_out_ch, offsets_in[-1])) - start = 0 - for i, s in enumerate(seq_lens): - end = start + s - packed[0, :, offsets_in[i] + pad : offsets_in[i + 1]] = qk[ - start:end - ].transpose(0, 1) - start = end - - packed_out = self._conv_qk_run(packed) # [1, C, offsets_in[-1] - pad] - - start = 0 - for i, s in enumerate(seq_lens): - end = start + s - a_i = offsets_in[i] - qk_out[start:end] = packed_out[0, :, a_i : a_i + s].transpose(0, 1) - new_state = packed[0, :, a_i + s : a_i + s + pad] - conv_state[mamba_idx_cpu[i]] = new_state.to(conv_state.dtype) - - hs_cur = hidden_states[start:end] - first = hidden_states.new_zeros((1, self.hidden_size)) - v2_input[start:end] = torch.cat([first, hs_cur[:-1]], dim=0) - prev_hs_state[mamba_idx_cpu[i]] = ( - hs_cur[-1].unsqueeze(-1).to(prev_hs_state.dtype) - ) - start = end - else: - start = 0 - for i, seq_len in enumerate(extend_seq_lens_cpu): - end = start + int(seq_len) - mamba_idx = mamba_idx_cpu[i] - has_prefix = int(extend_prefix_lens_cpu[i]) > 0 - - qk_cur = qk[start:end].transpose(0, 1).unsqueeze(0) # [1, C, S_cur] - if has_prefix: - left_pad = conv_state[mamba_idx].unsqueeze(0).to(dtype) - else: - left_pad = qk_cur.new_zeros((1, self.in_out_ch, self.total_padding)) - padded = torch.cat([left_pad, qk_cur], dim=-1) - - out = self._conv_qk_run(padded) # [1, C, S_cur] - qk_out[start:end] = out.squeeze(0).transpose(0, 1) - - new_state = padded[..., -self.total_padding :] - conv_state[mamba_idx] = new_state.squeeze(0).to(conv_state.dtype) - - hs_cur = hidden_states[start:end] - if has_prefix: - first = prev_hs_state[mamba_idx].squeeze(-1).to(dtype).unsqueeze(0) - else: - first = hidden_states.new_zeros((1, self.hidden_size)) - shifted = torch.cat([first, hs_cur[:-1]], dim=0) - v2_input[start:end] = shifted - - prev_hs_state[mamba_idx] = ( - hs_cur[-1].unsqueeze(-1).to(prev_hs_state.dtype) - ) - - start = end - - query_conv = qk_out[:, : self.latent_q_dim].view( - T, self.num_q_heads, self.head_dim - ) - key_conv = qk_out[:, self.latent_q_dim :].view( - T, self.num_k_heads, self.head_dim - ) - - query, key = self._add_grouped_qk_means( - query_conv, key_conv, query_pre, key_base - ) - query, key = self._normalize_qk(query, key) - - v1, _ = self.val_proj1(hidden_states) - v2, _ = self.val_proj2(v2_input) - value_full = torch.cat([v1, v2], dim=-1).view( - T, self.num_k_heads_full, self.head_dim - ) - value = self._slice_v_per_rank(value_full) - return query, key, value - - def _forward_decode( - self, - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Single-token decode path for a whole batch. - - Reads each request's cached conv state and ``prev_hs`` from the - centralized MambaPool via ``index_select``, runs the conv on the - small ``[T, C, total_padding+1]`` window, and writes back via - ``index_copy_``. - """ - T = hidden_states.shape[0] - dtype = hidden_states.dtype - - conv_state, prev_hs_state, mamba_indices = self._get_pool_state(forward_batch) - - q_raw, _ = self.linear_q(hidden_states) - k_raw, _ = self.linear_k(hidden_states) - qk = torch.cat([q_raw, k_raw], dim=-1) # [T, C] - - query_pre = q_raw.view(T, self.num_q_heads, self.head_dim) - key_base = k_raw.view(T, self.num_k_heads, self.head_dim) - - left_pad = conv_state.index_select(0, mamba_indices).to(dtype) - cur = qk.unsqueeze(-1) # [T, C, 1] - padded = torch.cat([left_pad, cur], dim=-1) # [T, C, total_padding+1] - out = self._conv_qk_run(padded) # [T, C, 1] - qk_out = out.squeeze(-1) # [T, C] - - new_state = padded[..., -self.total_padding :] - conv_state.index_copy_(0, mamba_indices, new_state.to(conv_state.dtype)) - - query_conv = qk_out[:, : self.latent_q_dim].view( - T, self.num_q_heads, self.head_dim - ) - key_conv = qk_out[:, self.latent_q_dim :].view( - T, self.num_k_heads, self.head_dim - ) - - query, key = self._add_grouped_qk_means( - query_conv, key_conv, query_pre, key_base - ) - query, key = self._normalize_qk(query, key) - - prev_hs = prev_hs_state.index_select(0, mamba_indices).squeeze(-1).to(dtype) - v1, _ = self.val_proj1(hidden_states) - v2, _ = self.val_proj2(prev_hs) - value_full = torch.cat([v1, v2], dim=-1).view( - T, self.num_k_heads_full, self.head_dim - ) - value = self._slice_v_per_rank(value_full) - prev_hs_state.index_copy_( - 0, mamba_indices, hidden_states.unsqueeze(-1).to(prev_hs_state.dtype) - ) - return query, key, value - def forward( self, hidden_states: torch.Tensor, @@ -779,6 +715,16 @@ class CCA(nn.Module): ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Project ``hidden_states`` into ``(q, k, v)`` honoring per-request state. + The per-request conv-state plumbing (slot gather/scatter, prefix mask, + cuda-graph buffers) is owned by :class:`ShortConvAttnBackend + `, + reached via ``get_attn_backend().conv_state_metadata``; CCA runs its own + two-stage grouped conv (:func:`cca_extend` / :func:`cca_decode`) against + that handle, so this module holds no pool access. Those functions return + the conv output ``qk_out`` and the ``val_proj2`` input ``v2_input`` (the + shifted / previous hidden state), updating the ``conv_state`` / + ``prev_hs`` pool slots in place. + ``q`` / ``k`` are returned in fp32 (the normalize step keeps fp32 for stability); ``v`` is returned in the input dtype since the caller casts everything back to ``hidden_states.dtype`` before rotary + @@ -798,10 +744,63 @@ class CCA(nn.Module): zero.view(0, self.num_k_heads, self.head_dim), ) + T = hidden_states.shape[0] + q_raw, _ = self.linear_q(hidden_states) # [T, latent_q] + k_raw, _ = self.linear_k(hidden_states) + qk = torch.cat([q_raw, k_raw], dim=-1) # [T, in_out_ch] + + query_pre = q_raw.view(T, self.num_q_heads, self.head_dim) + key_base = k_raw.view(T, self.num_k_heads, self.head_dim) + + # The backend hands out the per-request conv-state handle (slot indices, + # prefix mask, cuda-graph buffers); CCA runs its own two-stage grouped + # conv against it and gets back the conv output + val_proj2 input, with + # the conv_state / prev_hs pool slots updated in place. + meta = get_attn_backend().conv_state_metadata(self.layer_id, forward_batch) + conv_state = meta.layer_cache.conv[0] + prev_hs_state = meta.layer_cache.conv[1] if forward_batch.forward_mode.is_decode_or_idle(): - return self._forward_decode(hidden_states, forward_batch) - # EXTEND / MIXED / DLLM_EXTEND all share the prefill loop. - return self._forward_extend(hidden_states, forward_batch) + qk_out, v2_input = cca_decode( + qk, + hidden_states, + self.conv_qk, + conv_state, + prev_hs_state, + meta.cache_indices, + self.total_padding, + ) + else: + qk_out, v2_input = cca_extend( + qk, + hidden_states, + self.conv_qk, + conv_state, + prev_hs_state, + meta.slot_ids_cpu, + meta.has_prefix_cpu, + forward_batch.extend_seq_lens_cpu, + self.total_padding, + ) + + query_conv = qk_out[:, : self.latent_q_dim].view( + T, self.num_q_heads, self.head_dim + ) + key_conv = qk_out[:, self.latent_q_dim :].view( + T, self.num_k_heads, self.head_dim + ) + + query, key = self._add_grouped_qk_means( + query_conv, key_conv, query_pre, key_base + ) + query, key = self._normalize_qk(query, key) + + v1, _ = self.val_proj1(hidden_states) + v2, _ = self.val_proj2(v2_input) + value_full = torch.cat([v1, v2], dim=-1).view( + T, self.num_k_heads_full, self.head_dim + ) + value = self._slice_v_per_rank(value_full) + return query, key, value # --------------------------------------------------------------------------- diff --git a/test/registered/unit/models/test_zaya_cca.py b/test/registered/unit/models/test_zaya_cca.py index 17243280f..d3740cea2 100644 --- a/test/registered/unit/models/test_zaya_cca.py +++ b/test/registered/unit/models/test_zaya_cca.py @@ -127,19 +127,77 @@ class _MockReqToTokenPool: return req_pool_indices.to(torch.int32) +class _MockShortConvBackend: + """Stand-in for ``ShortConvHybridAttnBackend`` in the CPU unit tests. + + The CCA module reaches the conv-state plumbing via + ``get_attn_backend().conv_state_metadata(...)`` and runs its own conv + kernel. This mock exposes that accessor over a ``_MockReqToTokenPool``, + mirroring ``ShortConvAttnBackend``: the req -> slot mapping (and, for extend, + its host ``.tolist()`` mirror) is resolved once per step and shared across + all conv layers, while the decode path stays entirely on-device. + """ + + def __init__(self, pool: "_MockReqToTokenPool"): + self.req_to_token_pool = pool + self.token_to_kv_pool = None + # Per-forward-step memoization keyed on the ForwardBatch identity, + # mirroring ShortConvAttnBackend.init_forward_metadata. + self._step_indices = {} # id(forward_batch) -> device index tensor + self._step_slot_ids = {} # id(forward_batch) -> host list (extend only) + + def _resolve_indices(self, forward_batch): + key = id(forward_batch) + indices = self._step_indices.get(key) + if indices is None: + indices = self.req_to_token_pool.get_mamba_indices( + forward_batch.req_pool_indices + ).to(torch.long) + self._step_indices[key] = indices + return indices + + def _resolve_slot_ids(self, forward_batch, indices): + key = id(forward_batch) + slot_ids = self._step_slot_ids.get(key) + if slot_ids is None: + slot_ids = indices.tolist() + self._step_slot_ids[key] = slot_ids + return slot_ids + + def conv_state_metadata(self, layer_id, forward_batch): + from sglang.srt.layers.attention.linear.short_conv_backend import ( + ShortConvMetadata, + ) + + layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) + indices = self._resolve_indices(forward_batch) # already int64 + if forward_batch.forward_mode.is_decode_or_idle(): + return ShortConvMetadata(layer_cache=layer_cache, cache_indices=indices) + + slot_ids = self._resolve_slot_ids(forward_batch, indices) + has_prefix = [int(p) > 0 for p in forward_batch.extend_prefix_lens_cpu] + return ShortConvMetadata( + layer_cache=layer_cache, + cache_indices=indices, + slot_ids_cpu=slot_ids, + has_prefix_cpu=has_prefix, + ) + + @contextmanager def _mock_pool_context(pool: _MockReqToTokenPool): - """Install a mock ``ForwardContext`` whose ``req_to_token_pool`` is ``pool``.""" + """Install a mock ``ForwardContext`` whose ``attn_backend`` exposes both + ``req_to_token_pool`` and ``conv_state_metadata`` over ``pool``.""" from sglang.srt.model_executor.forward_context import ( ForwardContext, set_forward_context, ) - backend = SimpleNamespace(req_to_token_pool=pool, token_to_kv_pool=None) + backend = _MockShortConvBackend(pool) ctx = ForwardContext(attn_backend=backend) prev = set_forward_context(ctx) try: - yield pool + yield backend finally: set_forward_context(prev) @@ -457,17 +515,17 @@ class TestZayaCCA(CustomTestCase): ) pool = _CountingPool(pool_size=8, cca_config=config) - with _mock_pool_context(pool): + with _mock_pool_context(pool) as backend: fb = _fresh_fb() cca0.forward(hs, fb) cca2.forward(hs, fb) # Two CCA layers, one forward step -> one shared lookup, both the - # device tensor and its host mirror memoized on the ForwardBatch. + # device tensor and its host mirror memoized once per step on the + # backend (ShortConvAttnBackend does this in init_forward_metadata). self.assertEqual(pool.get_mamba_indices_calls, 1) - self.assertTrue(hasattr(fb, "_zaya_mamba_indices")) - self.assertTrue(hasattr(fb, "_zaya_mamba_indices_cpu")) - self.assertEqual(fb._zaya_mamba_indices_cpu, [0]) + self.assertIn(id(fb), backend._step_indices) + self.assertEqual(backend._step_slot_ids[id(fb)], [0]) # A new forward step (fresh ForwardBatch) resolves the mapping again. cca0.forward(hs, _fresh_fb()) @@ -479,16 +537,20 @@ class TestZayaCCA(CustomTestCase): cca, config = _make_tiny_cca(seed=7) pool = _MockReqToTokenPool(pool_size=8, cca_config=config) - with _mock_pool_context(pool): + with _mock_pool_context(pool) as backend: + # Keep a reference to the extend batch so its id() cannot be recycled + # by the later decode batch (the mock keys its per-step memo on + # id(forward_batch); a GC'd-then-reused address would false-collide). + fb_extend = _make_forward_batch( + is_decode=False, + extend_seq_lens_cpu=[3], + extend_prefix_lens_cpu=[0], + req_pool_indices=[0], + input_ids=torch.arange(3, dtype=torch.int64), + ) cca.forward( torch.randn(3, config.hidden_size, dtype=torch.float32) * 0.1, - _make_forward_batch( - is_decode=False, - extend_seq_lens_cpu=[3], - extend_prefix_lens_cpu=[0], - req_pool_indices=[0], - input_ids=torch.arange(3, dtype=torch.int64), - ), + fb_extend, ) fb_decode = _make_forward_batch( is_decode=True, @@ -502,10 +564,11 @@ class TestZayaCCA(CustomTestCase): fb_decode, ) - # Device indices are memoized, but the host ``.tolist()`` mirror is only - # built by the extend path. - self.assertTrue(hasattr(fb_decode, "_zaya_mamba_indices")) - self.assertFalse(hasattr(fb_decode, "_zaya_mamba_indices_cpu")) + # Decode resolves device indices, but the host ``.tolist()`` mirror + # is only built by the extend path -- so the decode step stays + # entirely on-device (CUDA-graph friendly). + self.assertIn(id(fb_decode), backend._step_indices) + self.assertNotIn(id(fb_decode), backend._step_slot_ids) class TestZayaCCATensorParallel(CustomTestCase):