feat(inkling): migrate short convs onto the ShortConv attention backend (#33023)

This commit is contained in:
Cheng Wan
2026-07-31 11:52:12 -07:00
committed by GitHub
parent d3222bcc3a
commit 77c77a3da8
16 changed files with 1195 additions and 435 deletions
@@ -232,6 +232,7 @@ class AscendHybridLinearAttnBackend(HybridLinearAttnBackend):
mamba_track_indices: Optional[torch.Tensor], mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor], mamba_steps_to_track: Optional[torch.Tensor],
model, model,
req_pool_indices: Optional[torch.Tensor] = None,
): ):
""" """
Update mamba states after MTP verify using fully fused Triton kernel. Update mamba states after MTP verify using fully fused Triton kernel.
@@ -242,6 +243,7 @@ class AscendHybridLinearAttnBackend(HybridLinearAttnBackend):
- index_select kernel launches - index_select kernel launches
- nonzero kernel launches - nonzero kernel launches
""" """
del req_pool_indices # accepted for hook parity; slots come from metadata
request_number = last_correct_step_indices.shape[0] request_number = last_correct_step_indices.shape[0]
state_indices_tensor = ( state_indices_tensor = (
@@ -277,6 +277,22 @@ def create_dual_chunk_flash_attn_backend(runner):
return DualChunkFlashAttentionBackend(runner) return DualChunkFlashAttentionBackend(runner)
def attn_backend_wrapper_for_draft_extend(
runner: "ModelRunner", full_attn_backend: "AttentionBackend"
):
"""Apply the model's attention wrapper to a DRAFT-EXTEND backend, if it needs one.
``DraftBackendFactory`` skips :func:`attn_backend_wrapper`, which is right for
the mamba hybrids whose MTP draft is all softmax attention. Inkling's draft has
its own short convs, so it must expose ``conv_state_metadata`` too.
"""
from sglang.srt.configs.inkling import InklingMMConfig, InklingModelConfig
if isinstance(runner.model_config.hf_config, (InklingModelConfig, InklingMMConfig)):
return attn_backend_wrapper(runner, full_attn_backend)
return full_attn_backend
def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBackend"): def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBackend"):
""" """
Wrapper for special models like hybrid GDN, so we don't Wrapper for special models like hybrid GDN, so we don't
@@ -305,7 +321,16 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
if isinstance( if isinstance(
runner.model_config.hf_config, (InklingModelConfig, InklingMMConfig) runner.model_config.hf_config, (InklingModelConfig, InklingMMConfig)
): ):
return full_attn_backend from sglang.srt.layers.attention.linear.inkling_sconv_backend import (
InklingShortConvAttnBackend,
InklingShortConvHybridAttnBackend,
)
return InklingShortConvHybridAttnBackend(
full_attn_backend,
InklingShortConvAttnBackend(runner),
cfg.full_attention_layer_ids,
)
from sglang.kernels.ops.attention.fla.utils import check_environments from sglang.kernels.ops.attention.fla.utils import check_environments
from sglang.srt.layers.attention.linear.kda_backend import KDAAttnBackend from sglang.srt.layers.attention.linear.kda_backend import KDAAttnBackend
@@ -1108,8 +1108,15 @@ class HybridLinearAttnBackend(AttentionBackend):
mamba_track_indices: Optional[torch.Tensor], mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor], mamba_steps_to_track: Optional[torch.Tensor],
model, model,
req_pool_indices: Optional[torch.Tensor] = None,
): ):
"""Update mamba states after MTP verify via a fused gather-scatter kernel.""" """Update mamba states after MTP verify via a fused gather-scatter kernel.
``req_pool_indices`` serves implementations that must re-derive the state
slot ids instead of reusing this step's ``forward_metadata``; the scatter
below reads the metadata it just planned.
"""
del req_pool_indices
request_number = last_correct_step_indices.shape[0] request_number = last_correct_step_indices.shape[0]
state_indices_tensor = ( state_indices_tensor = (
@@ -0,0 +1,572 @@
# 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.
# ==============================================================================
"""Inkling's short-conv state backend.
A :mod:`~sglang.srt.layers.attention.linear.short_conv_backend` sidecar. Four short
convs per decoder layer keep per-request conv state in the centralized
``MambaPool``; the model reaches this via :meth:`conv_state_metadata`, never
through ``forward_decode`` / ``forward_extend``.
On top of what :class:`ShortConvAttnBackend` owns, Inkling's kernels take a
precomputed ``cache_mask`` / ``safe_idx`` / ``cu`` / ``si`` set plus the extend
``track_conv_indices``. All of it is step-global, so it is resolved once per step
and shared by every conv module in the step (a decoder layer holds four).
The hook split is a decode-latency decision. ``init_forward_metadata_in_graph`` is
*recorded* into the decode / target-verify / draft-extend graphs, so prep placed
there replays for free; out of graph it lands on the per-step CPU path that a
captured graph exists to avoid. ``init_forward_metadata_out_graph`` therefore takes
only what a phase cannot record: full-cuda-graph prefill (no in-graph hook) and the
unified pool's slot translate. Prep consequently sits outside the graph, so every
tensor a captured kernel reads lives in a graph-static buffer refilled in place.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, NamedTuple, Optional
import torch
from sglang.kernels.ops.mamba.mamba_state_scatter_triton import (
scatter_mamba_states_after_mtp_verify,
)
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
ShortConvHybridAttnBackend,
)
from sglang.srt.layers.attention.linear.short_conv_backend import ShortConvAttnBackend
from sglang.srt.layers.attention.mamba.mamba2_metadata import ForwardMetadata
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.inkling_common.kernels.sconv import (
HIS_ONES,
HIS_PREFIX,
HIS_SEQ_MINUS_EXT,
HIS_ZEROS,
SconvDecodeMetadata,
SconvExtendMetadata,
SconvMetadataOut,
fused_decode_sconv_metadata,
fused_extend_sconv_metadata,
precompute_helion_extend_metadata,
)
from sglang.srt.runtime_context import get_server_args
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
class InklingShortConvMetadata(NamedTuple):
"""Per-(layer, step) conv-state handle handed to Inkling's conv kernels.
``layer_cache`` holds this layer's pool views indexed by ``SconvType``; the
rest is step-global, and on the graph path is a static buffer refilled in place.
"""
layer_cache: Any
cache_indices: torch.Tensor # per-request slot ids, int32
query_start_loc: Optional[torch.Tensor] = None # cu-seqlens, int32
has_initial_state: Optional[torch.Tensor] = None # "resumes a cached prefix"
precomputed: Optional[SconvExtendMetadata | SconvDecodeMetadata] = None
# [B, conv_kernel - 1] input positions whose conv window feeds the prefix
# cache. Extend only, and only when tracking is on.
track_conv_indices: Optional[torch.Tensor] = None
class InklingShortConvAttnBackend(ShortConvAttnBackend):
"""Owns Inkling's per-step short-conv state plumbing (see module docstring)."""
# int32 matches the pool and the conv kernels; an int64 view would re-run a
# narrowing cast in every conv layer.
cache_indices_dtype: torch.dtype = torch.int32
# Fully device-side extend path, so the ZAYA1-style host mirrors would only
# add a device->host sync per step.
needs_extend_host_mirrors: bool = False
def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner)
# conv[i] is [n_layers, n_slots, conv_kernel - 1, conv_dim].
self.conv_state_len: int = self.conv_states_shape[2]
self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
# A plain table lookup is recordable; the unified pool's translate is an
# allocator lookup and must stay in the out-of-graph replay prep.
self._slot_gather_recordable = (
type(self.req_to_token_pool).translate_mamba_indices
is HybridReqToTokenPool.translate_mamba_indices
)
self._query_start_loc: Optional[torch.Tensor] = None
self._precomputed: Optional[SconvExtendMetadata | SconvDecodeMetadata] = None
self._track_conv_indices: Optional[torch.Tensor] = None
self._alloc_graph_buffers()
def _alloc_graph_buffers(self):
"""Sized from the CONFIGURED capture shapes, once, never reallocated:
growing a buffer after a graph captured it moves the address that graph
reads, and prefill captures before the decode runner reports its bounds."""
server_args = get_server_args()
cuda_graph_config = server_args.cuda_graph_config
decode_bs: list[int] = []
prefill_tokens: list[int] = []
decode_max_bs = 0
if cuda_graph_config is not None:
decode_bs = list(cuda_graph_config.decode.bs or [])
prefill_tokens = list(cuda_graph_config.prefill.bs or [])
decode_max_bs = cuda_graph_config.decode.max_bs or 0
draft_token_num = server_args.speculative_num_draft_tokens or 1
# req_to_token_pool.size is the runner's max_bs for both graph phases.
max_bs = max([self.req_to_token_pool.size, decode_max_bs, *decode_bs])
max_tokens = max([max_bs, *prefill_tokens, max_bs * draft_token_num])
dev = self.device
self._graph_bufs = SconvMetadataOut(
query_start_loc=torch.empty(max_bs + 1, dtype=torch.int32, device=dev),
has_initial_state=torch.empty(max_bs, dtype=torch.bool, device=dev),
cache_mask=torch.empty((max_bs, 1, 1), dtype=torch.bool, device=dev),
safe_idx=torch.empty(max_bs, dtype=torch.int64, device=dev),
cu=torch.empty(max_bs + 1, dtype=torch.int64, device=dev),
si=torch.empty(max_tokens, dtype=torch.int32, device=dev),
)
self._graph_track_conv_indices = torch.zeros(
(max_bs, self.conv_state_len), dtype=torch.int64, device=dev
)
# Same address-stability requirement; the base only sizes this from
# init_cuda_graph_state, which the prefill graph never calls.
self._alloc_cache_indices_buf(max_bs)
self._track_window_offsets = torch.arange(
self.conv_state_len, dtype=torch.int64, device=dev
)
self._track_index_floor = torch.zeros((1,), dtype=torch.int64, device=dev)
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
super().init_cuda_graph_state(max_bs, max_num_tokens)
# Fail now, not at the first replay, if a phase outgrew __init__'s bounds.
self._graph_metadata_out(B=max_bs, T=max_num_tokens)
def _graph_metadata_out(self, *, B: int, T: int) -> SconvMetadataOut:
"""Graph-static destinations sliced to this step. Asserts rather than
allocating, which would leave captured kernels on a dead address."""
bufs = self._graph_bufs
assert B + 1 <= bufs["query_start_loc"].shape[0] and T <= bufs["si"].shape[0], (
f"short-conv metadata buffers too small for a captured shape: "
f"B={B}, T={T} vs bs bound {bufs['query_start_loc'].shape[0] - 1}, "
f"token bound {bufs['si'].shape[0]}"
)
return SconvMetadataOut(
query_start_loc=bufs["query_start_loc"][: B + 1],
has_initial_state=bufs["has_initial_state"][:B],
cache_mask=bufs["cache_mask"][:B],
safe_idx=bufs["safe_idx"][:B],
cu=bufs["cu"][: B + 1],
si=bufs["si"][:T],
)
def _forward_metadata(self, forward_batch: ForwardBatch) -> ForwardMetadata:
"""Slot ids only. Leaner than ``MambaAttnBackendBase._forward_metadata``
on purpose: no SSM state (whose track prep also syncs), a conv window on a
different axis, and ``query_start_loc`` from the fused kernel."""
return ForwardMetadata(
query_start_loc=None,
mamba_cache_indices=self._translate_mamba_indices(
self.req_to_token_pool.get_mamba_indices(forward_batch.req_pool_indices)
),
)
def _reset_step_state(self):
super()._reset_step_state()
self._query_start_loc = None
self._precomputed = None
self._track_conv_indices = None
@staticmethod
def _phase_records_metadata(forward_batch: ForwardBatch) -> bool:
"""True when this phase's runner records ``init_forward_metadata_in_graph``
(decode / target-verify / draft-extend do; full-cuda-graph prefill does
not)."""
mode = forward_batch.forward_mode
return (
mode.is_decode_or_idle()
or mode.is_target_verify()
or mode.is_draft_extend_v2()
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Eager path: nothing downstream is captured, so kernels may allocate."""
self._prepare_slot_indices(forward_batch)
self._refresh_sconv_metadata(forward_batch, on_graph_path=False)
def init_forward_metadata_out_graph(
self, forward_batch: ForwardBatch, in_capture: bool = False
):
"""Whatever this phase cannot record. Runs before EVERY replay, so the
common path is one predicate and a return."""
del in_capture
if self._phase_records_metadata(forward_batch):
if not self._slot_gather_recordable:
self._prepare_slot_indices(forward_batch)
return
self._prepare_slot_indices(forward_batch)
self._refresh_sconv_metadata(forward_batch, on_graph_path=True)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch):
"""Recorded into the graph: writes the static destinations, so the launches
refill them every replay and never allocate (the hook's contract)."""
if not self._phase_records_metadata(forward_batch):
return
if self._slot_gather_recordable:
self._prepare_slot_indices(forward_batch)
self._refresh_sconv_metadata(forward_batch, on_graph_path=True)
def init_forward_metadata_capture_cpu_graph(self, *args, **kwargs):
raise NotImplementedError(
"Inkling's short-conv backend has no CPU-graph path; its conv "
"kernels are CUDA/Triton only."
)
def _prepare_slot_indices(self, forward_batch: ForwardBatch):
self._reset_step_state()
req_pool_indices = forward_batch.req_pool_indices
n = req_pool_indices.shape[0]
buf = self._cache_indices_buf
if self._slot_gather_recordable and n <= buf.shape[0]:
# One launch; the base's gather-then-copy would add a second recorded
# kernel per step (the pool's table already has this dtype). No PAD
# sentinel needed: MambaSlotAllocator.clear reserves slot 0 as the dummy
# write target, so zero-filled padded rows already land there.
torch.index_select(
self.req_to_token_pool.req_index_to_mamba_index_mapping,
0,
req_pool_indices,
out=buf[:n],
)
self._cache_indices = buf[:n]
self.forward_metadata = ForwardMetadata(
query_start_loc=None, mamba_cache_indices=self._cache_indices
)
return
self.forward_metadata = self._forward_metadata(forward_batch)
self._refresh_cache_indices()
def _refresh_sconv_metadata(
self, forward_batch: ForwardBatch, *, on_graph_path: bool
):
if self._cache_indices is None:
return
mode = forward_batch.forward_mode
if mode.is_decode_or_idle():
self._refresh_decode_metadata(forward_batch, on_graph_path)
elif mode.is_target_verify():
self._refresh_extend_metadata(forward_batch, on_graph_path)
elif mode.is_extend(include_draft_extend_v2=True):
self._refresh_extend_metadata(forward_batch, on_graph_path)
self._refresh_track_conv_indices(forward_batch, on_graph_path)
else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}")
def _refresh_decode_metadata(
self, forward_batch: ForwardBatch, on_graph_path: bool
):
B = forward_batch.batch_size
(
self._query_start_loc,
self._has_initial_state,
self._precomputed,
) = fused_decode_sconv_metadata(
B=B,
cache_indices=self._cache_indices,
out=self._graph_metadata_out(B=B, T=B) if on_graph_path else None,
)
def _refresh_extend_metadata(
self, forward_batch: ForwardBatch, on_graph_path: bool
):
"""One fused launch; unfused fallback off-CUDA / past the batch bound."""
B = forward_batch.batch_size
if forward_batch.forward_mode.is_target_verify():
# target_verify has no extend_seq_lens/extend_prefix_lens; the lens are
# a constant draft_token_num per request.
draft_token_num = forward_batch.spec_info.draft_token_num
T = B * draft_token_num
his_kwargs = dict(his_mode=HIS_ONES, draft_token_num=draft_token_num)
else:
T = forward_batch.extend_num_tokens
spec_info = forward_batch.spec_info
if (
isinstance(spec_info, EagleDraftExtendInput)
and spec_info.num_front_tokens > 0
):
# Boundary-KV fix: run conv fresh so warm-up rows rebuild the window.
his_mode, his_src = HIS_ZEROS, None
elif forward_batch.extend_prefix_lens is not None:
his_mode, his_src = HIS_PREFIX, forward_batch.extend_prefix_lens
else:
# draft_extend_v2 capture has no extend_prefix_lens.
his_mode, his_src = HIS_SEQ_MINUS_EXT, forward_batch.seq_lens
his_kwargs = dict(
his_mode=his_mode,
extend_seq_lens=forward_batch.extend_seq_lens,
his_src=his_src,
)
# Captured kernels bake their token extent at CAPTURE (the prefill bucket)
# while replay reports only the live count, so fill the WHOLE seq-index
# buffer; the kernel clamps the tail to B - 1. target_verify's
# B * draft_token_num is exact either way.
fill_T = T
out = None
if on_graph_path:
if not forward_batch.forward_mode.is_target_verify():
fill_T = self._graph_bufs["si"].shape[0]
out = self._graph_metadata_out(B=B, T=fill_T)
fused = fused_extend_sconv_metadata(
B=B,
T=fill_T,
cache_indices=self._cache_indices,
out=out,
**his_kwargs,
)
if fused is not None:
query_start_loc, has_initial_state, precomputed = fused
else:
# The unfused fallback allocates, so it cannot serve a captured shape.
assert not on_graph_path, (
"the fused extend metadata kernel declined a captured shape "
f"(B={B}); its unfused fallback is not cuda-graph safe"
)
query_start_loc, has_initial_state = self._unfused_extend_metadata(
forward_batch
)
precomputed = precompute_helion_extend_metadata(
B=B,
T=T,
W=self.conv_state_len + 1,
cache_indices=self._cache_indices,
has_initial_state=has_initial_state,
query_start_loc=query_start_loc,
)
if fill_T != T:
# Hand back the live extent; only the address matters to the graph.
precomputed = SconvExtendMetadata(
cache_mask=precomputed["cache_mask"],
safe_idx=precomputed["safe_idx"],
cu=precomputed["cu"],
si=precomputed["si"][:T],
)
self._query_start_loc = query_start_loc
self._has_initial_state = has_initial_state
self._precomputed = precomputed
def _unfused_extend_metadata(self, forward_batch: ForwardBatch):
"""Unfused query_start_loc / has_initial_state prep; fallback only."""
device = forward_batch.req_pool_indices.device
if forward_batch.forward_mode.is_target_verify():
draft_token_num = forward_batch.spec_info.draft_token_num
query_start_loc = torch.arange(
0,
(forward_batch.batch_size + 1) * draft_token_num,
draft_token_num,
dtype=torch.int32,
device=device,
)
has_initial_state = torch.ones(
forward_batch.batch_size, dtype=torch.bool, device=device
)
return query_start_loc, has_initial_state
query_start_loc = torch.zeros(
forward_batch.batch_size + 1,
dtype=torch.int32,
device=device,
)
query_start_loc[1:] = forward_batch.extend_seq_lens.cumsum(dim=0)
spec_info = forward_batch.spec_info
if (
isinstance(spec_info, EagleDraftExtendInput)
and spec_info.num_front_tokens > 0
):
has_initial_state = torch.zeros(
forward_batch.batch_size, dtype=torch.bool, device=device
)
elif forward_batch.extend_prefix_lens is not None:
has_initial_state = forward_batch.extend_prefix_lens > 0
else:
has_initial_state = (
forward_batch.seq_lens[: forward_batch.batch_size]
- forward_batch.extend_seq_lens
) > 0
return query_start_loc, has_initial_state
def _refresh_track_conv_indices(
self, forward_batch: ForwardBatch, on_graph_path: bool
):
"""Input positions of the conv windows to snapshot for prefix caching: the
last ``conv_kernel - 1`` tokens up to the last complete
``mamba_cache_chunk_size`` boundary.
The padded tail is ZEROED, not left stale: the captured gather reads all
``batch_size`` rows while the track lengths cover only live requests, and
every row it may read must index inside *this* replay's token buffer.
"""
if forward_batch.mamba_track_mask is None:
return
rows = forward_batch.batch_size
query_start_loc = self._query_start_loc
live = min(
rows,
forward_batch.mamba_track_seqlens.shape[0],
forward_batch.extend_prefix_lens.shape[0],
)
lens_to_track = (
forward_batch.mamba_track_seqlens[:live]
- forward_batch.extend_prefix_lens[:live]
)
chunk_aligned = (
lens_to_track // self.mamba_cache_chunk_size
) * self.mamba_cache_chunk_size
start_indices = query_start_loc[:live] + chunk_aligned - self.conv_state_len
if on_graph_path:
assert rows <= self._graph_track_conv_indices.shape[0], (
f"track-index buffer too small for a captured shape: rows={rows} "
f"vs bound {self._graph_track_conv_indices.shape[0]}"
)
out = self._graph_track_conv_indices[:rows]
else:
out = torch.empty(
(rows, self.conv_state_len),
dtype=torch.int64,
device=start_indices.device,
)
torch.add(
start_indices.unsqueeze(-1).to(torch.int64),
self._track_window_offsets,
out=out[:live],
)
# 1-element tensors, never [-1]: a 0-d -> Python conversion would sync.
torch.clamp(
out[:live],
min=self._track_index_floor,
max=query_start_loc[-1:].to(torch.int64) - 1,
out=out[:live],
)
if live < rows:
out[live:].zero_()
self._track_conv_indices = out
def commit_conv_state_after_mtp_verify(
self,
*,
req_pool_indices: torch.Tensor,
last_correct_step_indices: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor],
) -> None:
"""Commit the TARGET_VERIFY conv windows at each request's last accepted step.
Slot ids come from ``req_pool_indices``, not the per-step
``self._cache_indices``: this runs after the forward context exits, so that
buffer may already belong to a later forward.
"""
pool = self.req_to_token_pool
scatter_mamba_states_after_mtp_verify(
pool.get_speculative_mamba2_params_all_layers(),
self._translate_mamba_indices(pool.get_mamba_indices(req_pool_indices)),
last_correct_step_indices,
mamba_track_indices,
mamba_steps_to_track,
)
def conv_state_metadata(
self, layer_id: int, forward_batch: ForwardBatch
) -> InklingShortConvMetadata:
"""``layer_id``'s handle for this step: a pure read, so every conv layer
shares one gather, one fused launch and one track-index build."""
del forward_batch
return InklingShortConvMetadata(
layer_cache=self.req_to_token_pool.mamba2_layer_cache(layer_id),
cache_indices=self._cache_indices,
query_start_loc=self._query_start_loc,
has_initial_state=self._has_initial_state,
precomputed=self._precomputed,
track_conv_indices=self._track_conv_indices,
)
class InklingShortConvHybridAttnBackend(ShortConvHybridAttnBackend):
"""Full-attention backend plus Inkling's conv-state sidecar.
Inkling has NO linear-attention layers, so every layer routes to the
full-attention child and the sidecar is reached only via
:meth:`conv_state_metadata`. Four departures from
:class:`ShortConvHybridAttnBackend`: every layer is full attention (including
the draft's, so the base's ``full_attn_layers = [0]`` does not hold);
DRAFT_EXTEND_V2 still inits the sidecar (the draft runs its own convs, unlike
the mamba models the base's skip was written for); the full-attention backend's
capability surface stays visible through the wrapper; and the MTP-verify commit
is Inkling's own, not the generic mamba scatter.
"""
def _is_full_attn(self, layer=None, layer_id: Optional[int] = None) -> bool:
del layer, layer_id
return True
def update_mamba_state_after_mtp_verify(
self,
last_correct_step_indices: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor],
model=None,
req_pool_indices: Optional[torch.Tensor] = None,
):
"""Overrides the generic mamba scatter, which sources slot ids from
``forward_metadata`` -- stale once the forward context has exited."""
del model
assert req_pool_indices is not None, (
"Inkling's conv-state commit needs req_pool_indices; the caller must "
"pass the verify batch's request slots."
)
self.short_conv_backend.commit_conv_state_after_mtp_verify(
req_pool_indices=req_pool_indices,
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
for attn_backend in self.attn_backend_list:
attn_backend.init_forward_metadata(forward_batch)
@property
def forward_metadata(self):
# The sidecar's is reached via conv_state_metadata, so this is the attention
# one (KV write locs, the SWA loc translate).
return self.full_attn_backend.forward_metadata
@property
def supports_ragged_verify_graph(self) -> bool:
return self.full_attn_backend.supports_ragged_verify_graph
@property
def supports_full_cuda_graph_chunked_prefix(self) -> bool:
return self.full_attn_backend.supports_full_cuda_graph_chunked_prefix
def prepare_full_cuda_graph_chunked_prefix(self, *args, **kwargs):
return self.full_attn_backend.prepare_full_cuda_graph_chunked_prefix(
*args, **kwargs
)
def draft_extend_metadata_captured_in_graph(self) -> bool:
return self.full_attn_backend.draft_extend_metadata_captured_in_graph()
@@ -85,6 +85,13 @@ class ShortConvAttnBackend(MambaAttnBackendBase):
# always populated for extend regardless of this flag.) # always populated for extend regardless of this flag.)
needs_cpu_seq_lens: bool = False needs_cpu_seq_lens: bool = False
# int64 is canonical (the CUDA causal_conv1d narrows at its own boundary); a
# subclass whose kernels take int32 sets int32 to skip that per-layer cast.
cache_indices_dtype: torch.dtype = torch.int64
# The host mirrors below cost a device->host sync per extend step, so only
# models with a host extend loop (ZAYA1 v1) ask for them.
needs_extend_host_mirrors: bool = True
def __init__(self, model_runner: ModelRunner): def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner) super().__init__(model_runner)
mamba_cache = self.req_to_token_pool.mamba_pool.mamba_cache mamba_cache = self.req_to_token_pool.mamba_pool.mamba_cache
@@ -107,19 +114,24 @@ class ShortConvAttnBackend(MambaAttnBackendBase):
self._has_prefix_cpu = None self._has_prefix_cpu = None
def _alloc_cache_indices_buf(self, max_bs: int): def _alloc_cache_indices_buf(self, max_bs: int):
# Persistent int64 index buffer, refilled in place per step so the # Refilled in place per step so a captured graph reads a stable address.
# captured (cuda or cpu) graph reads a stable address. # Grow-only, never reallocated at the same size: the cuda- and cpu-graph
# hooks can both run, in either order, after another phase captured.
buf = self._cache_indices_buf
if buf is not None and buf.shape[0] >= max_bs:
return
assert buf is None, (
f"cache-indices buffer must be sized before any graph capture: have "
f"{buf.shape[0]}, need {max_bs}"
)
self._cache_indices_buf = torch.empty( self._cache_indices_buf = torch.empty(
max_bs, dtype=torch.int64, device=self.device max_bs, dtype=self.cache_indices_dtype, device=self.device
) )
def _refresh_cache_indices(self): def _refresh_cache_indices(self):
# Resolve the int64 slot-index view ONCE per step, shared by every conv # ONCE per step, shared by every conv layer. With a graph buffer, refill IN
# layer. When a graph index buffer is allocated and large enough, refill # PLACE and hand out a view so the captured address stays current; otherwise
# it IN PLACE and hand out a view -- the captured graph then reads a # (eager, or bs past the buffer) a fresh cast is fine.
# 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 md = self.forward_metadata
idx = md.mamba_cache_indices if md is not None else None idx = md.mamba_cache_indices if md is not None else None
buf = self._cache_indices_buf buf = self._cache_indices_buf
@@ -130,7 +142,7 @@ class ShortConvAttnBackend(MambaAttnBackendBase):
buf[:n].copy_(idx) buf[:n].copy_(idx)
self._cache_indices = buf[:n] self._cache_indices = buf[:n]
else: else:
self._cache_indices = idx.to(torch.long) self._cache_indices = idx.to(self.cache_indices_dtype)
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
super().init_cuda_graph_state(max_bs, max_num_tokens) super().init_cuda_graph_state(max_bs, max_num_tokens)
@@ -153,7 +165,7 @@ class ShortConvAttnBackend(MambaAttnBackendBase):
and not mode.is_draft_extend_v2() and not mode.is_draft_extend_v2()
): ):
self._has_initial_state = forward_batch.extend_prefix_lens > 0 self._has_initial_state = forward_batch.extend_prefix_lens > 0
if self._cache_indices is not None: if self.needs_extend_host_mirrors and self._cache_indices is not None:
self._slot_ids_cpu = self._cache_indices.tolist() self._slot_ids_cpu = self._cache_indices.tolist()
self._has_prefix_cpu = [ self._has_prefix_cpu = [
int(p) > 0 for p in forward_batch.extend_prefix_lens_cpu int(p) > 0 for p in forward_batch.extend_prefix_lens_cpu
-32
View File
@@ -1198,38 +1198,6 @@ class InklingForConditionalGeneration(nn.Module):
), ),
) )
def update_conv_state_after_mtp_verify(
self,
req_to_token_pool,
req_pool_indices: torch.Tensor,
last_correct_step_indices: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor],
) -> None:
"""Commit the per-step sconv windows saved during TARGET_VERIFY into the
persistent conv caches at each request's last accepted step.
Inkling bypasses the HybridLinearAttnBackend wrapper (ShortConvolution reads
the mamba pool directly), so the model owns this commit instead of an
attention-backend hook. The pool is passed in because this runs from the
spec worker after the forward context has exited.
"""
from sglang.kernels.ops.mamba.mamba_state_scatter_triton import (
scatter_mamba_states_after_mtp_verify,
)
pool = req_to_token_pool
mamba_indices = pool.translate_mamba_indices(
pool.get_mamba_indices(req_pool_indices)
)
scatter_mamba_states_after_mtp_verify(
pool.get_speculative_mamba2_params_all_layers(),
mamba_indices,
last_correct_step_indices,
mamba_track_indices,
mamba_steps_to_track,
)
def _load_regular_param( def _load_regular_param(
self, self,
params_dict: dict[str, torch.nn.Parameter], params_dict: dict[str, torch.nn.Parameter],
@@ -1003,7 +1003,7 @@ def ar_scattered_sconv_fused(
**norm_kwargs, **norm_kwargs,
) )
if is_verify: if is_verify:
# Save the per-position windows for update_conv_state_after_mtp_verify. # Save the per-position windows for the backend's MTP-verify commit.
sconv.verify_fused_ar_finish(forward_batch, x_scratch, cache_indices) sconv.verify_fused_ar_finish(forward_batch, x_scratch, cache_indices)
if norm is not None: if norm is not None:
return norm_kwargs["norm_out"], norm_residual return norm_kwargs["norm_out"], norm_residual
@@ -23,6 +23,49 @@ class SconvExtendMetadata(TypedDict):
si: torch.Tensor si: torch.Tensor
class SconvMetadataOut(TypedDict):
"""Preallocated destinations for the fused metadata kernels.
A caller that needs the addresses to stay stable across cuda-graph replays
passes its static buffers, already sliced to this step's B / T, so the kernel
writes straight into them instead of allocating.
"""
query_start_loc: torch.Tensor # [B + 1] int32
has_initial_state: torch.Tensor # [B] bool
cache_mask: torch.Tensor # [B, 1, 1] bool
safe_idx: torch.Tensor # [B] int64
cu: torch.Tensor # [B + 1] int64
si: torch.Tensor # [T] int32
def _metadata_out(
out: "SconvMetadataOut | None", *, B: int, T: int, device: torch.device
) -> SconvMetadataOut:
"""Metadata destinations: freshly allocated, or ``out`` shape-checked."""
spec = (
("query_start_loc", (B + 1,), torch.int32),
("has_initial_state", (B,), torch.bool),
("cache_mask", (B, 1, 1), torch.bool),
("safe_idx", (B,), torch.int64),
("cu", (B + 1,), torch.int64),
("si", (T,), torch.int32),
)
if out is None:
return SconvMetadataOut(
**{
name: torch.empty(shape, dtype=dtype, device=device)
for name, shape, dtype in spec
}
)
for name, shape, dtype in spec:
t = out[name]
assert (
tuple(t.shape) == shape and t.dtype == dtype and t.is_contiguous()
), f"{name}: got {tuple(t.shape)}/{t.dtype}, want {shape}/{dtype} contiguous"
return out
CHUNK_SIZE = 64 CHUNK_SIZE = 64
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -260,23 +303,25 @@ def _fused_decode_metadata_kernel(
def fused_decode_sconv_metadata( def fused_decode_sconv_metadata(
B: int, cache_indices: torch.Tensor B: int, cache_indices: torch.Tensor, out: SconvMetadataOut | None = None
) -> tuple[torch.Tensor, torch.Tensor, SconvDecodeMetadata]: ) -> tuple[torch.Tensor, torch.Tensor, SconvDecodeMetadata]:
"""Single-launch replacement for the decode metadata prep: the two arange calls, """Single-launch replacement for the decode metadata prep: the two arange calls,
ones, `!= PAD`, `&`, `clamp` and `.long()` that ones, `!= PAD`, `&`, `clamp` and `.long()` that
``precompute_helion_decode_metadata`` (+ its callers) issued as ~7 tiny ``precompute_helion_decode_metadata`` (+ its callers) issued as ~7 tiny
elementwise kernels. Returns elementwise kernels. Returns
``(query_start_loc, has_initial_state, SconvDecodeMetadata)`` with tensors ``(query_start_loc, has_initial_state, SconvDecodeMetadata)`` with tensors
bit-identical to the unfused path. bit-identical to the unfused path. Pass ``out`` to write into preallocated
(e.g. cuda-graph-static) destinations instead of fresh allocations.
""" """
assert cache_indices.shape[0] == B and cache_indices.stride(0) == 1 assert cache_indices.shape[0] == B and cache_indices.stride(0) == 1
device = cache_indices.device device = cache_indices.device
query_start_loc = torch.empty(B + 1, dtype=torch.int32, device=device) dst = _metadata_out(out, B=B, T=B, device=device)
has_initial_state = torch.empty(B, dtype=torch.bool, device=device) query_start_loc = dst["query_start_loc"]
cache_mask = torch.empty((B, 1, 1), dtype=torch.bool, device=device) has_initial_state = dst["has_initial_state"]
safe_idx = torch.empty(B, dtype=torch.int64, device=device) cache_mask = dst["cache_mask"]
cu = torch.empty(B + 1, dtype=torch.int64, device=device) safe_idx = dst["safe_idx"]
si = torch.empty(B, dtype=torch.int32, device=device) cu = dst["cu"]
si = dst["si"]
BLOCK = 1024 BLOCK = 1024
_fused_decode_metadata_kernel[(triton.cdiv(B + 1, BLOCK),)]( _fused_decode_metadata_kernel[(triton.cdiv(B + 1, BLOCK),)](
cache_indices, cache_indices,
@@ -419,13 +464,15 @@ def fused_extend_sconv_metadata(
extend_seq_lens: torch.Tensor | None = None, extend_seq_lens: torch.Tensor | None = None,
his_src: torch.Tensor | None = None, his_src: torch.Tensor | None = None,
draft_token_num: int | None = None, draft_token_num: int | None = None,
out: SconvMetadataOut | None = None,
) -> tuple[torch.Tensor, torch.Tensor, SconvExtendMetadata] | None: ) -> tuple[torch.Tensor, torch.Tensor, SconvExtendMetadata] | None:
"""Single-launch replacement for the extend metadata prep: the """Single-launch replacement for the extend metadata prep: the
zeros + cumsum(+scan-init) + slice-copy + compare chain of zeros + cumsum(+scan-init) + slice-copy + compare chain the unfused
``_prepare_extend_common_metadata`` plus the != PAD, &, clamp, long, to, ``query_start_loc`` / ``has_initial_state`` build issues, plus the != PAD, &,
arange, searchsorted, clamp, int32 chain of clamp, long, to, arange, searchsorted, clamp, int32 chain of
``precompute_helion_extend_metadata`` (~10-14 tiny kernels, re-issued per ``precompute_helion_extend_metadata`` (~10-14 tiny kernels, and before the
owning sconv instance -- and per de-tied draft step under draft_extend_v2). conv-state backend owned this prep, re-issued once per conv module of the
owning layer).
Returns ``(query_start_loc, has_initial_state, SconvExtendMetadata)`` with Returns ``(query_start_loc, has_initial_state, SconvExtendMetadata)`` with
tensors bit-identical to the unfused path, or None when the shape falls tensors bit-identical to the unfused path, or None when the shape falls
outside the fused kernel's single-tile bound (caller runs unfused). outside the fused kernel's single-tile bound (caller runs unfused).
@@ -433,7 +480,8 @@ def fused_extend_sconv_metadata(
``his_mode`` selects the has_initial_state source: HIS_ZEROS (boundary-KV ``his_mode`` selects the has_initial_state source: HIS_ZEROS (boundary-KV
draft extend), HIS_PREFIX (``his_src`` = extend_prefix_lens), HIS_SEQ_MINUS_EXT draft extend), HIS_PREFIX (``his_src`` = extend_prefix_lens), HIS_SEQ_MINUS_EXT
(``his_src`` = seq_lens), HIS_ONES (target_verify; ``draft_token_num`` set, (``his_src`` = seq_lens), HIS_ONES (target_verify; ``draft_token_num`` set,
``extend_seq_lens`` unused). ``extend_seq_lens`` unused). Pass ``out`` to write into preallocated (e.g.
cuda-graph-static) destinations instead of fresh allocations.
""" """
if B > _FUSED_EXTEND_MAX_B or not cache_indices.is_cuda: if B > _FUSED_EXTEND_MAX_B or not cache_indices.is_cuda:
return None return None
@@ -444,12 +492,13 @@ def fused_extend_sconv_metadata(
else: else:
assert extend_seq_lens is not None and extend_seq_lens.stride(0) == 1 assert extend_seq_lens is not None and extend_seq_lens.stride(0) == 1
device = cache_indices.device device = cache_indices.device
query_start_loc = torch.empty(B + 1, dtype=torch.int32, device=device) dst = _metadata_out(out, B=B, T=T, device=device)
has_initial_state = torch.empty(B, dtype=torch.bool, device=device) query_start_loc = dst["query_start_loc"]
cache_mask = torch.empty((B, 1, 1), dtype=torch.bool, device=device) has_initial_state = dst["has_initial_state"]
safe_idx = torch.empty(B, dtype=torch.int64, device=device) cache_mask = dst["cache_mask"]
cu = torch.empty(B + 1, dtype=torch.int64, device=device) safe_idx = dst["safe_idx"]
si = torch.empty(T, dtype=torch.int32, device=device) cu = dst["cu"]
si = dst["si"]
BLOCK_T = 256 BLOCK_T = 256
dummy = cache_indices # never dereferenced thanks to masks/constexpr dummy = cache_indices # never dereferenced thanks to masks/constexpr
_fused_extend_metadata_kernel[(1 + triton.cdiv(T, BLOCK_T),)]( _fused_extend_metadata_kernel[(1 + triton.cdiv(T, BLOCK_T),)](
+59 -340
View File
@@ -1,5 +1,4 @@
from enum import IntEnum from enum import IntEnum
from typing import Any
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -10,24 +9,16 @@ from torch.nn.parameter import Parameter
from sglang.srt.mem_cache.memory_pool import MambaPool from sglang.srt.mem_cache.memory_pool import MambaPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch 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.models.inkling_common.kernels.sconv import ( from sglang.srt.models.inkling_common.kernels.sconv import (
HIS_ONES,
HIS_PREFIX,
HIS_SEQ_MINUS_EXT,
HIS_ZEROS,
SconvDecodeMetadata, SconvDecodeMetadata,
SconvExtendMetadata, SconvExtendMetadata,
causal_conv1d, causal_conv1d,
fused_causal_conv1d_update_decode, fused_causal_conv1d_update_decode,
fused_decode_sconv_metadata,
fused_extend_sconv_metadata,
precompute_helion_extend_metadata,
save_intermediate_conv_windows, save_intermediate_conv_windows,
update_sconv_cache, update_sconv_cache,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
from sglang.srt.utils import is_cuda, set_weight_attrs from sglang.srt.utils import is_cuda, set_weight_attrs
@@ -40,10 +31,6 @@ class SconvType(IntEnum):
MLP = 5 MLP = 5
# Module-level cache for sconv metadata (shared across layers in the same forward pass)
_metadata_cache: dict = {}
class ShortConvolution(nn.Module): class ShortConvolution(nn.Module):
"""Short convolution layer for efficient causal convolution operations. """Short convolution layer for efficient causal convolution operations.
@@ -135,145 +122,20 @@ class ShortConvolution(nn.Module):
) )
param_data.copy_(loaded_weight) param_data.copy_(loaded_weight)
def _owns_extend_metadata(self, forward_batch: ForwardBatch) -> bool: def _conv_state(self, forward_batch: ForwardBatch):
# layer 0 computes the shared _metadata_cache for all layers within one """This layer's conv-state handle for the current step.
# forward. Under de-tied draft_extend_v2 each STEP is its own forward
# against its own pool, and only step 0's model carries layer_id == 0 —
# steps 1..N-1 would silently reuse a previous forward's cached (freed
# or wrong-pool) tensors, so every step must own its metadata.
return self.layer_id == 0 or forward_batch.forward_mode.is_draft_extend_v2()
def _prepare_extend_common_metadata( ``InklingShortConvAttnBackend`` resolved the whole step-global metadata set
self, forward_batch: ForwardBatch, cache_indices: torch.Tensor once during metadata prep, so this is a pure read shared by every conv
): module in the step.
"""Compute ALL extend sconv metadata (query_start_loc, has_initial_state, """
and the SconvExtendMetadata) in one fused launch and stash it in return get_attn_backend().conv_state_metadata(self.layer_id, forward_batch)
_metadata_cache; _prepare_extend_sconv_metadata is then a cache read.
Falls back to the original unfused op sequence off-CUDA or past the
fused kernel's batch bound."""
if self._owns_extend_metadata(forward_batch):
B = forward_batch.batch_size
is_verify = forward_batch.forward_mode.is_target_verify()
if is_verify:
# target_verify does not populate extend_seq_lens/extend_prefix_lens;
# the lens are a constant draft_token_num per request.
draft_token_num = forward_batch.spec_info.draft_token_num
num_tokens = B * draft_token_num
fused = fused_extend_sconv_metadata(
B=B,
T=num_tokens,
cache_indices=cache_indices,
his_mode=HIS_ONES,
draft_token_num=draft_token_num,
)
else:
num_tokens = forward_batch.extend_num_tokens
spec_info = forward_batch.spec_info
if (
isinstance(spec_info, EagleDraftExtendInput)
and spec_info.num_front_tokens > 0
):
# Boundary-KV fix: run conv fresh so warm-up rows rebuild
# the window.
his_mode, his_src = HIS_ZEROS, None
elif forward_batch.extend_prefix_lens is not None:
his_mode, his_src = HIS_PREFIX, forward_batch.extend_prefix_lens
else:
# draft_extend_v2 capture has no extend_prefix_lens.
his_mode, his_src = HIS_SEQ_MINUS_EXT, forward_batch.seq_lens
fused = fused_extend_sconv_metadata(
B=B,
T=num_tokens,
cache_indices=cache_indices,
his_mode=his_mode,
extend_seq_lens=forward_batch.extend_seq_lens,
his_src=his_src,
)
if fused is not None:
query_start_loc, has_initial_state, precomputed = fused
else:
query_start_loc, has_initial_state = (
self._unfused_extend_common_metadata(forward_batch)
)
precomputed = precompute_helion_extend_metadata(
B=B,
T=num_tokens,
W=self.kernel_size[0],
cache_indices=cache_indices,
has_initial_state=has_initial_state,
query_start_loc=query_start_loc,
)
_metadata_cache["query_start_loc"] = query_start_loc
_metadata_cache["has_initial_state"] = has_initial_state
_metadata_cache["helion_precomputed_extend"] = precomputed
return _metadata_cache["query_start_loc"], _metadata_cache["has_initial_state"]
def _unfused_extend_common_metadata(self, forward_batch: ForwardBatch): def _sconv_cache(self, meta) -> torch.Tensor:
"""Original multi-kernel query_start_loc/has_initial_state prep; fused return meta.layer_cache.conv[self.sconv_type.value]
fallback only."""
device = forward_batch.req_pool_indices.device
if forward_batch.forward_mode.is_target_verify():
draft_token_num = forward_batch.spec_info.draft_token_num
query_start_loc = torch.arange(
0,
(forward_batch.batch_size + 1) * draft_token_num,
draft_token_num,
dtype=torch.int32,
device=device,
)
has_initial_state = torch.ones(
forward_batch.batch_size, dtype=torch.bool, device=device
)
return query_start_loc, has_initial_state
query_start_loc = torch.zeros(
forward_batch.batch_size + 1,
dtype=torch.int32,
device=device,
)
query_start_loc[1:] = forward_batch.extend_seq_lens.cumsum(dim=0)
spec_info = forward_batch.spec_info
if (
isinstance(spec_info, EagleDraftExtendInput)
and spec_info.num_front_tokens > 0
):
has_initial_state = torch.zeros(
forward_batch.batch_size, dtype=torch.bool, device=device
)
elif forward_batch.extend_prefix_lens is not None:
has_initial_state = forward_batch.extend_prefix_lens > 0
else:
has_initial_state = (
forward_batch.seq_lens[: forward_batch.batch_size]
- forward_batch.extend_seq_lens
) > 0
return query_start_loc, has_initial_state
def _prepare_extend_sconv_metadata( def _weight_2d(self) -> torch.Tensor:
self, forward_batch: ForwardBatch, cache_indices: torch.Tensor return rearrange(self.weight, "d 1 w -> d w")
) -> SconvExtendMetadata | Any:
# Filled by _prepare_extend_common_metadata, which every caller invokes
# first with the same cache_indices (the fused kernel produces the
# whole metadata set in one launch).
del forward_batch, cache_indices
return _metadata_cache["helion_precomputed_extend"]
def _prepare_decode_sconv_metadata(
self, forward_batch: ForwardBatch, cache_indices: torch.Tensor
):
if self.layer_id == 0:
query_start_loc, has_initial_state, precomputed = (
fused_decode_sconv_metadata(
B=forward_batch.batch_size, cache_indices=cache_indices
)
)
_metadata_cache["query_start_loc_decode"] = query_start_loc
_metadata_cache["has_initial_state_decode"] = has_initial_state
_metadata_cache["helion_precomputed_decode"] = precomputed
return (
_metadata_cache["query_start_loc_decode"],
_metadata_cache["has_initial_state_decode"],
_metadata_cache["helion_precomputed_decode"],
)
def _apply_training_sconv_kernel( def _apply_training_sconv_kernel(
self, self,
@@ -304,123 +166,22 @@ class ShortConvolution(nn.Module):
) )
return y return y
def _init_track_conv_indices(
self, query_start_loc: torch.Tensor, forward_batch: ForwardBatch
):
"""
Compute indices for extracting conv states from the input sequence during extend.
In Mamba models, the conv layer maintains a sliding window of recent inputs.
After processing a prefill chunk, we need to save the last `conv_state_len` tokens
of the processed region for prefix caching.
The key insight is that FLA (Flash Linear Attention) processes sequences in chunks
of FLA_CHUNK_SIZE. We only track the conv state up to the last complete chunk boundary
(aligned_len).
start_indices is the starting token index of the conv state to track in this extend batch.
indices include all pos to track in this extend batch, conv_state_len for each req that
needs to be tracked (i.e. mamba_track_mask is True)
Returns:
indices: Tensor of shape [num_tracked_requests, conv_state_len] containing
flattened positions into the packed input tensor.
"""
conv_state_len = self.kernel_size[0] - 1
# Calculate the end position of the last aligned chunk
lens_to_track = (
forward_batch.mamba_track_seqlens - forward_batch.extend_prefix_lens
)
mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
chunk_aligned_lens_to_track = (
lens_to_track // mamba_cache_chunk_size
) * mamba_cache_chunk_size
start_indices = (
query_start_loc[:-1] + chunk_aligned_lens_to_track - conv_state_len
)
# Create indices: [batch_size, conv_state_len] or padded batch_size in prefill cudagraph
indices = start_indices.unsqueeze(-1) + torch.arange(
conv_state_len,
device=forward_batch.req_pool_indices.device,
dtype=start_indices.dtype,
)
# Use slice [-1:] instead of [-1] to avoid 0-d tensor -> scalar conversion during graph capture
return torch.clamp(
indices,
min=torch.zeros(
(1,),
dtype=start_indices.dtype,
device=forward_batch.req_pool_indices.device,
),
max=query_start_loc[-1:] - 1,
)
def _prepare_extend_track_conv_indices(
self, query_start_loc: torch.Tensor, forward_batch: ForwardBatch
) -> torch.Tensor:
if self.layer_id == 0:
track_conv_indices = self._init_track_conv_indices(
query_start_loc, forward_batch
)
_metadata_cache["track_conv_indices_extend"] = track_conv_indices
return _metadata_cache["track_conv_indices_extend"]
def _prepare_cache_indices(
self, req_to_token_pool, forward_batch: ForwardBatch
) -> torch.Tensor:
"""Resolve the per-request mamba slot indices ONCE per forward step.
``get_mamba_indices`` is a GPU gather
(``req_index_to_mamba_index_mapping[req_pool_indices]``) that depends
only on ``forward_batch.req_pool_indices``, which is invariant across
every sconv layer within a step. Computing it in each layer's
``forward`` launched one redundant gather kernel per k_sconv/v_sconv
(``2 * num_attn_layers`` per step). Cache the layer-0 result in the
shared per-step metadata cache and hand it back to subsequent layers,
so all layers reuse the same resolved indices.
Cuda-graph-safe: on capture only layer 0's gather is recorded and
subsequent layers read that captured tensor; on replay layer 0's gather
re-runs into the same address, keeping it current -- the same mechanism
the other ``_metadata_cache`` entries already rely on.
Under de-tied DRAFT_EXTEND_V2 each per-step forward runs with
layer_id != 0 against its own draft pool, so every step must own its
gather instead of reusing another forward's cached tensor (same rule
as ``_owns_extend_metadata``).
"""
if self._owns_extend_metadata(forward_batch):
_metadata_cache["cache_indices"] = (
req_to_token_pool.translate_mamba_indices(
req_to_token_pool.get_mamba_indices(forward_batch.req_pool_indices)
)
)
return _metadata_cache["cache_indices"]
def _prepare_extend_sconv_cache( def _prepare_extend_sconv_cache(
self, self,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
sconv_cache: torch.Tensor, sconv_cache: torch.Tensor,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
query_start_loc: torch.Tensor, track_conv_indices: torch.Tensor | None,
): ):
if forward_batch.mamba_track_mask is not None: if track_conv_indices is not None:
# Track conv state for prefix caching. Fused gather→scatter writes # Fused gather->scatter straight into sconv_cache, with no intermediate
# directly into sconv_cache without an intermediate [B, W-1, D] buffer. # [B, W-1, D] buffer.
conv_dst = forward_batch.mamba_track_indices
# [B, W - 1]
track_conv_indices = self._prepare_extend_track_conv_indices(
query_start_loc, forward_batch
)
fused_gather_scatter_to_sconv_cache( fused_gather_scatter_to_sconv_cache(
hidden_states=hidden_states, hidden_states=hidden_states,
sconv_cache=sconv_cache, sconv_cache=sconv_cache,
track_conv_indices=track_conv_indices, track_conv_indices=track_conv_indices,
mask=forward_batch.mamba_track_mask, mask=forward_batch.mamba_track_mask,
dst_indices=conv_dst, dst_indices=forward_batch.mamba_track_indices,
) )
def _save_intermediate_conv_windows( def _save_intermediate_conv_windows(
@@ -436,8 +197,8 @@ class ShortConvolution(nn.Module):
Builds a padded sequence [initial_conv_state | draft_tokens] and extracts Builds a padded sequence [initial_conv_state | draft_tokens] and extracts
sliding windows of size (kernel_size - 1) after each draft token position. sliding windows of size (kernel_size - 1) after each draft token position.
These intermediate states are consumed by These intermediate states are consumed by
InklingForConditionalGeneration.update_conv_state_after_mtp_verify InklingShortConvAttnBackend.commit_conv_state_after_mtp_verify to restore
to restore the correct conv state for the number of accepted tokens. the correct conv state for the number of accepted tokens.
""" """
save_intermediate_conv_windows( save_intermediate_conv_windows(
sconv_cache=sconv_cache, sconv_cache=sconv_cache,
@@ -510,35 +271,30 @@ class ShortConvolution(nn.Module):
def decode_fused_ar_inputs(self, forward_batch: ForwardBatch): def decode_fused_ar_inputs(self, forward_batch: ForwardBatch):
"""Return inputs for fused decode all-reduce, convolution, and norm. """Return inputs for fused decode all-reduce, convolution, and norm.
These match the fused decode branch of ``forward``, including its These match the fused decode branch of ``forward``. Returns
per-step metadata cache behavior. Returns
``(sconv_cache, cache_indices, cache_mask, weight_2d)``.""" ``(sconv_cache, cache_indices, cache_mask, weight_2d)``."""
req_to_token_pool = get_req_to_token_pool() meta = self._conv_state(forward_batch)
cache = req_to_token_pool.mamba2_layer_cache(self.layer_id) return (
sconv_cache = cache.conv[self.sconv_type.value] self._sconv_cache(meta),
cache_indices = self._prepare_cache_indices(req_to_token_pool, forward_batch) meta.cache_indices,
_, _, precomputed = self._prepare_decode_sconv_metadata( meta.precomputed["cache_mask"],
forward_batch, cache_indices self._weight_2d(),
) )
weight = rearrange(self.weight, "d 1 w -> d w")
return sconv_cache, cache_indices, precomputed["cache_mask"], weight
def verify_fused_ar_inputs(self, forward_batch: ForwardBatch): def verify_fused_ar_inputs(self, forward_batch: ForwardBatch):
"""Return inputs for fused target-verify convolution and norm. """Return inputs for fused target-verify convolution and norm.
These mirror the target-verify branch of ``forward``. Returns ``(sconv_cache, These mirror the target-verify branch of ``forward``. Returns ``(sconv_cache,
cache_indices[B], has_initial_state[B], weight_2d, inter_out)``.""" cache_indices[B], has_initial_state[B], weight_2d, inter_out)``."""
req_to_token_pool = get_req_to_token_pool() meta = self._conv_state(forward_batch)
cache = req_to_token_pool.mamba2_layer_cache(self.layer_id)
sconv_cache = cache.conv[self.sconv_type.value]
cache_indices = self._prepare_cache_indices(req_to_token_pool, forward_batch)
_, has_initial_state = self._prepare_extend_common_metadata(
forward_batch, cache_indices
)
weight = rearrange(self.weight, "d 1 w -> d w")
inter_out = cache.intermediate_conv_window[self.sconv_type.value]
b = forward_batch.batch_size b = forward_batch.batch_size
return sconv_cache, cache_indices[:b], has_initial_state, weight, inter_out return (
self._sconv_cache(meta),
meta.cache_indices[:b],
meta.has_initial_state,
self._weight_2d(),
meta.layer_cache.intermediate_conv_window[self.sconv_type.value],
)
def extend_fused_ar_inputs(self, forward_batch: ForwardBatch): def extend_fused_ar_inputs(self, forward_batch: ForwardBatch):
"""Return inputs for fused extend all-reduce and scattered convolution. """Return inputs for fused extend all-reduce and scattered convolution.
@@ -551,32 +307,13 @@ class ShortConvolution(nn.Module):
prep); ``cache_indices``/``has_initial_state`` feed the in-kernel prep); ``cache_indices``/``has_initial_state`` feed the in-kernel
cache update, and ``track_rows``/``track_mask``/``track_dst`` feed the cache update, and ``track_rows``/``track_mask``/``track_dst`` feed the
in-kernel prefix-cache track.""" in-kernel prefix-cache track."""
req_to_token_pool = get_req_to_token_pool() meta = self._conv_state(forward_batch)
cache = req_to_token_pool.mamba2_layer_cache(self.layer_id) precomputed = meta.precomputed
sconv_cache = cache.conv[self.sconv_type.value] # The backend resolves track rows only for the extend modes that snapshot
cache_indices = self._prepare_cache_indices(req_to_token_pool, forward_batch) # windows -- never decode or target-verify.
weight = rearrange(self.weight, "d 1 w -> d w") dev = meta.cache_indices.device
if forward_batch.forward_mode.is_decode(): if meta.track_conv_indices is not None:
# Decode: every token its own sequence (arange qsl, has_init=ones). track_rows = meta.track_conv_indices
query_start_loc, has_initial_state, precomputed = (
self._prepare_decode_sconv_metadata(forward_batch, cache_indices)
)
else:
query_start_loc, has_initial_state = self._prepare_extend_common_metadata(
forward_batch, cache_indices
)
precomputed = self._prepare_extend_sconv_metadata(
forward_batch, cache_indices
)
# Prefix-cache track inputs (extend only; the kernel fuses the write).
dev = cache_indices.device
if (
forward_batch.mamba_track_mask is not None
and not forward_batch.forward_mode.is_decode()
):
track_rows = self._prepare_extend_track_conv_indices(
query_start_loc, forward_batch
).long()
track_mask = forward_batch.mamba_track_mask track_mask = forward_batch.mamba_track_mask
track_dst = forward_batch.mamba_track_indices track_dst = forward_batch.mamba_track_indices
else: else:
@@ -585,15 +322,15 @@ class ShortConvolution(nn.Module):
track_mask = torch.empty((0,), dtype=torch.bool, device=dev) track_mask = torch.empty((0,), dtype=torch.bool, device=dev)
track_dst = torch.empty((0,), dtype=torch.int64, device=dev) track_dst = torch.empty((0,), dtype=torch.int64, device=dev)
return ( return (
sconv_cache, self._sconv_cache(meta),
precomputed["safe_idx"], precomputed["safe_idx"],
precomputed["cache_mask"].view(-1), precomputed["cache_mask"].view(-1),
precomputed["cu"], precomputed["cu"],
precomputed["si"], precomputed["si"],
weight, self._weight_2d(),
query_start_loc, meta.query_start_loc,
cache_indices, meta.cache_indices,
has_initial_state, meta.has_initial_state,
track_rows, track_rows,
track_mask, track_mask,
track_dst, track_dst,
@@ -607,15 +344,13 @@ class ShortConvolution(nn.Module):
) -> None: ) -> None:
"""Target-verify finish for the fused {AR + scattered sconv} path: no """Target-verify finish for the fused {AR + scattered sconv} path: no
working-cache update; save the per-position windows (consumed by working-cache update; save the per-position windows (consumed by
update_conv_state_after_mtp_verify), exactly as the verify branch of commit_conv_state_after_mtp_verify), exactly as the verify branch of
``forward`` does -- on the reduced pre-conv x.""" ``forward`` does -- on the reduced pre-conv x."""
req_to_token_pool = get_req_to_token_pool() meta = self._conv_state(forward_batch)
cache = req_to_token_pool.mamba2_layer_cache(self.layer_id)
sconv_cache = cache.conv[self.sconv_type.value]
self._save_intermediate_conv_windows( self._save_intermediate_conv_windows(
forward_batch=forward_batch, forward_batch=forward_batch,
cache=cache, cache=meta.layer_cache,
sconv_cache=sconv_cache, sconv_cache=self._sconv_cache(meta),
cache_indices=cache_indices, cache_indices=cache_indices,
hidden_states=x_scratch, hidden_states=x_scratch,
) )
@@ -638,20 +373,13 @@ class ShortConvolution(nn.Module):
""" """
del positions del positions
req_to_token_pool = get_req_to_token_pool() meta = self._conv_state(forward_batch)
cache = req_to_token_pool.mamba2_layer_cache(self.layer_id) cache_indices = meta.cache_indices
sconv_cache = cache.conv[self.sconv_type.value] sconv_cache = self._sconv_cache(meta)
cache_indices = self._prepare_cache_indices(req_to_token_pool, forward_batch) precomputed = meta.precomputed
weight = self._weight_2d()
weight = rearrange(self.weight, "d 1 w -> d w")
if forward_batch.forward_mode.is_target_verify(): if forward_batch.forward_mode.is_target_verify():
query_start_loc, has_initial_state = self._prepare_extend_common_metadata(
forward_batch, cache_indices
)
precomputed = self._prepare_extend_sconv_metadata(
forward_batch, cache_indices
)
y = causal_conv1d( y = causal_conv1d(
x=hidden_states, x=hidden_states,
weight=weight, weight=weight,
@@ -663,23 +391,17 @@ class ShortConvolution(nn.Module):
) )
self._save_intermediate_conv_windows( self._save_intermediate_conv_windows(
forward_batch=forward_batch, forward_batch=forward_batch,
cache=cache, cache=meta.layer_cache,
sconv_cache=sconv_cache, sconv_cache=sconv_cache,
cache_indices=cache_indices, cache_indices=cache_indices,
hidden_states=hidden_states, hidden_states=hidden_states,
) )
elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True): elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
query_start_loc, has_initial_state = self._prepare_extend_common_metadata(
forward_batch, cache_indices
)
self._prepare_extend_sconv_cache( self._prepare_extend_sconv_cache(
forward_batch, sconv_cache, hidden_states, query_start_loc forward_batch, sconv_cache, hidden_states, meta.track_conv_indices
) )
precomputed = self._prepare_extend_sconv_metadata(
forward_batch, cache_indices
)
if forward_batch.forward_mode.is_draft_extend_v2(): if forward_batch.forward_mode.is_draft_extend_v2():
y = causal_conv1d( y = causal_conv1d(
x=hidden_states, x=hidden_states,
@@ -702,8 +424,8 @@ class ShortConvolution(nn.Module):
weight=weight, weight=weight,
sconv_cache=sconv_cache, sconv_cache=sconv_cache,
cache_indices=cache_indices, cache_indices=cache_indices,
query_start_loc=query_start_loc, query_start_loc=meta.query_start_loc,
has_initial_state=has_initial_state, has_initial_state=meta.has_initial_state,
precomputed=precomputed, precomputed=precomputed,
is_decode=False, is_decode=False,
) )
@@ -714,9 +436,6 @@ class ShortConvolution(nn.Module):
# into the persistent ping-pong slot in-register (no separate # into the persistent ping-pong slot in-register (no separate
# copy_if_needed launch). track_mask is None when prefix caching with the # copy_if_needed launch). track_mask is None when prefix caching with the
# mamba extra buffer is disabled, which disables the track-copy path. # mamba extra buffer is disabled, which disables the track-copy path.
_query_start_loc, _has_initial_state, precomputed = (
self._prepare_decode_sconv_metadata(forward_batch, cache_indices)
)
y = fused_causal_conv1d_update_decode( y = fused_causal_conv1d_update_decode(
x=hidden_states, x=hidden_states,
weight=weight, weight=weight,
@@ -328,15 +328,11 @@ class DFlashWorkerV2(BaseSpecWorker):
def init_attention_backends(self): def init_attention_backends(self):
self._draft_worker.init_attention_backends() self._draft_worker.init_attention_backends()
target_model = self.model_runner.model
self._need_mamba_verify_commit = mambaish_config( self._need_mamba_verify_commit = mambaish_config(
self.model_runner.model_config self.model_runner.model_config
) is not None and ( ) is not None and hasattr(
hasattr( self.model_runner.attn_backend,
self.model_runner.attn_backend, "update_mamba_state_after_mtp_verify",
"update_mamba_state_after_mtp_verify",
)
or hasattr(target_model, "update_conv_state_after_mtp_verify")
) )
def init_cuda_graphs(self): def init_cuda_graphs(self):
@@ -1293,17 +1289,7 @@ class DFlashWorkerV2(BaseSpecWorker):
mamba_track_indices=batch.mamba_track_indices, mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track, mamba_steps_to_track=mamba_steps_to_track,
model=model_runner.model, model=model_runner.model,
)
elif hasattr(model_runner.model, "update_conv_state_after_mtp_verify"):
# Inkling's short convolutions access the mamba pool directly, so
# their accepted verify state is committed by the model rather
# than an attention-backend wrapper.
model_runner.model.update_conv_state_after_mtp_verify(
req_to_token_pool=model_runner.req_to_token_pool,
req_pool_indices=batch.req_pool_indices[: commit_lens.shape[0]], req_pool_indices=batch.req_pool_indices[: commit_lens.shape[0]],
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
) )
def _ensure_accept_bonus_buffers(self, bs: int) -> None: def _ensure_accept_bonus_buffers(self, bs: int) -> None:
+26 -1
View File
@@ -9,6 +9,21 @@ from sglang.srt.utils.common import (
) )
def _assert_draft_needs_no_conv_sidecar(draft_model_runner) -> None:
"""Refuse a multi-step draft decode backend for a draft with conv layers."""
from sglang.srt.configs.inkling import InklingMMConfig, InklingModelConfig
if isinstance(
draft_model_runner.model_config.hf_config,
(InklingModelConfig, InklingMMConfig),
):
raise NotImplementedError(
"Inkling's draft model runs its own short convs, which need the "
"conv-state sidecar the multi-step draft decode backend cannot carry. "
"Use --enable-multi-layer-eagle."
)
class DraftBackendFactory: class DraftBackendFactory:
def __init__( def __init__(
self, self,
@@ -46,6 +61,10 @@ class DraftBackendFactory:
if self.speculative_num_steps <= 1: if self.speculative_num_steps <= 1:
return None return None
# Returns a per-step CONTAINER, not an AttentionBackend, so
# attn_backend_wrapper_for_draft_extend cannot give it a conv sidecar.
_assert_draft_needs_no_conv_sidecar(self.draft_model_runner)
backend_map = { backend_map = {
"flashinfer": self._create_flashinfer_decode_backend, "flashinfer": self._create_flashinfer_decode_backend,
"triton": self._create_triton_decode_backend, "triton": self._create_triton_decode_backend,
@@ -96,11 +115,17 @@ class DraftBackendFactory:
if self.server_args.speculative_attention_mode == "decode" if self.server_args.speculative_attention_mode == "decode"
else "prefill_attention_backend" else "prefill_attention_backend"
) )
return self._create_backend( backend = self._create_backend(
backend_name, backend_name,
backend_map, backend_map,
"EAGLE is not supported in attention backend {backend_type}", "EAGLE is not supported in attention backend {backend_type}",
) )
# A draft with conv layers of its own (Inkling) needs its sidecar here too.
from sglang.srt.layers.attention.attention_registry import (
attn_backend_wrapper_for_draft_extend,
)
return attn_backend_wrapper_for_draft_extend(self.draft_model_runner, backend)
def _create_dsa_decode_backend(self): def _create_dsa_decode_backend(self):
from sglang.srt.layers.attention.dsa_backend import ( from sglang.srt.layers.attention.dsa_backend import (
@@ -957,16 +957,7 @@ def commit_mamba_states_after_verify(
mamba_track_indices=batch.mamba_track_indices, mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track, mamba_steps_to_track=mamba_steps_to_track,
model=model_runner.model, model=model_runner.model,
)
elif hasattr(model_runner.model, "update_conv_state_after_mtp_verify"):
# Models whose conv layers bypass the attention-backend wrapper
# (Inkling) own the commit themselves.
model_runner.model.update_conv_state_after_mtp_verify(
req_to_token_pool=model_runner.req_to_token_pool,
req_pool_indices=batch.req_pool_indices[:bs], req_pool_indices=batch.req_pool_indices[:bs],
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
) )
@@ -1,6 +1,6 @@
"""fused_decode_sconv_metadata must be bit-identical to the unfused prep. """fused_decode_sconv_metadata must be bit-identical to the unfused prep.
The unfused reference is the exact op sequence `_prepare_decode_sconv_metadata` The unfused reference is the exact op sequence the decode metadata prep
used to launch: two arange calls + ones + precompute_helion_decode_metadata used to launch: two arange calls + ones + precompute_helion_decode_metadata
(!= PAD, &, clamp, long, arange x2). (!= PAD, &, clamp, long, arange x2).
""" """
@@ -1,6 +1,6 @@
"""fused_extend_sconv_metadata must be bit-identical to the unfused prep. """fused_extend_sconv_metadata must be bit-identical to the unfused prep.
The unfused reference is the exact op sequence _prepare_extend_common_metadata The unfused reference is the exact op sequence the extend metadata prep
+ precompute_helion_extend_metadata used to launch: zeros + cumsum + slice-copy + precompute_helion_extend_metadata used to launch: zeros + cumsum + slice-copy
(or arange + ones for verify) + the has_initial_state compare, then != PAD, &, (or arange + ones for verify) + the has_initial_state compare, then != PAD, &,
clamp, long, to(int64), arange, searchsorted, clamp, to(int32). clamp, long, to(int64), arange, searchsorted, clamp, to(int32).
@@ -0,0 +1,352 @@
"""Inkling's short-conv metadata must be resolved exactly ONCE per forward step.
A decoder layer holds FOUR ``ShortConvolution`` modules, and per-layer ownership
would recompute the whole set once per module. Pinned here: one resolution per step
however many modules ask, every module gets the *same* tensors, and the graph-path
destinations stay address-stable across steps -- including across a later
``init_cuda_graph_state``, where reallocating would move an address an
already-captured prefill graph reads.
"""
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
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-small")
NUM_LAYERS = 4
NUM_SCONV_STREAMS = 6 # pool-wide streams: k/v full, k/v local, attn, mlp
NUM_MODULES_PER_LAYER = 4 # k_sconv, v_sconv, attn_sconv, mlp_sconv
POOL_SLOTS = 32
CONV_KERNEL = 4
CONV_DIM = 8
class _MockMambaPool:
enable_linear_replayssm = False
def __init__(self):
conv = [
torch.zeros(
(NUM_LAYERS, POOL_SLOTS + 1, CONV_KERNEL - 1, CONV_DIM),
dtype=torch.bfloat16,
device="cuda",
)
for _ in range(NUM_SCONV_STREAMS)
]
self.mamba_cache = SimpleNamespace(conv=conv, temporal=None)
def mamba2_layer_cache(self, layer_id: int):
return SimpleNamespace(
conv=[c[layer_id] for c in self.mamba_cache.conv],
intermediate_conv_window=None,
)
class _MockReqToTokenPool:
"""The four methods the backend calls, plus ``size`` (its max-bs bound)."""
def __init__(self):
self.size = POOL_SLOTS
self.mamba_pool = _MockMambaPool()
self.req_index_to_mamba_index_mapping = torch.arange(
POOL_SLOTS + 1, dtype=torch.int32, device="cuda"
)
self.gather_calls = 0
def get_mamba_indices(self, req_indices: torch.Tensor) -> torch.Tensor:
self.gather_calls += 1
return self.req_index_to_mamba_index_mapping[req_indices]
def translate_mamba_indices(self, mamba_indices: torch.Tensor) -> torch.Tensor:
return mamba_indices
def mamba2_layer_cache(self, layer_id: int):
return self.mamba_pool.mamba2_layer_cache(layer_id)
def get_speculative_mamba2_params_all_layers(self):
return self.mamba_pool.mamba_cache
def _decode_batch(bs: int):
return SimpleNamespace(
forward_mode=ForwardMode.DECODE,
batch_size=bs,
req_pool_indices=torch.arange(bs, dtype=torch.int64, device="cuda"),
seq_lens=torch.full((bs,), 64, dtype=torch.int64, device="cuda"),
spec_info=None,
mamba_track_mask=None,
mamba_track_seqlens=None,
mamba_track_indices=None,
)
def _extend_batch(seq_lens):
bs = len(seq_lens)
lens = torch.tensor(seq_lens, dtype=torch.int64, device="cuda")
return SimpleNamespace(
forward_mode=ForwardMode.EXTEND,
batch_size=bs,
req_pool_indices=torch.arange(bs, dtype=torch.int64, device="cuda"),
seq_lens=lens,
extend_seq_lens=lens,
extend_prefix_lens=torch.zeros(bs, dtype=torch.int64, device="cuda"),
extend_num_tokens=int(sum(seq_lens)),
spec_info=None,
mamba_track_mask=torch.ones(bs, dtype=torch.bool, device="cuda"),
mamba_track_seqlens=lens,
mamba_track_indices=torch.arange(bs, dtype=torch.int64, device="cuda"),
)
class TestInklingSconvMetadataOnce(CustomTestCase):
@classmethod
def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("Inkling's conv metadata kernels are CUDA-only.")
server_args = ServerArgs(
model_path="dummy",
page_size=1,
# Skips the model-config load in the Inkling prefill-graph default.
disable_prefill_cuda_graph=True,
disable_cuda_graph=True,
)
# Pre-seed the cached property so it does not reach for a real HF config.
server_args._mamba_cache_chunk_size = 64
set_global_server_args_for_scheduler(server_args)
def _build_backend(self):
from sglang.srt.layers.attention.linear.inkling_sconv_backend import (
InklingShortConvAttnBackend,
)
pool = _MockReqToTokenPool()
from sglang.srt.runtime_context import get_server_args
runner = SimpleNamespace(
device="cuda",
server_args=get_server_args(),
is_draft_worker=False,
req_to_token_pool=pool,
token_to_kv_pool=None,
)
return InklingShortConvAttnBackend(runner), pool
def _count_fused_calls(self, backend):
"""Wrap the two fused metadata entry points with counters."""
import sglang.srt.layers.attention.linear.inkling_sconv_backend as mod
counts = {"decode": 0, "extend": 0}
real_decode = mod.fused_decode_sconv_metadata
real_extend = mod.fused_extend_sconv_metadata
def decode(*a, **kw):
counts["decode"] += 1
return real_decode(*a, **kw)
def extend(*a, **kw):
counts["extend"] += 1
return real_extend(*a, **kw)
mod.fused_decode_sconv_metadata = decode
mod.fused_extend_sconv_metadata = extend
self.addCleanup(setattr, mod, "fused_decode_sconv_metadata", real_decode)
self.addCleanup(setattr, mod, "fused_extend_sconv_metadata", real_extend)
return counts
def _drain_all_conv_modules(self, backend, forward_batch):
"""Mimic every ShortConvolution in the model asking for its handle."""
handles = []
for layer_id in range(NUM_LAYERS):
for _module in range(NUM_MODULES_PER_LAYER):
handles.append(backend.conv_state_metadata(layer_id, forward_batch))
return handles
def test_decode_resolves_once_per_step(self):
backend, pool = self._build_backend()
counts = self._count_fused_calls(backend)
fb = _decode_batch(bs=3)
backend.init_forward_metadata(fb)
handles = self._drain_all_conv_modules(backend, fb)
self.assertEqual(counts["decode"], 1)
self.assertEqual(pool.gather_calls, 1)
self.assertEqual(len(handles), NUM_LAYERS * NUM_MODULES_PER_LAYER)
first = handles[0]
for h in handles[1:]:
self.assertIs(h.cache_indices, first.cache_indices)
self.assertIs(h.precomputed, first.precomputed)
self.assertIs(h.query_start_loc, first.query_start_loc)
self.assertIs(h.has_initial_state, first.has_initial_state)
def test_extend_resolves_once_per_step(self):
backend, pool = self._build_backend()
counts = self._count_fused_calls(backend)
fb = _extend_batch([7, 5, 3])
backend.init_forward_metadata(fb)
handles = self._drain_all_conv_modules(backend, fb)
self.assertEqual(counts["extend"], 1)
self.assertEqual(pool.gather_calls, 1)
first = handles[0]
self.assertIsNotNone(first.track_conv_indices)
self.assertEqual(tuple(first.track_conv_indices.shape), (3, CONV_KERNEL - 1))
for h in handles[1:]:
self.assertIs(h.track_conv_indices, first.track_conv_indices)
self.assertIs(h.precomputed, first.precomputed)
def test_each_step_re_resolves(self):
"""A second forward must recompute; nothing may leak across steps."""
backend, pool = self._build_backend()
counts = self._count_fused_calls(backend)
fb = _decode_batch(bs=2)
for _ in range(3):
backend.init_forward_metadata(fb)
self._drain_all_conv_modules(backend, fb)
self.assertEqual(counts["decode"], 3)
self.assertEqual(pool.gather_calls, 3)
def test_graph_destinations_are_address_stable(self):
for slots_in_graph in (False, True):
with self.subTest(slots_in_graph=slots_in_graph):
self._check_address_stable(slots_in_graph)
def _check_address_stable(self, slots_in_graph: bool):
"""A captured graph holds each metadata tensor's address, so steps refill in
place and a later ``init_cuda_graph_state`` must not reallocate."""
backend, _pool = self._build_backend()
# Cover both halves of the slot split (the mock's translate is not the base
# one, so slots would otherwise always stay eager).
backend._slot_gather_recordable = slots_in_graph
fb = _decode_batch(bs=2)
# Mirrors the decode runner: out-of-graph prep, then the recorded hook.
backend.init_forward_metadata_out_graph(fb, in_capture=True)
backend.init_forward_metadata_in_graph(fb)
h0 = backend.conv_state_metadata(0, fb)
ptrs = (
h0.cache_indices.data_ptr(),
h0.query_start_loc.data_ptr(),
h0.has_initial_state.data_ptr(),
h0.precomputed["cache_mask"].data_ptr(),
h0.precomputed["safe_idx"].data_ptr(),
h0.precomputed["cu"].data_ptr(),
h0.precomputed["si"].data_ptr(),
)
backend.init_cuda_graph_state(max_bs=8, max_num_tokens=8)
backend.init_forward_metadata_out_graph(fb)
backend.init_forward_metadata_in_graph(fb)
h1 = backend.conv_state_metadata(0, fb)
self.assertEqual(
ptrs,
(
h1.cache_indices.data_ptr(),
h1.query_start_loc.data_ptr(),
h1.has_initial_state.data_ptr(),
h1.precomputed["cache_mask"].data_ptr(),
h1.precomputed["safe_idx"].data_ptr(),
h1.precomputed["cu"].data_ptr(),
h1.precomputed["si"].data_ptr(),
),
)
class TestInklingMtpVerifyCommit(CustomTestCase):
"""The commit runs after the forward context exits, so the per-step slot buffer
may already belong to a later forward. Sourcing slot ids from
``forward_metadata`` (as the generic mamba path does) therefore mismatches the
verify batch; they must come from the passed ``req_pool_indices``.
"""
@classmethod
def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("Inkling's conv-state kernels are CUDA-only.")
TestInklingSconvMetadataOnce.setUpClass()
def _build_wrapper(self):
from sglang.srt.layers.attention.linear.inkling_sconv_backend import (
InklingShortConvAttnBackend,
InklingShortConvHybridAttnBackend,
)
from sglang.srt.runtime_context import get_server_args
pool = _MockReqToTokenPool()
runner = SimpleNamespace(
device="cuda",
server_args=get_server_args(),
is_draft_worker=False,
req_to_token_pool=pool,
token_to_kv_pool=None,
)
sidecar = InklingShortConvAttnBackend(runner)
full = SimpleNamespace(
token_to_kv_pool=None,
req_to_token_pool=pool,
needs_cpu_seq_lens=True,
)
wrapper = InklingShortConvHybridAttnBackend(
full, sidecar, list(range(NUM_LAYERS))
)
return wrapper, sidecar, pool
def test_commit_uses_passed_req_pool_indices_not_step_metadata(self):
wrapper, sidecar, pool = self._build_wrapper()
# The hazard: a later forward left a SHORTER slot buffer than the verify
# batch this commit is for.
sidecar.init_forward_metadata(_decode_batch(bs=3))
self.assertEqual(sidecar._cache_indices.shape[0], 3)
seen = {}
def fake_scatter(caches, state_indices, last_correct, track, steps):
seen["state_indices"] = state_indices
import sglang.srt.layers.attention.linear.inkling_sconv_backend as mod
real = mod.scatter_mamba_states_after_mtp_verify
mod.scatter_mamba_states_after_mtp_verify = fake_scatter
self.addCleanup(setattr, mod, "scatter_mamba_states_after_mtp_verify", real)
req_pool_indices = torch.arange(5, dtype=torch.int64, device="cuda")
wrapper.update_mamba_state_after_mtp_verify(
last_correct_step_indices=torch.zeros(5, dtype=torch.int64, device="cuda"),
mamba_track_indices=None,
mamba_steps_to_track=None,
model=None,
req_pool_indices=req_pool_indices,
)
# 5 rows from req_pool_indices, not the 3 on the step buffer.
self.assertEqual(seen["state_indices"].shape[0], 5)
self.assertTrue(
torch.equal(seen["state_indices"], pool.get_mamba_indices(req_pool_indices))
)
def test_commit_requires_req_pool_indices(self):
"""The generic caller signature makes it optional; Inkling cannot guess it."""
wrapper, _sidecar, _pool = self._build_wrapper()
with self.assertRaises(AssertionError):
wrapper.update_mamba_state_after_mtp_verify(
last_correct_step_indices=torch.zeros(
2, dtype=torch.int64, device="cuda"
),
mamba_track_indices=None,
mamba_steps_to_track=None,
model=None,
)
if __name__ == "__main__":
unittest.main()
@@ -3,6 +3,7 @@ from unittest.mock import MagicMock, patch
import torch import torch
import sglang.srt
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -370,5 +371,56 @@ class TestConvWindowDedupLayout(CustomTestCase):
) )
class TestMtpVerifyHookSignature(CustomTestCase):
"""Every ``update_mamba_state_after_mtp_verify`` override must accept the full
keyword call the spec workers make, or it raises TypeError at verify time on
whatever hardware it serves.
Parses sources rather than importing: the accelerator backends defining
overrides are exactly the ones whose deps are absent on most hosts, so an
import-based check would skip the cases that matter.
"""
CALL_KWARGS = {
"last_correct_step_indices",
"mamba_track_indices",
"mamba_steps_to_track",
"model",
"req_pool_indices",
}
HOOK = "update_mamba_state_after_mtp_verify"
def test_all_overrides_accept_the_call_kwargs(self):
import ast
import pathlib
srt = pathlib.Path(next(iter(sglang.srt.__path__)))
found = []
for path in srt.rglob("*.py"):
try:
tree = ast.parse(path.read_text())
except SyntaxError:
continue
for node in ast.walk(tree):
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
if node.name != self.HOOK:
continue
args = node.args
if args.kwarg is not None:
continue # **kwargs passthrough accepts everything
names = {a.arg for a in args.args} | {a.arg for a in args.kwonlyargs}
found.append((path.relative_to(srt), node.lineno, names))
self.assertTrue(found, f"no {self.HOOK} definitions found under sglang.srt")
for rel, lineno, names in found:
missing = self.CALL_KWARGS - names
self.assertFalse(
missing,
f"{rel}:{lineno} {self.HOOK} is missing {sorted(missing)}; "
"the spec workers call this hook by keyword.",
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()