Inference Support Mamba 2 and 1 (#34556)

This commit is contained in:
desmond-intel
2026-09-12 20:51:18 +08:00
committed by GitHub
parent fd32226706
commit 6dc7b3421b
16 changed files with 1544 additions and 12 deletions
@@ -268,6 +268,11 @@ in the GitHub search bar.
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>tiiuae/Falcon-H1-34B-Instruct</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>TII's hybrid Mamba-Transformer architecture combining attention and state-space models for efficient long-context inference.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong>Mamba-Codestral</strong> (7B, Mamba2)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>mistralai/Mamba-Codestral-7B-v0.1</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Mistral AI's pure Mamba2 state-space model for code generation; uses selective state spaces instead of attention, so it holds no KV cache.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong>Hunyuan-Large</strong> (389B, MoE)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>tencent/Tencent-Hunyuan-Large</code></td>
+5
View File
@@ -46,6 +46,8 @@ from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig
from sglang.srt.configs.lfm2_vl import Lfm2VlConfig
from sglang.srt.configs.locate_anything import LocateAnythingConfig
from sglang.srt.configs.longcat_flash import LongcatFlashConfig
from sglang.srt.configs.mamba import FalconMambaConfig, MambaConfig
from sglang.srt.configs.mamba2 import Mamba2Config
from sglang.srt.configs.minicpm import MiniCPMHybridConfig
from sglang.srt.configs.minicpmv4_6 import MiniCPMV4_6Config, MiniCPMV4_6VisionConfig
from sglang.srt.configs.minimax_vl import MiniMaxM3VLConfig
@@ -126,8 +128,11 @@ __all__ = [
"DotsOCRConfig",
"Dots3Config",
"FalconH1Config",
"FalconMambaConfig",
"GraniteMoeHybridConfig",
"HYV4Config",
"MambaConfig",
"Mamba2Config",
"Lfm2Config",
"Lfm2MoeConfig",
"Lfm2VlConfig",
+7 -1
View File
@@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Any
from sglang.srt.configs import (
BailingHybridConfig,
FalconH1Config,
FalconMambaConfig,
GraniteMoeHybridConfig,
InklingMMConfig,
InklingModelConfig,
@@ -15,6 +16,8 @@ from sglang.srt.configs import (
Lfm2Config,
Lfm2MoeConfig,
Lfm2VlConfig,
Mamba2Config,
MambaConfig,
MiniCPMHybridConfig,
NemotronH_Nano_VL_V2_Config,
NemotronHConfig,
@@ -78,7 +81,10 @@ def mamba2_config(model_config: ModelConfig):
| Lfm2Config
| Lfm2MoeConfig
| Lfm2VlConfig
| ZayaConfig,
| ZayaConfig
| Mamba2Config
| MambaConfig
| FalconMambaConfig,
):
return config
if isinstance(config, InklingModelConfig):
+80
View File
@@ -0,0 +1,80 @@
# coding=utf-8
# Copyright 2024 The HuggingFace Inc. team. All rights reserved.
#
# 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.
"""Mamba (Mamba-1) model configuration for SGLang."""
from transformers import FalconMambaConfig as HFFalconMambaConfig
from transformers import MambaConfig as HFMambaConfig
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
# Mamba-1 has no chunk size; the Mamba2 backend only reads mamba_chunk_size to
# bound the conv window, so a constant is enough.
_MAMBA1_CACHE_CHUNK_SIZE = 256
def _mamba1_cache_params(config) -> Mamba2CacheParams:
from sglang.srt.runtime_context import get_parallel
parallel = get_parallel()
tp_world_size = parallel.tp_size if parallel else 1
shape = Mamba2StateShape.create_full_rank(
tp_world_size=tp_world_size,
intermediate_size=config.intermediate_size,
state_size=config.state_size,
conv_kernel=config.conv_kernel,
)
return Mamba2CacheParams(shape=shape, layers=list(range(config.num_hidden_layers)))
class MambaConfig(HFMambaConfig):
"""Config for pure Mamba-1 models (state-spaces Mamba, -hf and raw).
Subclasses the transformers MambaConfig and adds the same SSM hooks as
Mamba2Config. Mamba-1 runs on the Mamba2 backend through a full-rank
(head_dim == 1) state layout.
"""
model_type = "mamba"
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.mamba_chunk_size = _MAMBA1_CACHE_CHUNK_SIZE
@property
def full_attention_layer_ids(self) -> list[int]:
return []
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
return _mamba1_cache_params(self)
class FalconMambaConfig(HFFalconMambaConfig):
"""Config for Falcon-Mamba. Same as Mamba-1 aside from the B/C/dt RMS norm,
which the model file handles; the config hooks are identical."""
model_type = "falcon_mamba"
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.mamba_chunk_size = _MAMBA1_CACHE_CHUNK_SIZE
@property
def full_attention_layer_ids(self) -> list[int]:
return []
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
return _mamba1_cache_params(self)
+57
View File
@@ -0,0 +1,57 @@
# coding=utf-8
# Copyright 2024 Mistral AI and the HuggingFace Inc. team. All rights reserved.
#
# 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.
"""Mamba2 model configuration for SGLang."""
from transformers import Mamba2Config as HFMamba2Config
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
class Mamba2Config(HFMamba2Config):
"""Config for pure Mamba-2 models such as Mamba-Codestral-7B.
Subclasses the transformers Mamba2Config and adds the SSM hooks the Mamba2
attention backend expects, following NemotronHConfig.
"""
model_type = "mamba2"
def __init__(self, **kwargs):
super().__init__(**kwargs)
# Mamba2AttnBackend reads mamba_chunk_size; alias it to chunk_size.
self.mamba_chunk_size = self.chunk_size
@property
def full_attention_layer_ids(self) -> list[int]:
return []
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
from sglang.srt.runtime_context import get_parallel
parallel = get_parallel()
tp_world_size = parallel.tp_size if parallel else 1
shape = Mamba2StateShape.create(
tp_world_size=tp_world_size,
intermediate_size=self.intermediate_size,
n_groups=self.n_groups,
num_heads=self.num_heads,
head_dim=self.head_dim,
state_size=self.state_size,
conv_kernel=self.conv_kernel,
)
return Mamba2CacheParams(
shape=shape, layers=list(range(self.num_hidden_layers))
)
+47
View File
@@ -242,6 +242,53 @@ class Mamba2StateShape:
conv_shard_groups=conv_shard_groups,
)
@staticmethod
def create_full_rank(
*,
tp_world_size: int,
intermediate_size: int,
state_size: int,
conv_kernel: int,
) -> "Mamba2StateShape":
"""State shape for a full-rank (``head_dim == 1``) selective-scan mixer.
This is the layout used by Mamba-1 mixers (e.g. Falcon-Mamba,
state-spaces Mamba).
Two things differ from Mamba-2 (:meth:`create`):
- The causal conv is applied over ``intermediate_size`` ONLY. In Mamba-1
the ``B``/``C`` selection matrices are produced by ``x_proj`` *after*
the conv, so (unlike Mamba-2) they are not part of the conv input and
``conv_dim == intermediate_size``.
- The SSM ``A`` matrix / state is full-rank per channel with shape
``(intermediate_size, state_size)``. We express this on the Mamba-2
head layout as ``num_heads = intermediate_size`` and ``head_dim = 1``
(``n_groups`` implicitly 1, ``B``/``C`` shared across channels) so the
shared Mamba2 attention backend, memory pool, and
``selective_state_update`` kernel drive it unchanged.
"""
assert intermediate_size % tp_world_size == 0, (
f"Mamba-1 intermediate_size ({intermediate_size}) must be divisible "
f"by tp_world_size ({tp_world_size})"
)
conv_dim = intermediate_size
conv_state_shape = (divide(conv_dim, tp_world_size), conv_kernel - 1)
# (num_heads // tp, head_dim, state_size) with head_dim == 1.
temporal_state_shape = (divide(intermediate_size, tp_world_size), 1, state_size)
return Mamba2StateShape(
conv=[conv_state_shape],
temporal=temporal_state_shape,
intermediate_size=intermediate_size,
conv_dim=conv_dim,
ssm_state_size=state_size,
num_heads=intermediate_size,
head_dim=1,
state_size=state_size,
conv_kernel=conv_kernel,
num_k_heads_per_tp=1,
)
@dataclass(kw_only=True, frozen=True)
class Mamba2CacheParams(BaseLinearStateParams):
+42 -8
View File
@@ -101,6 +101,20 @@ def get_mimo_v2_fused_qkv_expected_tp_size(hf_config):
class AttentionArch(IntEnum):
MLA = auto()
MHA = auto()
SSM = auto() # State Space Models (Mamba, Mamba2)
# Pure Mamba-1 (selective-scan) archs; same mixer/state layout, differing only
# in cosmetic details handled in their model files.
PURE_MAMBA1_ARCHITECTURES = (
"FalconMambaForCausalLM",
"MambaForCausalLM",
)
# Pure state-space (SSM) causal-LMs: no attention, so no num_attention_heads /
# head_dim in their HF config. Used for head-dim derivation and attention-arch
# detection below.
PURE_SSM_ARCHITECTURES = ("Mamba2ForCausalLM",) + PURE_MAMBA1_ARCHITECTURES
class ModelImpl(str, Enum):
@@ -1043,14 +1057,23 @@ class ModelConfig:
def _derive_model_shapes(self):
from sglang.srt.configs.dots3 import Dots3Config
# Pure SSM models have no attention heads; use head_dim == 0 so the
# KV-cell size is 0 rather than a division on a missing head count.
is_pure_ssm = any(
arch in self.hf_config.architectures for arch in PURE_SSM_ARCHITECTURES
)
# Unify the config keys for hf_text_config
self.head_dim = getattr(self.hf_text_config, "head_dim", None)
if self.head_dim is None:
self.head_dim = (
self.hf_text_config.hidden_size
// self.hf_text_config.num_attention_heads
)
setattr(self.hf_text_config, "head_dim", self.head_dim)
if is_pure_ssm:
self.head_dim = 0
else:
self.head_dim = (
self.hf_text_config.hidden_size
// self.hf_text_config.num_attention_heads
)
setattr(self.hf_text_config, "head_dim", self.head_dim)
self.v_head_dim = getattr(self.hf_text_config, "v_head_dim", None)
if self.v_head_dim is None or self.v_head_dim == 0:
@@ -1228,9 +1251,16 @@ class ModelConfig:
elif "BaichuanForCausalLM" in self.hf_config.architectures:
self.use_alibi = self.hf_config.hidden_size != 4096
self.attention_arch = AttentionArch.MHA
# Pure Mamba SSMs have no attention (head_dim set to 0 above).
if is_pure_ssm:
self.attention_arch = AttentionArch.SSM
else:
self.attention_arch = AttentionArch.MHA
self.num_attention_heads = self.hf_text_config.num_attention_heads
# Mamba2 has no num_attention_heads.
self.num_attention_heads = getattr(
self.hf_text_config, "num_attention_heads", None
)
self.num_key_value_heads = getattr(
self.hf_text_config, "num_key_value_heads", None
)
@@ -1304,7 +1334,8 @@ class ModelConfig:
return self.num_attention_heads
def get_num_attention_heads(self, tensor_parallel_size) -> int:
total_num_attention_heads = self.num_attention_heads
# Pure-SSM (Mamba) models have no attention; num_attention_heads is None.
total_num_attention_heads = self.num_attention_heads or 0
return max(1, total_num_attention_heads // tensor_parallel_size)
# adapted from https://github.com/vllm-project/vllm/blob/main/vllm/config.py#L289
@@ -1367,6 +1398,9 @@ class ModelConfig:
if num_kv_heads is not None:
return num_kv_heads
# Mamba SSMs have no attention, so no KV heads.
if self.attention_arch == AttentionArch.SSM:
return 0
# For non-grouped-query attention models, the number of KV heads is
# equal to the number of attention heads.
return self.hf_text_config.num_attention_heads
@@ -0,0 +1,472 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright 2023-2025 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.
# ==============================================================================
"""Mamba-1 (selective-scan) SSM mixer for SGLang, e.g. Falcon-Mamba.
Unlike Mamba-2 (:class:`MambaMixer2`, SSD / chunked-scan, scalar per-head ``A``),
Mamba-1 keeps a **full-rank** ``A`` of shape ``(intermediate_size, state_size)``
and derives the selective parameters ``dt``/``B``/``C`` from ``x_proj`` applied
*after* the causal conv (so the conv is over ``intermediate_size`` only). This
matches HuggingFace ``FalconMambaMixer`` / the original ``MambaMixer``.
Reuse strategy (to ride the existing Mamba2 attention backend, memory pool and
kernels unchanged): the full-rank state is expressed on the Mamba2 head layout
as ``num_heads == intermediate_size`` and ``head_dim == 1`` (see
``Mamba2StateShape.create_full_rank``). Then:
- the causal conv uses the shared ``causal_conv1d_fn`` / ``causal_conv1d_update``
(Triton variants on XPU), exactly like Mamba2;
- single-token **decode** uses the shared ``selective_state_update`` kernel,
which already supports a full-rank ``A`` of shape ``(nheads, dim, dstate)``
and applies the ``silu(z)`` output gate;
- multi-token **prefill** runs a portable pure-torch selective scan (there is
no Mamba-1 chunked-scan kernel in-tree), which is device-agnostic and works
on Intel XPU.
Falcon-Mamba adds a weightless RMSNorm to ``B``, ``C`` and ``dt`` (the "Falcon"
stabilization trick), applied here via :func:`rms_normalize` gated on
``use_bc_dt_rms``.
"""
import logging
from typing import Optional, Tuple
import torch
from torch import nn
from sglang.kernels.ops.mamba.triton_ops import selective_state_update
from sglang.srt.distributed import divide
from sglang.srt.layers.attention.mamba.mamba import (
causal_conv1d_fn,
causal_conv1d_fn_triton,
causal_conv1d_update,
causal_conv1d_update_triton,
)
from sglang.srt.layers.linear import (
ColumnParallelLinear,
MergedColumnParallelLinear,
RowParallelLinear,
)
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.model_loader.weight_utils import sharded_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import set_weight_attrs
logger = logging.getLogger(__name__)
def rms_normalize(hidden_states: torch.Tensor, eps: float) -> torch.Tensor:
"""Weightless RMSNorm (matches HF ``falcon_mamba.rms_forward``).
Falcon-Mamba normalizes ``B``, ``C`` and the time step with a *non-learnable*
RMSNorm (no weight) before discretization; other Mamba-1 models skip this.
"""
input_dtype = hidden_states.dtype
hidden_states = hidden_states.to(torch.float32)
variance = hidden_states.pow(2).mean(-1, keepdim=True)
hidden_states = hidden_states * torch.rsqrt(variance + eps)
return hidden_states.to(input_dtype)
class MambaMixer1(nn.Module):
"""Mamba-1 selective-scan mixer.
Weight names match the HF checkpoint (``in_proj``, ``conv1d``, ``x_proj``,
``dt_proj``, ``A_log``, ``D``, ``out_proj``) so the model loader maps them
directly.
"""
def __init__(
self,
*,
hidden_size: int,
intermediate_size: int,
state_size: int,
conv_kernel: int,
time_step_rank: int,
use_conv_bias: bool,
use_bias: bool,
activation: str = "silu",
use_bc_dt_rms: bool = False,
rms_eps: float = 1e-6,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.tp_size = get_parallel().tp_size
self.hidden_size = hidden_size
self.intermediate_size = intermediate_size
self.ssm_state_size = state_size
self.conv_kernel_size = conv_kernel
self.time_step_rank = time_step_rank
self.activation = activation
self.use_bc_dt_rms = use_bc_dt_rms
self.rms_eps = rms_eps
assert intermediate_size % self.tp_size == 0, (
f"Mamba-1 intermediate_size ({intermediate_size}) must be divisible "
f"by tp_size ({self.tp_size})"
)
self.intermediate_size_per_tp = divide(intermediate_size, self.tp_size)
# in_proj -> [x, gate], each of size intermediate_size (column-sharded).
self.in_proj = MergedColumnParallelLinear(
input_size=hidden_size,
output_sizes=[intermediate_size, intermediate_size],
bias=use_bias,
quant_config=quant_config,
prefix=f"{prefix}.in_proj",
)
# Depthwise causal conv over the intermediate channels only (column-sharded).
self.conv1d = ColumnParallelLinear(
input_size=conv_kernel,
output_size=intermediate_size,
bias=use_conv_bias,
quant_config=None,
prefix=f"{prefix}.conv1d",
)
# Checkpoint stores conv1d.weight as (dim, 1, K); ColumnParallelLinear
# allocates (dim, K). Re-view to (dim, 1, K) so the conv kernel and the
# default weight loader agree on shape (same trick as MambaMixer2).
self.conv1d.weight.data = self.conv1d.weight.data.unsqueeze(1)
# x_proj: intermediate -> [dt_rank, B(state), C(state)]. Input dim is
# sharded across TP, so this reduces (RowParallel) to full dt/B/C.
self.x_proj = RowParallelLinear(
input_size=intermediate_size,
output_size=time_step_rank + 2 * state_size,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.x_proj",
)
# dt_proj: dt_rank -> intermediate (column-sharded). dt_rank input is
# replicated (small), so keep the input unsharded.
self.dt_proj = ColumnParallelLinear(
input_size=time_step_rank,
output_size=intermediate_size,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.dt_proj",
)
# Full-rank A (stored as A_log) and D, sharded along the intermediate dim.
self.A_log = nn.Parameter(
torch.empty(self.intermediate_size_per_tp, state_size, dtype=torch.float32)
)
self.D = nn.Parameter(torch.ones(self.intermediate_size_per_tp))
set_weight_attrs(self.A_log, {"weight_loader": sharded_weight_loader(0)})
set_weight_attrs(self.D, {"weight_loader": sharded_weight_loader(0)})
# The time-step bias is folded into dt_proj.bias (applied before the
# scan), so the selective_state_update kernel gets a zero dt_bias. Keep
# it as a registered buffer of shape (nheads=dim, head_dim=1); passing a
# real tensor also avoids a `dt_bias is None` unpack path in the kernel.
self.register_buffer(
"dt_bias_zero",
torch.zeros(self.intermediate_size_per_tp, 1),
persistent=False,
)
# out_proj: intermediate -> hidden (input sharded, RowParallel reduces).
self.out_proj = RowParallelLinear(
input_size=intermediate_size,
output_size=hidden_size,
bias=use_bias,
quant_config=quant_config,
prefix=f"{prefix}.out_proj",
)
def _ssm_params(
self, conv_out: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""From convolved x (tokens, dim_per_tp) -> (dt_per_tp, B, C).
``dt`` is per-TP-channel (intermediate/tp); ``B``/``C`` are the full,
replicated state selection vectors. Falcon-Mamba RMS-normalizes all three.
"""
ssm_params, _ = self.x_proj(conv_out)
time_step, B, C = torch.split(
ssm_params,
[self.time_step_rank, self.ssm_state_size, self.ssm_state_size],
dim=-1,
)
if self.use_bc_dt_rms:
B = rms_normalize(B, self.rms_eps)
C = rms_normalize(C, self.rms_eps)
time_step = rms_normalize(time_step, self.rms_eps)
dt, _ = self.dt_proj(time_step) # (tokens, intermediate/tp)
return dt, B, C
def forward(
self,
*,
hidden_states: torch.Tensor,
output: Optional[torch.Tensor],
layer_cache,
metadata,
mup_vector: Optional[torch.Tensor] = None,
use_triton_causal_conv: bool = False,
) -> Tuple[torch.Tensor, None, None]:
# Matches Mamba2AttnBackend.forward's (out, intermediate_states, track_states)
# contract; Mamba-1 has neither extra state, so both are None.
assert not metadata.is_target_verify, (
"Mamba-1 (Falcon-Mamba) does not support speculative decoding yet"
)
# The per-token scan writes no radix track states, so a tracked prefix
# would read back unwritten; require --disable-radix-cache.
assert not metadata.has_mamba_track_mask, (
"Mamba-1 (Falcon-Mamba) does not support radix mamba-state tracking; "
"serve with --disable-radix-cache"
)
conv_state = layer_cache.conv[0]
ssm_state = layer_cache.temporal # (slots, intermediate/tp, 1, state)
state_indices = metadata.mamba_cache_indices
query_start_loc = metadata.query_start_loc
dim = self.intermediate_size_per_tp
num_prefills = metadata.num_prefills
num_prefill_tokens = metadata.num_prefill_tokens
num_decodes = metadata.num_decodes
num_actual_tokens = num_prefill_tokens + num_decodes
# Project and split into x (to be convolved+scanned) and gate.
projected, _ = self.in_proj(hidden_states)
if mup_vector is not None:
projected = projected * mup_vector
x_in, gate = projected.split([dim, dim], dim=-1)
x_in = x_in[:num_actual_tokens]
gate = gate[:num_actual_tokens]
conv_weights = self.conv1d.weight.view(
self.conv1d.weight.size(0), self.conv1d.weight.size(2)
)
# Split varlen tokens into prefill (front) then decode (back).
x_p, x_d = torch.split(x_in, [num_prefill_tokens, num_decodes], dim=0)
gate_p, gate_d = torch.split(gate, [num_prefill_tokens, num_decodes], dim=0)
state_indices_p = state_indices[:num_prefills]
state_indices_d = state_indices[num_prefills : num_prefills + num_decodes]
out = torch.empty(
(num_actual_tokens, dim), dtype=hidden_states.dtype, device=x_in.device
)
out_p, out_d = torch.split(out, [num_prefill_tokens, num_decodes], dim=0)
A = -torch.exp(self.A_log.float()) # (dim, state)
if num_prefills > 0:
self._forward_prefill(
x=x_p,
gate=gate_p,
out=out_p,
A=A,
conv_state=conv_state,
ssm_state=ssm_state,
conv_weights=conv_weights,
state_indices=state_indices_p,
query_start_loc=query_start_loc[: num_prefills + 1],
metadata=metadata,
use_triton_causal_conv=use_triton_causal_conv,
)
if num_decodes > 0:
self._forward_decode(
x=x_d,
gate=gate_d,
out=out_d,
A=A,
conv_state=conv_state,
ssm_state=ssm_state,
conv_weights=conv_weights,
state_indices=state_indices_d,
use_triton_causal_conv=use_triton_causal_conv,
)
mixer_out, _ = self.out_proj(out)
if output is not None:
output[:num_actual_tokens].copy_(mixer_out)
return mixer_out, None, None
def _forward_prefill(
self,
*,
x,
gate,
out,
A,
conv_state,
ssm_state,
conv_weights,
state_indices,
query_start_loc,
metadata,
use_triton_causal_conv,
):
mixed = metadata.mixed_metadata
has_initial = mixed.has_initial_states if mixed is not None else None
# Per-sequence prefill token counts; required by the Triton causal-conv
# varlen kernel (used on XPU). Fall back to deriving from query_start_loc.
seq_lens_cpu = mixed.extend_seq_lens_cpu if mixed is not None else None
if seq_lens_cpu is None:
seq_lens_cpu = (query_start_loc[1:] - query_start_loc[:-1]).cpu().tolist()
# The causal-conv kernel needs input, weights and the conv-state cache in
# one dtype. The cache dtype (SGLANG_MAMBA_CONV_DTYPE) is independent of
# the model dtype, so cast the conv inputs to it and the result back.
act_dtype = x.dtype
conv_dtype = conv_state.dtype
ccfn = causal_conv1d_fn_triton if use_triton_causal_conv else causal_conv1d_fn
conv_out = (
ccfn(
x.transpose(0, 1).to(conv_dtype), # (dim, tokens)
conv_weights.to(conv_dtype),
(
self.conv1d.bias.to(conv_dtype)
if self.conv1d.bias is not None
else None
),
activation=self.activation,
conv_states=conv_state,
has_initial_state=has_initial,
cache_indices=state_indices,
query_start_loc=query_start_loc,
seq_lens_cpu=seq_lens_cpu,
)
.transpose(0, 1)[: x.shape[0]]
.to(act_dtype)
) # (tokens, dim)
dt, B, C = self._ssm_params(conv_out)
# Sequential selective scan per sequence (portable, device-agnostic).
seq_lens = (query_start_loc[1:] - query_start_loc[:-1]).tolist()
for i, seqlen in enumerate(seq_lens):
start = int(query_start_loc[i])
end = start + seqlen
slot = int(state_indices[i])
if has_initial is not None and bool(has_initial[i]):
h = ssm_state[slot, :, 0, :].float() # (dim, state)
else:
h = torch.zeros(
A.shape[0], A.shape[1], dtype=torch.float32, device=x.device
)
# softplus matches HF `discrete_time_step = softplus(dt_proj(time_step))`
# (the decode kernel applies this internally via dt_softplus=True).
dt_seq = nn.functional.softplus(dt[start:end].float())
h, y = self._selective_scan(
x=conv_out[start:end].float(), # (seqlen, dim)
dt=dt_seq, # (seqlen, dim)
A=A, # (dim, state)
B=B[start:end].float(), # (seqlen, state)
C=C[start:end].float(), # (seqlen, state)
h0=h,
)
y = y + conv_out[start:end].float() * self.D.float()[None, :]
y = y * nn.functional.silu(gate[start:end].float())
out[start:end].copy_(y.to(out.dtype))
# Persist the final recurrent state for subsequent decode.
ssm_state[slot, :, 0, :].copy_(h.to(ssm_state.dtype))
@staticmethod
def _selective_scan(*, x, dt, A, B, C, h0):
"""Reference Mamba-1 recurrence for one sequence.
Shapes: x/dt (T, dim); A (dim, state); B/C (T, state); h0 (dim, state).
Returns (final_state (dim, state), y (T, dim)).
Discretization is computed *per timestep* rather than materializing the
full (T, dim, state) tensors up front: with dim==intermediate_size (8192)
and prefill chunks up to 2048 tokens, a materialized (T, dim, state) is
~1 GB in fp32 and OOMs the XPU under concurrent prefill. The per-step
form bounds peak activation to O(dim * state).
"""
h = h0
ys = []
for t in range(x.shape[0]):
# dA = exp(dt * A), dBx = dt * B * x (all (dim, state) for this step).
dt_t = dt[t][:, None] # (dim, 1)
dA_t = torch.exp(dt_t * A) # (dim, state)
dBx_t = (dt_t * B[t][None, :]) * x[t][:, None] # (dim, state)
h = dA_t * h + dBx_t
ys.append((h * C[t][None, :]).sum(-1)) # (dim,)
y = torch.stack(ys, dim=0) # (T, dim)
return h, y
def _forward_decode(
self,
*,
x,
gate,
out,
A,
conv_state,
ssm_state,
conv_weights,
state_indices,
use_triton_causal_conv,
):
# Match the conv-state cache dtype (see _forward_prefill), then cast the
# result back to the activation dtype before the x_proj matmul.
act_dtype = x.dtype
conv_dtype = conv_state.dtype
ccu = (
causal_conv1d_update_triton
if use_triton_causal_conv
else causal_conv1d_update
)
conv_out = ccu(
x.to(conv_dtype),
conv_state,
conv_weights.to(conv_dtype),
self.conv1d.bias.to(conv_dtype) if self.conv1d.bias is not None else None,
self.activation,
conv_state_indices=state_indices,
).to(act_dtype)
dt, B, C = self._ssm_params(conv_out)
# Map onto the shared selective_state_update kernel with the
# (nheads=dim, head_dim=1, ngroups=1) full-rank layout. z=gate applies
# the silu output gate; D is the per-channel skip connection.
n_decode = x.shape[0]
dim = self.intermediate_size_per_tp
A_k = A[:, None, :] # (dim, 1, state)
D_k = self.D.float()[:, None] # (dim, 1)
x_k = conv_out.view(n_decode, dim, 1)
dt_k = dt.view(n_decode, dim, 1)
gate_k = gate.view(n_decode, dim, 1)
B_k = B.view(n_decode, 1, self.ssm_state_size)
C_k = C.view(n_decode, 1, self.ssm_state_size)
out_k = out.view(n_decode, dim, 1)
selective_state_update(
ssm_state,
x_k,
dt_k,
A_k,
B_k,
C_k,
D_k,
z=gate_k,
dt_bias=self.dt_bias_zero,
dt_softplus=True,
state_batch_indices=state_indices,
out=out_k,
)
@property
def mamba_type(self) -> str:
return "mamba1"
@@ -68,11 +68,31 @@ def configure_aux_hidden_state_capture(
def build_attention_backends(*, model_runner: ModelRunner) -> AttentionBackends:
"""Init attention kernel backend."""
from sglang.srt.configs.model_config import AttentionArch
# TODO: Refactor device-specific init branches into platform interface (separate PR).
# Must run before the SSM early-return below; Mamba mixers still issue GEMMs.
if model_runner.device in ("cuda", "musa"):
init_cublas()
# SSM models use the Mamba backend, not attention. Import inside the branch so
# non-SSM models don't load the Mamba-specific backend deps.
if model_runner.model_config.attention_arch == AttentionArch.SSM:
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
Mamba2AttnBackend,
)
mamba_backend = Mamba2AttnBackend(model_runner)
return AttentionBackends(
attn_backend=mamba_backend,
decode_attn_backend=None,
decode_attn_backend_group=[],
prefill_attention_backend_str="mamba2",
decode_attention_backend_str="mamba2",
)
server_args = model_runner.server_args
# Already resolved and stamped on the runner before this call.
resolved = ResolvedAttentionBackendStr(
prefill=model_runner.prefill_attention_backend_str,
+252
View File
@@ -0,0 +1,252 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright 2023-2025 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.
# ==============================================================================
"""Inference-only Falcon-Mamba model (tiiuae/falcon-mamba-7b) for SGLang.
Falcon-Mamba is a pure **Mamba-1** (selective-scan) state-space model: each
decoder layer is a pre-norm followed by a Mamba-1 mixer, with no MLP sublayer
and no attention. It differs from Mamba-2 (see ``models/mamba2.py``) in the SSM
math (full-rank ``A``, low-rank ``dt`` via ``x_proj``/``dt_proj``, conv over the
intermediate channels only) and adds a weightless RMSNorm on ``B``/``C``/``dt``.
The Mamba-1 mixer (``MambaMixer1``) rides the shared Mamba2 attention backend,
memory pool and kernels; see ``layers/attention/mamba/mamba1.py``.
Reference: https://huggingface.co/tiiuae/falcon-mamba-7b
"""
import logging
from typing import Iterable, Optional, Set, Tuple
import torch
from torch import nn
from sglang.srt.layers.attention.mamba.mamba1 import MambaMixer1
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.vocab_parallel_embedding import (
ParallelLMHead,
VocabParallelEmbedding,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__)
class FalconMambaDecoderLayer(nn.Module):
"""Falcon-Mamba decoder layer: pre-norm + Mamba-1 mixer (no MLP)."""
def __init__(
self,
config,
layer_id: int,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
use_bc_dt_rms: bool = True,
):
super().__init__()
self.layer_id = layer_id
# Pre-normalization (checkpoint key: backbone.layers.N.norm)
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
self.mixer = MambaMixer1(
hidden_size=config.hidden_size,
intermediate_size=config.intermediate_size,
state_size=config.state_size,
conv_kernel=config.conv_kernel,
time_step_rank=config.time_step_rank,
use_conv_bias=config.use_conv_bias,
use_bias=config.use_bias,
activation=config.hidden_act,
# Falcon-Mamba stabilizes with a weightless RMSNorm on B/C/dt; plain
# Mamba (MambaForCausalLM) sets this False.
use_bc_dt_rms=use_bc_dt_rms,
rms_eps=config.mixer_rms_eps if use_bc_dt_rms else 1e-6,
quant_config=quant_config,
prefix=add_prefix("mixer", prefix),
)
def forward(
self,
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
residual: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
# Pre-norm (with fused residual add) -> Mamba-1 mixer.
if residual is None:
residual = hidden_states
hidden_states = self.norm(hidden_states)
else:
hidden_states, residual = self.norm(hidden_states, residual)
# Run the mixer through the Mamba2 attention backend (owns the conv/ssm
# state cache). use_triton_causal_conv keeps the conv portable off-CUDA.
attn_backend = get_attn_backend()
output = torch.empty_like(hidden_states)
attn_backend.forward(
self.mixer,
hidden_states,
output,
layer_id=self.layer_id,
forward_batch=forward_batch,
use_triton_causal_conv=True,
)
return output, residual
class FalconMambaModel(nn.Module):
"""Falcon-Mamba backbone (no LM head)."""
def __init__(
self,
config,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
use_bc_dt_rms: bool = True,
):
super().__init__()
self.config = config
self.padding_idx = config.pad_token_id
self.vocab_size = config.vocab_size
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
prefix=add_prefix("embed_tokens", prefix),
)
self.layers = make_layers(
config.num_hidden_layers,
lambda idx, prefix: FalconMambaDecoderLayer(
config=config,
layer_id=idx,
quant_config=quant_config,
prefix=prefix,
use_bc_dt_rms=use_bc_dt_rms,
),
prefix=add_prefix("layers", prefix),
)
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
def forward(
self,
input_ids: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
hidden_states = self.embed_tokens(input_ids)
residual = None
for layer in self.layers:
hidden_states, residual = layer(hidden_states, forward_batch, residual)
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
class FalconMambaForCausalLM(nn.Module):
"""Falcon-Mamba (Mamba-1) model with a language modeling head.
Also serves as the base for plain Mamba (models/mamba.py), which only flips
``use_bc_dt_rms`` off; the tied vs untied LM head follows
``config.tie_word_embeddings``.
"""
# Falcon-Mamba applies the weightless B/C/dt RMSNorm; plain Mamba sets False.
use_bc_dt_rms: bool = True
def __init__(
self,
config,
quant_config: Optional[QuantizationConfig] = None,
cache_config=None,
):
super().__init__()
self.config = config
self.quant_config = quant_config
self.model = FalconMambaModel(
config=config,
quant_config=quant_config,
prefix="model",
use_bc_dt_rms=self.use_bc_dt_rms,
)
self.lm_head = ParallelLMHead(
config.vocab_size,
config.hidden_size,
prefix="lm_head",
)
# Tie to the input embeddings when the checkpoint has no separate lm_head
# (state-spaces Mamba); Falcon-Mamba is untied (tie_word_embeddings=False).
if config.tie_word_embeddings:
self.lm_head = self.lm_head.tie_weights(self.model.embed_tokens)
self.logits_processor = LogitsProcessor(config)
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
hidden_states = self.model(input_ids, forward_batch)
return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch
)
def get_input_embeddings(self) -> nn.Module:
return self.model.embed_tokens
def get_output_embeddings(self) -> nn.Module:
return self.lm_head
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
for name, loaded_weight in weights:
# Remap checkpoint names to SGLang modules: backbone.* -> model.*,
# embeddings./embedding. (plural -hf / singular raw state-spaces) ->
# embed_tokens., norm_f -> norm. Keep A_log as-is (the mixer computes
# A = -exp(A_log)).
if name.startswith("backbone."):
name = "model." + name[len("backbone.") :]
name = name.replace("embeddings.", "embed_tokens.")
name = name.replace("embedding.", "embed_tokens.")
name = name.replace("norm_f.", "norm.")
if name not in params_dict:
logger.warning(f"Skipping parameter {name} - not found in model")
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
unloaded_params = set(params_dict.keys()) - loaded_params
# A tied lm_head is legitimately absent from the checkpoint.
unloaded_params = {p for p in unloaded_params if not p.startswith("lm_head")}
if unloaded_params:
logger.warning(
f"The following parameters were not loaded: {unloaded_params}"
)
return loaded_params
EntryClass = FalconMambaForCausalLM
+33
View File
@@ -0,0 +1,33 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright 2023-2025 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.
# ==============================================================================
"""Inference-only Mamba model (state-spaces/mamba-*) for SGLang.
The canonical Mamba-1 (selective-scan) state-space model. It is Falcon-Mamba
without the weightless B/C/dt RMSNorm, so it reuses FalconMambaForCausalLM and
only flips ``use_bc_dt_rms`` off; the tied LM head follows
``config.tie_word_embeddings``.
Reference: https://huggingface.co/state-spaces/mamba-130m-hf
"""
from sglang.srt.models.falcon_mamba import FalconMambaForCausalLM
class MambaForCausalLM(FalconMambaForCausalLM):
# Plain Mamba has no B/C/dt RMSNorm (that is Falcon-Mamba's variant).
use_bc_dt_rms: bool = False
EntryClass = MambaForCausalLM
+262
View File
@@ -0,0 +1,262 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright 2023-2025 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.
# ==============================================================================
"""Inference-only Mamba2 model (e.g. Mamba-Codestral-7B) for SGLang.
Mamba2 is a pure state-space model (SSM) architecture that replaces attention
with selective state spaces. Each decoder layer is a pre-norm followed by a
Mamba2 mixer; there is no MLP sublayer.
Reference: https://huggingface.co/mistralai/Mamba-Codestral-7B-v0.1
"""
import logging
from typing import Iterable, Optional, Set, Tuple
import torch
from torch import nn
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.layers.attention.mamba.mamba import MambaMixer2
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.vocab_parallel_embedding import (
ParallelLMHead,
VocabParallelEmbedding,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__)
class Mamba2DecoderLayer(nn.Module):
"""
Mamba2 decoder layer with SSM mixer instead of attention.
"""
def __init__(
self,
config,
layer_id: int,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.layer_id = layer_id
self.hidden_size = config.hidden_size
# Pre-normalization (checkpoint key: backbone.layers.N.norm)
self.norm = RMSNorm(
config.hidden_size,
eps=config.layer_norm_epsilon,
)
state_shape = Mamba2StateShape.create(
tp_world_size=get_parallel().tp_size,
intermediate_size=config.intermediate_size,
n_groups=config.n_groups,
num_heads=config.num_heads,
head_dim=config.head_dim,
state_size=config.state_size,
conv_kernel=config.conv_kernel,
)
cache_params = Mamba2CacheParams(
shape=state_shape,
layers=list(range(config.num_hidden_layers)),
)
self.mixer = MambaMixer2(
cache_params=cache_params,
hidden_size=config.hidden_size,
use_conv_bias=config.use_conv_bias,
use_bias=config.use_bias,
n_groups=config.n_groups,
rms_norm_eps=config.layer_norm_epsilon,
activation=config.hidden_act,
use_rms_norm=config.rms_norm,
quant_config=quant_config,
prefix=add_prefix("mixer", prefix),
)
def forward(
self,
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
residual: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
# Pre-norm (fused residual add) -> SSM mixer; no MLP sublayer.
if residual is None:
residual = hidden_states
hidden_states = self.norm(hidden_states)
else:
hidden_states, residual = self.norm(hidden_states, residual)
# Run the mixer through the Mamba2 attention backend (owns the conv/ssm
# state cache). use_triton_causal_conv keeps the conv portable off-CUDA.
attn_backend = get_attn_backend()
output = torch.empty_like(hidden_states)
attn_backend.forward(
self.mixer,
hidden_states,
output,
layer_id=self.layer_id,
forward_batch=forward_batch,
use_triton_causal_conv=True,
)
hidden_states = output
return hidden_states, residual
class Mamba2Model(nn.Module):
"""Mamba2 model without the language modeling head."""
def __init__(
self,
config,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.config = config
self.padding_idx = config.pad_token_id
self.vocab_size = config.vocab_size
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
prefix=add_prefix("embed_tokens", prefix),
)
self.layers = make_layers(
config.num_hidden_layers,
lambda idx, prefix: Mamba2DecoderLayer(
config=config,
layer_id=idx,
quant_config=quant_config,
prefix=prefix,
),
prefix=add_prefix("layers", prefix),
)
self.norm = RMSNorm(
config.hidden_size,
eps=config.layer_norm_epsilon,
)
def forward(
self,
input_ids: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
hidden_states = self.embed_tokens(input_ids)
residual = None
for layer in self.layers:
hidden_states, residual = layer(hidden_states, forward_batch, residual)
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
class Mamba2ForCausalLM(nn.Module):
"""
Mamba2 model with a language modeling head.
This is the main model class for Mamba2-based models like Mamba-Codestral-7B.
"""
def __init__(
self,
config,
quant_config: Optional[QuantizationConfig] = None,
cache_config=None,
):
super().__init__()
self.config = config
self.quant_config = quant_config
self.model = Mamba2Model(
config=config,
quant_config=quant_config,
prefix="model",
)
self.lm_head = ParallelLMHead(
config.vocab_size,
config.hidden_size,
prefix="lm_head",
)
self.logits_processor = LogitsProcessor(config)
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
hidden_states = self.model(input_ids, forward_batch)
return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch
)
def get_input_embeddings(self) -> nn.Module:
return self.model.embed_tokens
def get_output_embeddings(self) -> nn.Module:
return self.lm_head
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
for name, loaded_weight in weights:
# Skip rotary embedding inverse frequencies
if "rotary_emb.inv_freq" in name or "inv_freq" in name:
continue
# Remap checkpoint names to SGLang modules: backbone.* -> model.*,
# embeddings -> embed_tokens, norm_f -> norm, A_log -> A.
if name.startswith("backbone."):
name = "model." + name[len("backbone.") :]
name = name.replace("embeddings.", "embed_tokens.")
name = name.replace("norm_f.", "norm.")
if "A_log" in name:
name = name.replace("A_log", "A")
if name not in params_dict:
logger.warning(f"Skipping parameter {name} - not found in model")
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
unloaded_params = set(params_dict.keys()) - loaded_params
if unloaded_params:
logger.warning(
f"The following parameters were not loaded: {unloaded_params}"
)
return loaded_params
EntryClass = Mamba2ForCausalLM
@@ -37,6 +37,7 @@ from sglang.srt.configs import (
DotsVLMConfig,
ExaoneConfig,
FalconH1Config,
FalconMambaConfig,
Glm5NextConfig,
Glm5NextTextConfig,
GraniteMoeHybridConfig,
@@ -58,6 +59,8 @@ from sglang.srt.configs import (
LagunaConfig,
LocateAnythingConfig,
LongcatFlashConfig,
Mamba2Config,
MambaConfig,
MiniCPMHybridConfig,
MiniCPMV4_6Config,
MiniCPMV4_6VisionConfig,
@@ -132,6 +135,9 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
Qwen4ExpConfig,
Qwen4ExpTextConfig,
FalconH1Config,
FalconMambaConfig,
Mamba2Config,
MambaConfig,
GraniteMoeHybridConfig,
HYV4Config,
DotsVLMConfig,
@@ -74,6 +74,47 @@ def _try_load_longcat_config(model, revision: Optional[str], **kwargs):
)
def _try_load_raw_mamba_config(model, revision: Optional[str], **kwargs):
"""Recognize the original state-spaces Mamba-1 checkpoints.
The raw `state-spaces/mamba-*` repos (e.g. mamba-130m/790m/2.8b, as opposed
to the `-hf` conversions) ship a minimal `config.json` with `d_model` /
`n_layer` / `ssm_cfg` and NO `model_type` / `architectures`, so
`AutoConfig.from_pretrained` rejects them with "Unrecognized model ...".
Detect that shape and build our `MambaConfig` (model_type `mamba`, arch
`MambaForCausalLM`) with the field-name mapping the SGLang Mamba model
expects. Uses `get_config_dict` (which does not require a model_type) so
this runs before the failing `AutoConfig` path.
"""
config_dict, _ = PretrainedConfig.get_config_dict(
model, revision=revision, **kwargs
)
# Raw state-spaces Mamba: has d_model + ssm_cfg, and no model_type/arch.
if config_dict.get("model_type") or config_dict.get("architectures"):
return None
if "d_model" not in config_dict or "ssm_cfg" not in config_dict:
return None
from sglang.srt.configs.mamba import MambaConfig
d_model = config_dict["d_model"]
# The embedding is padded up to a multiple of pad_vocab_size_multiple; match
# the checkpoint (e.g. 50277 -> 50280) so weight shapes line up.
pad = config_dict.get("pad_vocab_size_multiple", 1)
vocab_size = config_dict.get("vocab_size", 50280)
if pad > 1:
vocab_size = ((vocab_size + pad - 1) // pad) * pad
return MambaConfig(
vocab_size=vocab_size,
hidden_size=d_model,
num_hidden_layers=config_dict["n_layer"],
state_size=config_dict.get("ssm_cfg", {}).get("d_state", 16),
layer_norm_epsilon=config_dict.get("layer_norm_epsilon", 1e-5),
residual_in_fp32=config_dict.get("residual_in_fp32", True),
architectures=["MambaForCausalLM"],
)
@register_model_config_parser("hf")
class HfModelConfigParser(ModelConfigParserBase):
def parse(
@@ -84,6 +125,8 @@ class HfModelConfigParser(ModelConfigParserBase):
**kwargs,
):
config = _try_load_longcat_config(model, revision, **kwargs)
if config is None:
config = _try_load_raw_mamba_config(model, revision, **kwargs)
if config is None:
config = AutoConfig.from_pretrained(
model,
@@ -140,9 +183,17 @@ class HfModelConfigParser(ModelConfigParserBase):
model_type = config.model_type
if model_type == "deepseek_vl_v2" and is_ocr:
model_type = "deepseek-ocr"
config = _CONFIG_REGISTRY[model_type].from_pretrained(
model, revision=revision
)
# Raw state-spaces Mamba configs are built by
# _try_load_raw_mamba_config with architectures injected; reloading
# from the checkpoint would drop them, so skip it when the config is
# already one of our classes.
from sglang.srt.configs.mamba import FalconMambaConfig, MambaConfig
from sglang.srt.configs.mamba2 import Mamba2Config
if not isinstance(config, (Mamba2Config, MambaConfig, FalconMambaConfig)):
config = _CONFIG_REGISTRY[model_type].from_pretrained(
model, revision=revision
)
# Re-check after reloading config from registry
if _is_deepseek_ocr_model(config) or _is_deepseek_ocr2_model(config):
@@ -0,0 +1,85 @@
"""
Prefill -> decode runtime test for the Mamba-1 mixer on Intel XPU.
Guards the MambaMixer1 <-> Mamba2AttnBackend contract (3-tuple return) and the
SSM conv/state cache end to end; the CPU weight-remap test cannot catch either.
Uses a real server so the scheduler initializes the mamba selective-scan backend.
Usage:
python3 -m unittest test_xpu_mamba1_runtime.TestXPUMamba1Runtime
"""
import unittest
import requests
from sglang.test.ci.ci_register import register_xpu_ci
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
kill_process_tree,
popen_launch_server,
)
register_xpu_ci(est_time=600, suite="stage-b-test-1-gpu-xpu")
# Small Mamba-1 (state-spaces) checkpoint; exercises MambaMixer1 on XPU.
MODEL = "state-spaces/mamba-130m-hf"
class TestXPUMamba1Runtime(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
MODEL,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
device="xpu",
other_args=[
"--device",
"xpu",
"--attention-backend",
"intel_xpu",
"--disable-radix-cache", # Mamba-1 has no radix track state
"--max-total-tokens",
"65536",
"--mem-fraction-static",
"0.9",
"--trust-remote-code",
],
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def _generate(self, prompt, max_new_tokens=32):
resp = requests.post(
f"{self.base_url}/generate",
json={
"text": prompt,
"sampling_params": {"temperature": 0, "max_new_tokens": max_new_tokens},
},
timeout=120,
)
resp.raise_for_status()
return resp.json()["text"]
def test_prefill_to_decode(self):
# Multi-token prompt forces a real prefill scan; max_new_tokens>1 forces the
# decode-step recurrence. Before the 3-tuple fix this raised
# "not enough values to unpack" on the first forward.
out = self._generate("The capital of France is")
self.assertTrue(out and out.strip(), "empty completion")
def test_greedy_is_deterministic(self):
# Identical greedy requests must match; a corrupted conv/ssm state cache
# across requests would make them diverge.
prompt = "Count: one two three"
self.assertEqual(self._generate(prompt), self._generate(prompt))
if __name__ == "__main__":
unittest.main()
+117
View File
@@ -0,0 +1,117 @@
"""CPU unit test for Mamba2 (Mamba-Codestral) checkpoint weight-name remapping.
Exercises ``Mamba2ForCausalLM.load_weights`` name translation from the
HuggingFace ``backbone.*`` checkpoint layout to SGLang module names, without
constructing the full model or requiring a GPU:
- ``backbone.`` -> ``model.``
- ``embeddings.`` -> ``embed_tokens.``
- ``norm_f.`` -> ``norm.``
- ``...mixer.A_log`` -> ``...mixer.A``
- ``lm_head.weight`` kept as-is
- ``*inv_freq`` entries skipped
Run: python3 test/srt/models/test_mamba2.py
"""
import unittest
import torch
from sglang.srt.models.mamba2 import Mamba2ForCausalLM
def _param(like: torch.Tensor) -> torch.nn.Parameter:
"""A parameter whose weight_loader copies in place (SGLang loader contract)."""
p = torch.nn.Parameter(torch.zeros_like(like), requires_grad=False)
p.weight_loader = lambda param, loaded: param.data.copy_(loaded)
return p
class _FakeMamba2:
"""Minimal stand-in exposing named_parameters() with SGLang-side names."""
def __init__(self, params):
self._params = params
def named_parameters(self):
return list(self._params.items())
# Exercise the real method as an unbound function (no full model build).
load_weights = Mamba2ForCausalLM.load_weights
class TestMamba2WeightRemap(unittest.TestCase):
def test_backbone_names_are_remapped_and_loaded(self):
# SGLang-side parameters (load targets).
sgl = {
"model.embed_tokens.weight": _param(torch.empty(4, 3)),
"model.layers.0.norm.weight": _param(torch.empty(3)),
"model.layers.0.mixer.A": _param(torch.empty(2)),
"model.norm.weight": _param(torch.empty(3)),
"lm_head.weight": _param(torch.empty(4, 3)),
}
model = _FakeMamba2(sgl)
# HuggingFace checkpoint names (sources) with distinct values.
hf = {
"backbone.embeddings.weight": torch.arange(12, dtype=torch.float32).reshape(
4, 3
),
"backbone.layers.0.norm.weight": torch.tensor([1.0, 2.0, 3.0]),
"backbone.layers.0.mixer.A_log": torch.tensor([5.0, 6.0]),
"backbone.norm_f.weight": torch.tensor([7.0, 8.0, 9.0]),
"lm_head.weight": torch.full((4, 3), 4.0),
}
loaded = model.load_weights(list(hf.items()))
# Every source mapped onto exactly its SGLang target.
self.assertEqual(loaded, set(sgl))
torch.testing.assert_close(
sgl["model.embed_tokens.weight"].data, hf["backbone.embeddings.weight"]
)
torch.testing.assert_close(
sgl["model.layers.0.norm.weight"].data,
hf["backbone.layers.0.norm.weight"],
)
# A_log -> A: the raw checkpoint tensor lands in the A parameter.
torch.testing.assert_close(
sgl["model.layers.0.mixer.A"].data, hf["backbone.layers.0.mixer.A_log"]
)
# norm_f -> norm (final norm), distinct from the per-layer norm above.
torch.testing.assert_close(
sgl["model.norm.weight"].data, hf["backbone.norm_f.weight"]
)
torch.testing.assert_close(sgl["lm_head.weight"].data, hf["lm_head.weight"])
def test_inv_freq_entries_are_skipped(self):
sgl = {"model.layers.0.mixer.A": _param(torch.empty(2))}
model = _FakeMamba2(sgl)
hf = [
("backbone.layers.0.mixer.A_log", torch.tensor([1.0, 2.0])),
("backbone.layers.0.mixer.inv_freq", torch.tensor([0.0, 0.0])),
("rotary_emb.inv_freq", torch.tensor([0.0])),
]
loaded = model.load_weights(hf)
# Only the A parameter is loaded; inv_freq sources are ignored.
self.assertEqual(loaded, {"model.layers.0.mixer.A"})
def test_unmatched_source_is_ignored_not_fatal(self):
sgl = {"model.norm.weight": _param(torch.empty(2))}
model = _FakeMamba2(sgl)
hf = [
("backbone.norm_f.weight", torch.tensor([1.0, 2.0])),
("backbone.this.does.not.exist", torch.tensor([9.0])),
]
loaded = model.load_weights(hf)
self.assertEqual(loaded, {"model.norm.weight"})
if __name__ == "__main__":
unittest.main()