Inference Support Mamba 2 and 1 (#34556)
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
)
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user