feat(inkling): migrate short convs onto the ShortConv attention backend (#33023)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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),)](
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user