diff --git a/docs/docs/supported-models/generative_models.mdx b/docs/docs/supported-models/generative_models.mdx index b180966c6..5c7cc94e7 100644 --- a/docs/docs/supported-models/generative_models.mdx +++ b/docs/docs/supported-models/generative_models.mdx @@ -199,9 +199,9 @@ in the GitHub search bar. IBM's open dense foundation models optimized for reasoning, code, and business AI use cases. Integrated with Red Hat and watsonx systems. - **Granite 3.0 MoE** (IBM) - `ibm-granite/granite-3.0-3b-a800m-instruct` - IBM’s Mixture-of-Experts models offering strong performance with cost-efficiency. MoE expert routing designed for enterprise deployment at scale. + **Granite 3.0, 3.1 MoE** (IBM) + `ibm-granite/granite-3.0-3b-a800m-instruct`, `ibm-granite/granite-3.1-3b-a800m-instruct` + IBM’s Mixture-of-Experts models offering strong performance with cost-efficiency. MoE expert routing designed for enterprise deployment at scale. 3.1 adds shared experts. **GPT-J** (6B) @@ -278,6 +278,16 @@ in the GitHub search bar. ibm-granite/granite-4.0-h-micro, ibm-granite/granite-4.0-micro IBM Granite 4.0 micro models: hybrid Mamba–MoE (h-micro) and dense (micro) variants. Enterprise-focused reasoning models + + Granite SWA (IBM) + ibm-granite/granite-swash-2b + IBM Granite dense model with per-layer sliding window attention and learnable per-head attention sinks. Served by the Granite implementation; needs transformers ≥ 5.15. + + + Granite MoE SWA (IBM) + ibm-granite/granite-swash-3b-a600m + IBM Granite Mixture-of-Experts model with per-layer sliding window attention, learnable attention sinks, and optional shared experts. Served by the GraniteMoE implementation; needs transformers ≥ 5.15. + Sarvam 2 (30B-A2B, 105B-A10B) sarvamai/sarvam-2 diff --git a/python/sglang/srt/configs/granitemoehybrid.py b/python/sglang/srt/configs/granitemoehybrid.py index 205666b36..e34dec2a1 100644 --- a/python/sglang/srt/configs/granitemoehybrid.py +++ b/python/sglang/srt/configs/granitemoehybrid.py @@ -26,6 +26,12 @@ MAMBA = "mamba" ATTENTION = "attention" +def _to_sglang_layer_types(layer_types: list[str]) -> list[str]: + # transformers >= 5.15 uses "linear_attention" / "full_attention" + aliases = {"linear_attention": MAMBA, "full_attention": ATTENTION} + return [aliases.get(t, t) for t in layer_types] + + class GraniteMoeHybridConfig(PretrainedConfig): r""" This is the configuration class to store the configuration of a [`GraniteMoeHybridModel`]. It is used to instantiate a @@ -186,7 +192,7 @@ class GraniteMoeHybridConfig(PretrainedConfig): else: self.layer_types.append(MAMBA) else: - self.layer_types = layer_types + self.layer_types = _to_sglang_layer_types(layer_types) # Validate layer_types if len(self.layer_types) != self.num_hidden_layers: @@ -266,6 +272,10 @@ class GraniteMoeHybridConfig(PretrainedConfig): **kwargs, ) + # HF's `PreTrainedConfig.validate_layer_type` rewrites layer_types to + # transformers >= 5.15 schema, so map them back afterwards. + self.layer_types = _to_sglang_layer_types(self.layer_types) + @property def mamba_layer_ids(self): """Returns the indices of layers that are Mamba layers.""" diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 31426267a..fa8da9951 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -49,6 +49,14 @@ MIMO_V2_MODEL_ARCHS = ( ) MIMO_V2_MULTIMODAL_ARCHS = ("MiMoV2ForCausalLM",) +SWA_SINK_ARCHS = frozenset( + { + "GptOssForCausalLM", + "GraniteSWAForCausalLM", + "GraniteMoeSWAForCausalLM", + } +) + def get_mimo_v2_fused_qkv_expected_tp_size(hf_config): layout = getattr(hf_config, "attention_projection_layout", None) @@ -813,8 +821,7 @@ class ModelConfig: attention. Not every hybrid-SWA model uses them. """ archs = self.hf_config.architectures or [] - # GptOss always creates sinks unconditionally. - if "GptOssForCausalLM" in archs: + if any(a in SWA_SINK_ARCHS for a in archs): return True # MiMoV2 creates sinks only when the config flags are set. @@ -2066,7 +2073,7 @@ def is_hybrid_swa_model( "DeepseekV4ForCausalLM", "DeepseekV4ForCausalLMNextN", "DeepseekV4ForCausalLMDSpark", - "GptOssForCausalLM", + *SWA_SINK_ARCHS, *MIMO_V2_MODEL_ARCHS, "MiMoV2MTP", "Step3p5ForCausalLM", @@ -2111,7 +2118,7 @@ def get_hybrid_layer_ids( full_attention_layer_ids = [ i for i in range(num_hidden_layers) if (i + 1) % 4 == 0 ] - elif "GptOssForCausalLM" in model_architectures: + elif any(arch in SWA_SINK_ARCHS for arch in model_architectures): layer_types = getattr(hf_text_config, "layer_types", []) swa_attention_layer_ids = [ i for i, x in enumerate(layer_types) if x == "sliding_attention" diff --git a/python/sglang/srt/models/granite.py b/python/sglang/srt/models/granite.py index 21fb338cb..1fcb41acd 100644 --- a/python/sglang/srt/models/granite.py +++ b/python/sglang/srt/models/granite.py @@ -24,7 +24,7 @@ from typing import Any, Dict, Iterable, Optional, Tuple import torch from torch import nn -from transformers import GraniteConfig +from transformers import GraniteConfig, PretrainedConfig from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.layernorm import RMSNorm @@ -43,13 +43,44 @@ from sglang.srt.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.model_loader.weight_utils import default_weight_loader +from sglang.srt.model_loader.weight_utils import ( + default_weight_loader, + sharded_weight_loader, +) from sglang.srt.runtime_context import get_parallel -from sglang.srt.utils import add_prefix +from sglang.srt.server_args import get_global_server_args +from sglang.srt.utils import add_prefix, set_weight_attrs from sglang.utils import get_exception_traceback logger = logging.getLogger(__name__) +SWA_MODEL_TYPES = frozenset({"granite_swa", "granitemoe_swa"}) + + +def granite_layer_attn_params( + config: PretrainedConfig, layer_id: int +) -> Tuple[int, float, bool]: + """Extract flags for sliding window, rope theta, attention sink.""" + if config.model_type not in SWA_MODEL_TYPES: + return -1, config.rope_parameters["rope_theta"], False + + # SGLang's window bound is exclusive, hence `- 1` (matching gpt_oss). + sliding_window_size = ( + config.sliding_window - 1 + if config.layer_types[layer_id] == "sliding_attention" + else -1 + ) + return sliding_window_size, config.layer_rope_theta[layer_id], True + + +def build_attention_sinks(num_heads: int) -> nn.Parameter: + # trtllm_mha requires float32 sinks, other backends use bfloat16. + attn_backend = get_global_server_args().attention_backend + sinks_dtype = torch.float32 if attn_backend == "trtllm_mha" else torch.bfloat16 + sinks = nn.Parameter(torch.empty(num_heads, dtype=sinks_dtype), requires_grad=False) + set_weight_attrs(sinks, {"weight_loader": sharded_weight_loader(0)}) + return sinks + class GraniteMLP(nn.Module): def __init__( @@ -97,7 +128,6 @@ class GraniteAttention(nn.Module): num_heads: int, num_kv_heads: int, layer_id: int = 0, - rope_theta: float = 10000, rope_scaling: Optional[Dict[str, Any]] = None, rope_is_neox_style: bool = True, max_position_embeddings: int = 8192, @@ -127,7 +157,9 @@ class GraniteAttention(nn.Module): self.q_size = self.num_heads * self.head_dim self.kv_size = self.num_kv_heads * self.head_dim self.scaling = config.attention_multiplier - self.rope_theta = rope_theta + sliding_window_size, self.rope_theta, has_sink = granite_layer_attn_params( + config, layer_id + ) self.max_position_embeddings = max_position_embeddings self.qkv_proj = QKVParallelLinear( @@ -147,14 +179,19 @@ class GraniteAttention(nn.Module): prefix=add_prefix("o_proj", prefix), ) - self.rotary_emb = get_rope( - self.head_dim, - rotary_dim=self.head_dim, - max_position=max_position_embeddings, - base=rope_theta, - rope_scaling=rope_scaling, - is_neox_style=rope_is_neox_style, + self.rotary_emb = ( + get_rope( + self.head_dim, + rotary_dim=self.head_dim, + max_position=max_position_embeddings, + base=self.rope_theta, + rope_scaling=rope_scaling, + is_neox_style=rope_is_neox_style, + ) + if self.rope_theta + else None ) + self.sinks = build_attention_sinks(self.num_heads) if has_sink else None self.attn = RadixAttention( self.num_heads, self.head_dim, @@ -163,6 +200,7 @@ class GraniteAttention(nn.Module): layer_id=layer_id, quant_config=quant_config, prefix=add_prefix("attn", prefix), + sliding_window_size=sliding_window_size, ) def forward( @@ -173,8 +211,12 @@ class GraniteAttention(nn.Module): ) -> torch.Tensor: qkv, _ = self.qkv_proj(hidden_states) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - q, k = self.rotary_emb(positions, q, k) - attn_output = self.attn(q, k, v, forward_batch) + if self.rotary_emb is not None: + q, k = self.rotary_emb(positions, q, k) + if self.sinks is None: + attn_output = self.attn(q, k, v, forward_batch) + else: + attn_output = self.attn(q, k, v, forward_batch, sinks=self.sinks) output, _ = self.o_proj(attn_output) return output @@ -190,7 +232,6 @@ class GraniteDecoderLayer(nn.Module): super().__init__() self.hidden_size = config.hidden_size self.residual_multiplier = config.residual_multiplier - rope_theta = config.rope_parameters["rope_theta"] rope_scaling = config.rope_parameters if rope_scaling is not None and getattr( config, "original_max_position_embeddings", None @@ -206,7 +247,6 @@ class GraniteDecoderLayer(nn.Module): num_heads=config.num_attention_heads, num_kv_heads=config.num_key_value_heads, layer_id=layer_id, - rope_theta=rope_theta, rope_scaling=rope_scaling, rope_is_neox_style=rope_is_neox_style, max_position_embeddings=max_position_embeddings, @@ -505,4 +545,8 @@ class GraniteForCausalLM(nn.Module): return None -EntryClass = [GraniteForCausalLM] +class GraniteSWAForCausalLM(GraniteForCausalLM): + pass + + +EntryClass = [GraniteForCausalLM, GraniteSWAForCausalLM] diff --git a/python/sglang/srt/models/granitemoe.py b/python/sglang/srt/models/granitemoe.py index 822c335dd..17989711b 100644 --- a/python/sglang/srt/models/granitemoe.py +++ b/python/sglang/srt/models/granitemoe.py @@ -1,13 +1,15 @@ """Inference-only GraniteMoe model.""" -from typing import Iterable, Optional +from typing import Iterable, Iterator, Optional import torch from torch import nn from transformers import GraniteConfig +from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( + MergedColumnParallelLinear, QKVParallelLinear, ReplicatedLinear, RowParallelLinear, @@ -25,9 +27,32 @@ from sglang.srt.layers.vocab_parallel_embedding import ( ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models import mixtral +from sglang.srt.models.granite import build_attention_sinks, granite_layer_attn_params +from sglang.srt.models.utils import WeightsMapper from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import add_prefix +SHARED_MOE_MODEL_TYPES = frozenset({"granitemoeshared", "granitemoe_swa"}) + + +def granitemoe_split_expert_weights( + weights: Iterable[tuple[str, torch.Tensor]], +) -> Iterator[tuple[str, torch.Tensor]]: + """Split the packed expert tensors into per-expert w1/w2/w3.""" + for name, weight in weights: + if name.endswith(".experts.gate_up_proj"): + for e in range(weight.size(0)): + expert = name.replace(".experts.gate_up_proj", f".experts.{e}") + w1, w3 = weight[e].chunk(2, dim=0) + yield f"{expert}.w1.weight", w1 + yield f"{expert}.w3.weight", w3 + elif name.endswith(".experts.down_proj"): + for e in range(weight.size(0)): + expert = name.replace(".experts.down_proj", f".experts.{e}") + yield f"{expert}.w2.weight", weight[e] + else: + yield name, weight + class GraniteMoeMoE(nn.Module): """A tensor-parallel MoE implementation for GraniteMoe that shards each @@ -89,16 +114,56 @@ class GraniteMoeMoE(nn.Module): return final_hidden_states.view(orig_shape) +class GraniteMoeSharedMLP(nn.Module): + + def __init__( + self, + config: GraniteConfig, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + + self.input_size = config.hidden_size + self.hidden_size = config.shared_intermediate_size + self.input_linear = MergedColumnParallelLinear( + input_size=self.input_size, + output_sizes=[self.hidden_size] * 2, + bias=False, + quant_config=quant_config, + prefix=add_prefix("input_linear", prefix), + ) + self.output_linear = RowParallelLinear( + self.hidden_size, + self.input_size, + bias=False, + quant_config=quant_config, + prefix=add_prefix("output_linear", prefix), + ) + if config.hidden_act != "silu": + raise ValueError( + f"Unsupported activation: {config.hidden_act}. " + "Only silu is supported for now." + ) + self.act_fn = SiluAndMul() + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + gate_up, _ = self.input_linear(hidden_states) + x = self.act_fn(gate_up) + x, _ = self.output_linear(x) + return x + + class GraniteMoeAttention(nn.Module): def __init__( self, + config: GraniteConfig, hidden_size: int, num_heads: int, num_kv_heads: int, max_position: int = 4096 * 32, layer_id: int = 0, - rope_theta: float = 10000, quant_config: Optional[QuantizationConfig] = None, attention_multiplier: Optional[float] = None, prefix: str = "", @@ -127,7 +192,9 @@ class GraniteMoeAttention(nn.Module): if attention_multiplier is not None else self.head_dim**-1 ) - self.rope_theta = rope_theta + sliding_window_size, self.rope_theta, has_sink = granite_layer_attn_params( + config, layer_id + ) self.qkv_proj = QKVParallelLinear( hidden_size, @@ -145,13 +212,18 @@ class GraniteMoeAttention(nn.Module): quant_config=quant_config, prefix=f"{prefix}.o_proj", ) - self.rotary_emb = get_rope( - self.head_dim, - rotary_dim=self.head_dim, - max_position=max_position, - base=int(self.rope_theta), - is_neox_style=True, + self.rotary_emb = ( + get_rope( + self.head_dim, + rotary_dim=self.head_dim, + max_position=max_position, + base=int(self.rope_theta), + is_neox_style=True, + ) + if self.rope_theta + else None ) + self.sinks = build_attention_sinks(self.num_heads) if has_sink else None self.attn = RadixAttention( self.num_heads, self.head_dim, @@ -160,6 +232,7 @@ class GraniteMoeAttention(nn.Module): layer_id=layer_id, quant_config=quant_config, prefix=f"{prefix}.attn", + sliding_window_size=sliding_window_size, ) def forward( @@ -170,8 +243,12 @@ class GraniteMoeAttention(nn.Module): ) -> torch.Tensor: qkv, _ = self.qkv_proj(hidden_states) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - q, k = self.rotary_emb(positions, q, k) - attn_output = self.attn(q, k, v, forward_batch) + if self.rotary_emb is not None: + q, k = self.rotary_emb(positions, q, k) + if self.sinks is None: + attn_output = self.attn(q, k, v, forward_batch) + else: + attn_output = self.attn(q, k, v, forward_batch, sinks=self.sinks) output, _ = self.o_proj(attn_output) return output @@ -187,13 +264,12 @@ class GraniteMoeDecoderLayer(nn.Module): ) -> None: super().__init__() self.hidden_size = config.hidden_size - rope_theta = config.rope_parameters["rope_theta"] self.self_attn = GraniteMoeAttention( + config=config, hidden_size=self.hidden_size, num_heads=config.num_attention_heads, max_position=config.max_position_embeddings, num_kv_heads=config.num_key_value_heads, - rope_theta=rope_theta, layer_id=layer_id, quant_config=quant_config, prefix=f"{prefix}.self_attn", @@ -208,6 +284,18 @@ class GraniteMoeDecoderLayer(nn.Module): quant_config=quant_config, prefix=f"{prefix}.block_sparse_moe", ) + shared_intermediate_size = ( + config.shared_intermediate_size + if config.model_type in SHARED_MOE_MODEL_TYPES + else 0 + ) + self.shared_mlp = ( + None + if shared_intermediate_size == 0 + else GraniteMoeSharedMLP( + config, quant_config=quant_config, prefix=f"{prefix}.shared_mlp" + ) + ) self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.post_attention_layernorm = RMSNorm( @@ -233,7 +321,12 @@ class GraniteMoeDecoderLayer(nn.Module): hidden_states = residual + hidden_states * self.residual_multiplier residual = hidden_states hidden_states = self.post_attention_layernorm(hidden_states) - hidden_states = self.block_sparse_moe(hidden_states) + if self.shared_mlp is None: + hidden_states = self.block_sparse_moe(hidden_states) + else: + # Routed experts consume `hidden_states`, so compute shared first + shared_output = self.shared_mlp(hidden_states) + hidden_states = self.block_sparse_moe(hidden_states) + shared_output hidden_states = residual + hidden_states * self.residual_multiplier return hidden_states @@ -297,6 +390,16 @@ class GraniteMoeModel(nn.Module): class GraniteMoeForCausalLM(nn.Module): + # Legacy and current HF expert / router names with otherwise shared layout + hf_to_sglang_mapper = WeightsMapper( + orig_to_new_suffix={ + ".block_sparse_moe.input_linear.weight": ".block_sparse_moe.experts.gate_up_proj", + ".block_sparse_moe.output_linear.weight": ".block_sparse_moe.experts.down_proj", + ".block_sparse_moe.router.layer.weight": ".block_sparse_moe.gate.weight", + ".block_sparse_moe.router.weight": ".block_sparse_moe.gate.weight", + } + ) + def __init__( self, config: GraniteConfig, @@ -346,42 +449,22 @@ class GraniteMoeForCausalLM(nn.Module): return self.pooler(hidden_states, forward_batch) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - new_weights = {} - for n, p in weights: - if n.endswith(".block_sparse_moe.input_linear.weight"): - for e in range(p.size(0)): - w1_name = n.replace( - ".block_sparse_moe.input_linear.weight", - f".block_sparse_moe.experts.{e}.w1.weight", - ) - w3_name = n.replace( - ".block_sparse_moe.input_linear.weight", - f".block_sparse_moe.experts.{e}.w3.weight", - ) - w1_param, w3_param = p[e].chunk(2, dim=0) - assert w1_name not in new_weights - assert w3_name not in new_weights - new_weights[w1_name] = w1_param - new_weights[w3_name] = w3_param - elif n.endswith(".block_sparse_moe.output_linear.weight"): - for e in range(p.size(0)): - w2_name = n.replace( - ".block_sparse_moe.output_linear.weight", - f".block_sparse_moe.experts.{e}.w2.weight", - ) - w2_param = p[e] - assert w2_name not in new_weights - new_weights[w2_name] = w2_param - elif n.endswith(".block_sparse_moe.router.layer.weight"): - gate_name = n.replace( - ".block_sparse_moe.router.layer.weight", - ".block_sparse_moe.gate.weight", - ) - assert gate_name not in new_weights - new_weights[gate_name] = p - else: - new_weights[n] = p - mixtral.MixtralForCausalLM.load_weights(self, new_weights.items()) + weights = granitemoe_split_expert_weights( + self.hf_to_sglang_mapper.apply(weights) + ) + mixtral.MixtralForCausalLM.load_weights(self, weights) -EntryClass = [GraniteMoeForCausalLM] +class GraniteMoeSharedForCausalLM(GraniteMoeForCausalLM): + pass + + +class GraniteMoeSWAForCausalLM(GraniteMoeForCausalLM): + pass + + +EntryClass = [ + GraniteMoeForCausalLM, + GraniteMoeSharedForCausalLM, + GraniteMoeSWAForCausalLM, +] diff --git a/python/sglang/srt/models/granitemoehybrid.py b/python/sglang/srt/models/granitemoehybrid.py index a1c7fa7ec..ed6f37df5 100644 --- a/python/sglang/srt/models/granitemoehybrid.py +++ b/python/sglang/srt/models/granitemoehybrid.py @@ -2,22 +2,16 @@ from typing import Iterable, Optional import torch from torch import nn -from transformers.models.granitemoeshared import GraniteMoeSharedConfig from sglang.srt.configs.granitemoehybrid import GraniteMoeHybridConfig from sglang.srt.distributed import get_pp_group -from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( HybridLinearAttnBackend, Mamba2AttnBackend, ) from sglang.srt.layers.attention.mamba.mamba import MambaMixer2 from sglang.srt.layers.layernorm import RMSNorm -from sglang.srt.layers.linear import ( - MergedColumnParallelLinear, - QKVParallelLinear, - RowParallelLinear, -) +from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.pooler import Pooler, PoolingType from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -35,47 +29,7 @@ from sglang.srt.models.transformers import maybe_prefix from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import make_layers -from .granitemoe import GraniteMoeMoE - - -# in vLLM this is in a separate file, but keeping it here for decoupling -class GraniteMoeSharedMLP(nn.Module): - def __init__( - self, - config: GraniteMoeSharedConfig, - quant_config: QuantizationConfig | None = None, - prefix: str = "", - ): - super().__init__() - - self.input_size = config.hidden_size - self.hidden_size = config.shared_intermediate_size - self.input_linear = MergedColumnParallelLinear( - input_size=self.input_size, - output_sizes=[self.hidden_size] * 2, - bias=False, - quant_config=quant_config, - prefix=f"{prefix}.input_linear", - ) - self.output_linear = RowParallelLinear( - self.hidden_size, - self.input_size, - bias=False, - quant_config=quant_config, - prefix=f"{prefix}.output_linear", - ) - if config.hidden_act != "silu": - raise ValueError( - f"Unsupported activation: {config.hidden_act}. " - "Only silu is supported for now." - ) - self.act_fn = SiluAndMul() - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - gate_up, _ = self.input_linear(hidden_states) - x = self.act_fn(gate_up) - x, _ = self.output_linear(x) - return x +from .granitemoe import GraniteMoeMoE, GraniteMoeSharedMLP class GraniteMoeHybridMambaDecoderLayer(nn.Module):