From b63f8416b3b73bafdec029005c5db36bad207b44 Mon Sep 17 00:00:00 2001 From: Stanislav Petrov <50638944+GungnirAP@users.noreply.github.com> Date: Mon, 21 Sep 2026 09:37:55 +0300 Subject: [PATCH] [Feature] Gigachat 3.5 support (#29189) Co-authored-by: Stanislav Petrov Co-authored-by: Viacheslav Barinov Co-authored-by: Viacheslav Co-authored-by: Xinyuan Tong Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> --- .../advanced_features/server_arguments.mdx | 4 +- .../supported-models/generative_models.mdx | 5 + .../arg_groups/model_overrides/__init__.py | 1 + .../arg_groups/model_overrides/gigachat35.py | 26 + python/sglang/srt/configs/__init__.py | 2 + python/sglang/srt/configs/gigachat35.py | 252 +++++++ python/sglang/srt/configs/model_config.py | 7 + .../srt/function_call/function_call_parser.py | 2 + .../srt/function_call/gigachat35_detector.py | 151 ++++ .../layers/attention/attention_registry.py | 3 +- python/sglang/srt/models/gigachat35.py | 668 ++++++++++++++++++ python/sglang/srt/models/gigachat35_mtp.py | 214 ++++++ python/sglang/srt/parser/reasoning_parser.py | 1 + .../sglang/srt/parser/template_detection.py | 6 + .../multi_layer_eagle_worker_v2.py | 13 +- .../srt/utils/hf_transformers/common.py | 2 + .../unit/parser/test_template_manager.py | 16 + 17 files changed, 1366 insertions(+), 7 deletions(-) create mode 100644 python/sglang/srt/arg_groups/model_overrides/gigachat35.py create mode 100644 python/sglang/srt/configs/gigachat35.py create mode 100644 python/sglang/srt/function_call/gigachat35_detector.py create mode 100644 python/sglang/srt/models/gigachat35.py create mode 100644 python/sglang/srt/models/gigachat35_mtp.py diff --git a/docs/docs/advanced_features/server_arguments.mdx b/docs/docs/advanced_features/server_arguments.mdx index c386e6d30..608fe3166 100644 --- a/docs/docs/advanced_features/server_arguments.mdx +++ b/docs/docs/advanced_features/server_arguments.mdx @@ -1154,13 +1154,13 @@ Combining `--enable-response-store` with `--disaggregation-mode=prefill` or `dec `--reasoning-parser` Specify the parser for reasoning models. Use `auto` to detect the parser from the model's chat template. `None` - auto, apertus2509, deepseek-r1, deepseek-v3, deepseek-v4, dots, glm45, ling3, hunyuan, gpt-oss, k2_horizon, kimi, kimi_k2, kimi_k3, mimo, muse, poolside_v1, qwen3, qwen3-thinking, minimax, minimax-append-think, minimax-m3, step3, step3p5, mistral, nemotron_3, interns1, gemma4, inkling, cohere_command4 + auto, apertus2509, deepseek-r1, deepseek-v3, deepseek-v4, dots, glm45, ling3, hunyuan, gpt-oss, k2_horizon, kimi, kimi_k2, kimi_k3, mimo, muse, poolside_v1, qwen3, qwen3-thinking, minimax, minimax-append-think, minimax-m3, step3, step3p5, mistral, nemotron_3, interns1, gemma4, gigachat35, inkling, cohere_command4 `--tool-call-parser` Specify the parser for handling tool-call interactions. Use `auto` to detect the parser from the model's chat template. `None` - auto, apertus2509, cohere_command4, deepseekv3, deepseekv31, deepseekv32, deepseekv4, dots, glm, glm45, glm47, gpt-oss, k2_horizon, kimi_k2, kimi_k3, lfm2, ling3, llama3, mimo, minicpm5, mistral, muse, poolside_v1, pythonic, qwen, qwen25, qwen3_coder, spark25, step3, step3p5, minimax-m2, minimax-m3, trinity, interns1, hermes, hunyuan, gigachat3, gemma4, inkling + auto, apertus2509, cohere_command4, deepseekv3, deepseekv31, deepseekv32, deepseekv4, dots, glm, glm45, glm47, gpt-oss, k2_horizon, kimi_k2, kimi_k3, lfm2, ling3, llama3, mimo, minicpm5, mistral, muse, poolside_v1, pythonic, qwen, qwen25, qwen3_coder, spark25, step3, step3p5, minimax-m2, minimax-m3, trinity, interns1, hermes, hunyuan, gigachat3, gigachat35, gemma4, inkling `--tool-server` diff --git a/docs/docs/supported-models/generative_models.mdx b/docs/docs/supported-models/generative_models.mdx index b3fdff683..6811d2954 100644 --- a/docs/docs/supported-models/generative_models.mdx +++ b/docs/docs/supported-models/generative_models.mdx @@ -308,5 +308,10 @@ in the GitHub search bar. JetBrains/Mellum2-12B-A2.5B-Base, JetBrains/Mellum2-12B-A2.5B-Thinking JetBrains' Qwen3-MoE-based code generation model with interleaved sliding-window/full attention, per-layer-type RoPE, and per-layer dense/sparse MLP routing. + + GigaChat 3.5 + ai-sage/GigaChat3.5-432B-A28B + GigaChat's hybrid MoE model: per-layer mix of DeepSeek-style MLA full attention and Qwen3 Gated-Delta-Net (GDN) linear attention, DeepSeek-style routed + shared experts, and multiple multi-token-prediction (MTP) heads for speculative decoding. Emits tool calls in GCML format. + diff --git a/python/sglang/srt/arg_groups/model_overrides/__init__.py b/python/sglang/srt/arg_groups/model_overrides/__init__.py index 871b71012..1605d1770 100644 --- a/python/sglang/srt/arg_groups/model_overrides/__init__.py +++ b/python/sglang/srt/arg_groups/model_overrides/__init__.py @@ -16,6 +16,7 @@ from sglang.srt.arg_groups.model_overrides import exaone # noqa: F401 from sglang.srt.arg_groups.model_overrides import falcon_h1 # noqa: F401 from sglang.srt.arg_groups.model_overrides import gemma2_gemma3 # noqa: F401 from sglang.srt.arg_groups.model_overrides import gemma4 # noqa: F401 +from sglang.srt.arg_groups.model_overrides import gigachat35 # noqa: F401 from sglang.srt.arg_groups.model_overrides import glm4_moe # noqa: F401 from sglang.srt.arg_groups.model_overrides import gpt_oss # noqa: F401 from sglang.srt.arg_groups.model_overrides import granitemoehybrid # noqa: F401 diff --git a/python/sglang/srt/arg_groups/model_overrides/gigachat35.py b/python/sglang/srt/arg_groups/model_overrides/gigachat35.py new file mode 100644 index 000000000..77a19e8b2 --- /dev/null +++ b/python/sglang/srt/arg_groups/model_overrides/gigachat35.py @@ -0,0 +1,26 @@ +"""Config-time override declarations for gigachat35. + +Architectures: GigaChat35ForCausalLM, GigaChat35ForCausalLMNextN. +""" + +import logging +from typing import Any, Dict + +from sglang.srt.arg_groups.model_override_base import ( + _register_for, + resolving_view, +) + +logger = logging.getLogger(__name__) + + +@_register_for("GigaChat35ForCausalLM", "GigaChat35ForCausalLMNextN") +def _gigachat35_overrides(server_args: Any, hf_config: Any) -> dict: + cfg = resolving_view(server_args) + overrides: Dict[str, Any] = {"disable_shared_experts_fusion": True} + if cfg.speculative_algorithm == "EAGLE": + logger.info( + "Enable multi-layer EAGLE speculative decoding for GigaChat 3.5 model." + ) + overrides["enable_multi_layer_eagle"] = True + return overrides diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py index 476983b04..6c404b17d 100644 --- a/python/sglang/srt/configs/__init__.py +++ b/python/sglang/srt/configs/__init__.py @@ -17,6 +17,7 @@ from sglang.srt.configs.dots_ocr import DotsOCRConfig from sglang.srt.configs.dots_vlm import DotsVLMConfig from sglang.srt.configs.exaone import ExaoneConfig from sglang.srt.configs.falcon_h1 import FalconH1Config +from sglang.srt.configs.gigachat35 import GigaChat35Config from sglang.srt.configs.glm5_next import Glm5NextConfig, Glm5NextTextConfig from sglang.srt.configs.granitemoehybrid import GraniteMoeHybridConfig from sglang.srt.configs.hy_v4 import HYV4Config @@ -132,6 +133,7 @@ __all__ = [ "Dots3Config", "FalconH1Config", "FalconMambaConfig", + "GigaChat35Config", "GraniteMoeHybridConfig", "HYV4Config", "MambaConfig", diff --git a/python/sglang/srt/configs/gigachat35.py b/python/sglang/srt/configs/gigachat35.py new file mode 100644 index 000000000..a18efe06e --- /dev/null +++ b/python/sglang/srt/configs/gigachat35.py @@ -0,0 +1,252 @@ +from __future__ import annotations + +from typing import Optional + +import torch +from transformers.configuration_utils import PretrainedConfig + +from sglang.srt.configs.mamba_utils import ( + Mamba2CacheParams, + Mamba2StateShape, + mamba2_state_dtype, +) + +FULL_ATTENTION = "attention" +LINEAR_ATTENTION = "linear_attention" +_FULL_ATTENTION_ALIASES = {FULL_ATTENTION, "full_attention"} + +_REQUIRED_LINEAR_ATTRS = ( + "linear_conv_kernel_dim", + "linear_key_head_dim", + "linear_value_head_dim", + "linear_num_key_heads", + "linear_num_value_heads", +) + + +class GigaChat35Config(PretrainedConfig): + model_type = "gigachat3_5" + keys_to_ignore_at_inference = ["past_key_values"] + + def __init__( + self, + vocab_size: int = 128256, + hidden_size: int = 1280, + intermediate_size: int = 896, + moe_intermediate_size: int = 896, + num_hidden_layers: int = 12, + num_attention_heads: int = 16, + num_key_value_heads: int = 16, + hidden_act: str = "silu", + max_position_embeddings: int = 4096, + initializer_range: float = 0.006, + rms_norm_eps: float = 1e-6, + use_cache: bool = True, + rope_theta: float = 100000.0, + rope_scaling: Optional[dict] = None, + rope_interleave: bool = True, + attention_bias: bool = False, + tie_word_embeddings: bool = False, + q_lora_rank: Optional[int] = 1536, + kv_lora_rank: int = 512, + qk_nope_head_dim: int = 128, + qk_rope_head_dim: int = 64, + v_head_dim: int = 128, + head_dim: int = 64, + n_routed_experts: int = 64, + n_shared_experts: int = 2, + num_experts_per_tok: int = 6, + moe_layer_freq: int = 1, + first_k_dense_replace: int = 1, + routed_scaling_factor: float = 2.5, + n_group: int = 1, + topk_group: int = 1, + topk_method: str = "noaux_tc", + scoring_func: str = "sigmoid", + norm_topk_prob: bool = True, + use_shared_expert_sigmoid: bool = False, + linear_attention_type: str = "Qwen3NextGatedDeltaNet", + layer_types: Optional[list[str]] = None, + full_attention_layers: Optional[list[int]] = None, + linear_conv_kernel_dim: int = 4, + linear_key_head_dim: int = 128, + linear_value_head_dim: int = 128, + linear_num_key_heads: int = 8, + linear_num_value_heads: int = 16, + linear_sigmoid_gate_scale: float = 2.0, + output_gate_type: str = "sigmoid", + norm_type: str = "ZeroCenteredGatedNorm", + layernorm_type: str = "pre_post", + layernorm_gating_weight: float = 2.0, + gated_attention: bool = True, + use_mla_scaling_factor: bool = True, + swiglu_limit: Optional[float] = None, + num_nextn_predict_layers: int = 0, + nextn_is_sparse: bool = False, + **kwargs, + ): + self.vocab_size = vocab_size + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.moe_intermediate_size = moe_intermediate_size + self.num_hidden_layers = num_hidden_layers + self.num_attention_heads = num_attention_heads + self.num_key_value_heads = num_key_value_heads + self.hidden_act = hidden_act + self.max_position_embeddings = max_position_embeddings + self.initializer_range = initializer_range + self.rms_norm_eps = rms_norm_eps + self.use_cache = use_cache + self.rope_theta = rope_theta + self.rope_scaling = rope_scaling + self.rope_interleave = rope_interleave + self.attention_bias = attention_bias + + self.q_lora_rank = q_lora_rank + self.kv_lora_rank = kv_lora_rank + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.head_dim = head_dim + + self.n_routed_experts = n_routed_experts + self.n_shared_experts = n_shared_experts + self.num_experts_per_tok = num_experts_per_tok + self.moe_layer_freq = moe_layer_freq + self.first_k_dense_replace = first_k_dense_replace + self.routed_scaling_factor = routed_scaling_factor + self.n_group = n_group + self.topk_group = topk_group + self.topk_method = topk_method + self.scoring_func = scoring_func + self.norm_topk_prob = norm_topk_prob + self.use_shared_expert_sigmoid = use_shared_expert_sigmoid + + self.linear_attention_type = linear_attention_type + self.layer_types = layer_types + self.full_attention_layers = full_attention_layers + self.linear_conv_kernel_dim = linear_conv_kernel_dim + self.linear_key_head_dim = linear_key_head_dim + self.linear_value_head_dim = linear_value_head_dim + self.linear_num_key_heads = linear_num_key_heads + self.linear_num_value_heads = linear_num_value_heads + self.linear_sigmoid_gate_scale = linear_sigmoid_gate_scale + self.output_gate_type = output_gate_type + self.linear_num_key_heads_cpu = linear_num_key_heads + self.linear_num_value_heads_cpu = linear_num_value_heads + + self.norm_type = norm_type + self.layernorm_type = layernorm_type + self.layernorm_gating_weight = layernorm_gating_weight + self.gated_attention = gated_attention + self.use_mla_scaling_factor = use_mla_scaling_factor + self.swiglu_limit = ( + swiglu_limit if (swiglu_limit is not None and swiglu_limit > 0) else None + ) + + self.num_nextn_predict_layers = num_nextn_predict_layers + self.nextn_is_sparse = nextn_is_sparse + + super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs) + + _dtype = getattr(self, "torch_dtype", None) or getattr(self, "dtype", None) + if isinstance(_dtype, str): + _dtype = getattr(torch, _dtype, torch.bfloat16) + self.torch_dtype = _dtype or torch.bfloat16 + + def _resolve_layer_types(self) -> list[str]: + """Return a normalized per-layer list of FULL_ATTENTION / LINEAR_ATTENTION. + + Resolution order: explicit ``layer_types`` -> ``full_attention_layers``. + Defaults to all-linear if nothing is specified (degenerate, but + well-defined). + """ + n = self.num_hidden_layers + + if self.layer_types is not None: + if len(self.layer_types) != n: + raise ValueError( + f"layer_types must have length num_hidden_layers ({n}), " + f"got {len(self.layer_types)}." + ) + resolved = [] + for idx, lt in enumerate(self.layer_types): + if lt == LINEAR_ATTENTION: + resolved.append(LINEAR_ATTENTION) + elif lt in _FULL_ATTENTION_ALIASES: + resolved.append(FULL_ATTENTION) + else: + raise ValueError(f"Unsupported layer type {lt!r} at index {idx}.") + return resolved + + resolved = [LINEAR_ATTENTION] * n + if self.full_attention_layers is not None: + for lid in self.full_attention_layers: + resolved[lid] = FULL_ATTENTION + return resolved + + @property + def layers_block_type(self) -> list[str]: + return self._resolve_layer_types() + + @property + def linear_layer_ids(self) -> list[int]: + return [ + i for i, lt in enumerate(self.layers_block_type) if lt == LINEAR_ATTENTION + ] + + @property + def full_attention_layer_ids(self) -> list[int]: + return [ + i for i, lt in enumerate(self.layers_block_type) if lt == FULL_ATTENTION + ] + + def is_linear_attention_layer(self, layer_id: int) -> bool: + if layer_id >= self.num_hidden_layers: + return False + return self.layers_block_type[layer_id] == LINEAR_ATTENTION + + @property + def mamba2_cache_params(self) -> Mamba2CacheParams: + from sglang.srt.runtime_context import get_parallel + + missing = [a for a in _REQUIRED_LINEAR_ATTRS if getattr(self, a, None) is None] + if missing: + raise ValueError( + "GigaChat35 hybrid GDN config is missing required linear-attention " + "fields: " + ", ".join(missing) + ) + + shape = Mamba2StateShape.create( + tp_world_size=get_parallel().attn_tp_size, + intermediate_size=self.linear_value_head_dim * self.linear_num_value_heads, + n_groups=self.linear_num_key_heads, + num_heads=self.linear_num_value_heads, + head_dim=self.linear_value_head_dim, + state_size=self.linear_key_head_dim, + conv_kernel=self.linear_conv_kernel_dim, + ) + return Mamba2CacheParams( + shape=shape, + layers=self.linear_layer_ids, + dtype=mamba2_state_dtype(self), + ) + + +from sglang.srt.configs.linear_attn_model_registry import ( # noqa: E402 + LinearAttnModelSpec, + register_linear_attn_model, +) + +register_linear_attn_model( + LinearAttnModelSpec( + config_class=GigaChat35Config, + backend_class_name="sglang.srt.layers.attention.linear.gdn_backend.GDNAttnBackend", + arch_names=[ + "GigaChat35ForCausalLM", + "GigaChat35ForCausalLMNextN", + ], + uses_mamba_radix_cache=True, + support_mamba_cache=True, + ) +) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 96161ab8a..1f5496d15 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -927,6 +927,11 @@ class ModelConfig: and self.hf_config.architectures[0] == "InklingForConditionalGeneration" ): self.hf_config.architectures[0] = "InklingForConditionalGenerationMTP" + if ( + is_draft_model + and self.hf_config.architectures[0] == "GigaChat35ForCausalLM" + ): + self.hf_config.architectures[0] = "GigaChat35ForCausalLMNextN" if ( is_draft_model and self.hf_config.architectures[0] == "Step3p7ForConditionalGeneration" @@ -1191,6 +1196,8 @@ class ModelConfig: or "MistralLarge3ForCausalLMEagle" in self.hf_config.architectures or "KimiK25ForConditionalGeneration" in self.hf_config.architectures or "Eagle3DeepseekV2ForCausalLM" in self.hf_config.architectures + or "GigaChat35ForCausalLM" in self.hf_config.architectures + or "GigaChat35ForCausalLMNextN" in self.hf_config.architectures ): self.head_dim = 256 self.attention_arch = AttentionArch.MLA diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 0bb26b918..2932cfc12 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -23,6 +23,7 @@ from sglang.srt.function_call.deepseekv41_detector import DeepSeekV41Detector from sglang.srt.function_call.dots_detector import DotsToolDetector from sglang.srt.function_call.gemma4_detector import Gemma4Detector from sglang.srt.function_call.gigachat3_detector import GigaChat3Detector +from sglang.srt.function_call.gigachat35_detector import GigaChat35Detector from sglang.srt.function_call.glm4_moe_detector import ( Glm4MoeDetector, GlmSpecialTokenConfig, @@ -109,6 +110,7 @@ class FunctionCallParser: "hermes": HermesDetector, "hunyuan": HunyuanDetector, "gigachat3": GigaChat3Detector, + "gigachat35": GigaChat35Detector, "gemma4": Gemma4Detector, "inkling": InklingDetector, } diff --git a/python/sglang/srt/function_call/gigachat35_detector.py b/python/sglang/srt/function_call/gigachat35_detector.py new file mode 100644 index 000000000..3f07a3ce8 --- /dev/null +++ b/python/sglang/srt/function_call/gigachat35_detector.py @@ -0,0 +1,151 @@ +import json +import logging +import re +from typing import Any, Dict, List + +from sglang.srt.entrypoints.openai.protocol import Tool +from sglang.srt.function_call.base_format_detector import BaseFormatDetector +from sglang.srt.function_call.core_types import ( + StreamingParseResult, + ToolCallItem, + _GetInfoFunc, +) + +logger = logging.getLogger(__name__) + +_GCML_OPEN = "<|GCML|tool_calls>" +_GCML_CLOSE = "" +_GCML_INVOKE_RE = re.compile( + r"<|GCML|invoke\s+name=\"(?P[^\"]*)\"\s*>" + r"(?P.*?)" + r"", + re.DOTALL, +) +_GCML_PARAM_RE = re.compile( + r"<|GCML|parameter\s+name=\"(?P[^\"]*)\"\s+" + r"string=\"(?Ptrue|false)\"\s*>" + r"(?P.*?)" + r"", + re.DOTALL, +) +_TRAILING_MARKER_RE = re.compile(r"(?:<\|message_sep\|>|)+\s*$") + + +def _strip_trailing_markers(text: str) -> str: + return _TRAILING_MARKER_RE.sub("", text) + + +def _parse_gcml_invoke_body(body: str) -> Dict[str, Any]: + """Parse the <|GCML|parameter ...> entries inside one invoke block.""" + args: Dict[str, Any] = {} + for m in _GCML_PARAM_RE.finditer(body): + raw = m.group("value") + if m.group("is_string") == "true": + args[m.group("name")] = raw + continue + try: + args[m.group("name")] = json.loads(raw) + except (json.JSONDecodeError, ValueError): + args[m.group("name")] = raw + return args + + +class GigaChat35Detector(BaseFormatDetector): + def __init__(self) -> None: + super().__init__() + self.bot_token = _GCML_OPEN + self.eot_token = _GCML_CLOSE + self._tool_region_started: bool = False + self._emitted_invokes: int = 0 + + def has_tool_call(self, text: str) -> bool: + """True if the text contains the GCML tool-call opening marker.""" + return self.bot_token in text + + def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult: + """Non-streaming parse of a complete model output.""" + if self.bot_token not in text: + return StreamingParseResult( + normal_text=_strip_trailing_markers(text), calls=[] + ) + + leading, _, after_open = text.partition(self.bot_token) + invokes_body, _, _ = after_open.partition(self.eot_token) + + actions = [ + { + "name": m.group("name"), + "arguments": _parse_gcml_invoke_body(m.group("body")), + } + for m in _GCML_INVOKE_RE.finditer(invokes_body) + ] + if not actions: + return StreamingParseResult( + normal_text=_strip_trailing_markers(text), calls=[] + ) + + calls = self.parse_base_json(actions, tools) + normal_text = _strip_trailing_markers(leading).rstrip("\n") + return StreamingParseResult(normal_text=normal_text, calls=calls) + + def parse_streaming_increment( + self, new_text: str, tools: List[Tool] + ) -> StreamingParseResult: + self._buffer += new_text + current_text = self._buffer + + if not self._tool_region_started and self.bot_token not in current_text: + if self._ends_with_partial_token(current_text, self.bot_token): + return StreamingParseResult() + self._buffer = "" + return StreamingParseResult(normal_text=current_text) + + if not self._tool_region_started: + leading, _, rest = current_text.partition(self.bot_token) + self._tool_region_started = True + self._buffer = self.bot_token + rest + current_text = self._buffer + if leading: + return StreamingParseResult(normal_text=leading) + + if not hasattr(self, "_tool_indices"): + self._tool_indices = self._get_tool_indices(tools) + + _, _, after_open = current_text.partition(self.bot_token) + invokes_body, _, _ = after_open.partition(self.eot_token) + matches = list(_GCML_INVOKE_RE.finditer(invokes_body)) + calls: List[ToolCallItem] = [] + for i in range(self._emitted_invokes, len(matches)): + m = matches[i] + name = m.group("name") + if name not in self._tool_indices: + logger.warning(f"[GigaChat35] undefined function call: {name}") + continue + args = _parse_gcml_invoke_body(m.group("body")) + args_json = json.dumps(args, ensure_ascii=False) + calls.append(ToolCallItem(tool_index=i, name=name, parameters=args_json)) + while len(self.prev_tool_call_arr) <= i: + self.prev_tool_call_arr.append({}) + self.prev_tool_call_arr[i] = {"name": name, "arguments": args} + while len(self.streamed_args_for_tool) <= i: + self.streamed_args_for_tool.append("") + self.streamed_args_for_tool[i] = args_json + + self._emitted_invokes = len(matches) + return StreamingParseResult(calls=calls) + + def finish(self, tools: List[Tool]) -> StreamingParseResult: + result = self.parse_streaming_increment("", tools) + leftover, self._buffer = self._buffer, "" + if leftover and not self._tool_region_started: + result.normal_text += leftover + return result + + def supports_structural_tag(self) -> bool: + """GigaChat 3.5 GCML does not use structural tags.""" + return False + + def structure_info(self) -> _GetInfoFunc: + raise NotImplementedError( + "GigaChat35Detector does not support structural_tag format." + ) diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index 2ed16e52a..165e7165c 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -522,8 +522,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac else: spec_result = get_linear_attn_config(runner.model_config.hf_config) if spec_result is not None: - spec, _ = spec_result - cfg = runner.model_config + spec, cfg = spec_result BackendClass = import_backend_class(spec.backend_class_name) linear_attn_backend = BackendClass(runner) if spec.hybrid_backend_class_name is not None: diff --git a/python/sglang/srt/models/gigachat35.py b/python/sglang/srt/models/gigachat35.py new file mode 100644 index 000000000..724791800 --- /dev/null +++ b/python/sglang/srt/models/gigachat35.py @@ -0,0 +1,668 @@ +"""GigaChat 3.5 model. + +GigaChat 3.5 is a DeepSeek-V3-style model (MLA attention + DeepSeek MoE) with a +*hybrid* attention stack: most layers use a Qwen3-Next Gated-Delta-Net (GDN) +linear-attention block, while a periodic subset use full MLA attention. On top of +the DeepSeek base it adds: + + * gated RMSNorm (optionally zero-centered) with a low-rank gating bottleneck, + used as a 4-norm sandwich (pre/post attention, pre/post MLP); + * gated attention (a sigmoid gate applied to the attention output before + ``o_proj``); + * an MLA query/key scaling factor (``alpha_q`` / ``alpha_kv``). +""" + +from __future__ import annotations + +from typing import Any, Dict, Iterable, Optional + +import torch +import torch.nn.functional as F +from torch import nn +from transformers import PretrainedConfig + +import sglang.srt.models.deepseek_v2 as deepseek_v2 +from sglang.srt.configs.gigachat35 import GigaChat35Config +from sglang.srt.distributed import get_pp_group +from sglang.srt.layers.communicator import get_attn_tp_context +from sglang.srt.layers.layernorm import GemmaRMSNorm, RMSNorm +from sglang.srt.layers.linear import ColumnParallelLinear +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, PPProxyTensors +from sglang.srt.models.deepseek_common.deepseek_weight_loader import ( + DeepseekV2WeightLoaderMixin, +) +from sglang.srt.models.qwen3_next import Qwen3GatedDeltaNet +from sglang.srt.runtime_context import get_forward, get_parallel +from sglang.srt.utils import BumpAllocator, add_prefix, make_layers + +_GATED_NORM_LOW_RANK = 16 + + +class GigaChat35GatedRMSNorm(nn.Module): + """RMSNorm (optionally zero-centered) followed by a low-rank sigmoid gate. + + ``scale`` folds the optional MLA ``alpha`` factor into the q/kv a-norms. + Supports the fused residual-add convention used by sglang decoder layers: + ``forward(x, residual)`` returns ``(normed, residual + x)``. + """ + + def __init__( + self, + hidden_size: int, + eps: float = 1e-6, + layernorm_gating_weight: float = 2.0, + zero_centered: bool = True, + scale: float = 1.0, + low_rank: int = _GATED_NORM_LOW_RANK, + ) -> None: + super().__init__() + self.hidden_size = hidden_size + self.variance_epsilon = eps + self.zero_centered = zero_centered + self.layernorm_gating_weight = layernorm_gating_weight + self.scale = scale + self.r = low_rank + self.weight = nn.Parameter( + torch.zeros(hidden_size) if zero_centered else torch.ones(hidden_size) + ) + self.gate_up_lowrank = nn.Linear(hidden_size, self.r, bias=False) + self.gate_down_lowrank = nn.Linear(self.r, hidden_size, bias=False) + + def _norm_gate(self, x: torch.Tensor) -> torch.Tensor: + out = x.float() + out = out * torch.rsqrt( + out.pow(2).mean(-1, keepdim=True) + self.variance_epsilon + ) + if self.zero_centered: + out = (out * (1.0 + self.weight.float())).to(x.dtype) + else: + out = (out * self.weight.float()).to(x.dtype) + gate = F.linear(out, self.gate_up_lowrank.weight.to(out.dtype)) + gate = F.silu(gate.float()).to(out.dtype) + gate = F.linear(gate, self.gate_down_lowrank.weight.to(out.dtype)) + gate = torch.sigmoid(gate.float()).to(out.dtype) + out = (out * self.layernorm_gating_weight * gate).to(x.dtype) + if self.scale != 1.0: + out = (out.float() * self.scale).to(x.dtype) + return out + + def forward( + self, + x: torch.Tensor, + residual: Optional[torch.Tensor] = None, + post_residual_addition: Optional[torch.Tensor] = None, + ): + if post_residual_addition is not None: + x = x + post_residual_addition + if residual is not None: + x = x + residual + residual = x + out = self._norm_gate(x) + return out if residual is None else (out, residual) + + +def build_norm( + config: GigaChat35Config, hidden_size: int, scale: float = 1.0 +) -> nn.Module: + """Construct the configured norm for a GigaChat 3.5 module.""" + norm_type = getattr(config, "norm_type", "LlamaRMSNorm") + eps = config.rms_norm_eps + if norm_type == "LlamaRMSNorm": + return RMSNorm(hidden_size, eps=eps) + if norm_type == "ZeroCenteredRMSNorm": + return GemmaRMSNorm(hidden_size, eps=eps) + if norm_type in ("ZeroCenteredGatedNorm", "GatedNorm"): + return GigaChat35GatedRMSNorm( + hidden_size, + eps=eps, + layernorm_gating_weight=getattr(config, "layernorm_gating_weight", 2.0), + zero_centered=(norm_type == "ZeroCenteredGatedNorm"), + scale=scale, + ) + raise ValueError(f"Unsupported norm_type for GigaChat 3.5: {norm_type!r}") + + +class GigaChat35PassthroughNorm(nn.Module): + """Identity norm with the fused residual-add convention. + + Used as the ``input_layernorm`` slot when a layer has no pre-norm + (``layernorm_type="post"``); it only folds the residual so the + LayerCommunicator's prepare-step bookkeeping still works. + """ + + def forward( + self, + x: torch.Tensor, + residual: Optional[torch.Tensor] = None, + post_residual_addition: Optional[torch.Tensor] = None, + ): + if post_residual_addition is not None: + x = x + post_residual_addition + if residual is None: + return x, x + merged = x + residual + return merged, merged + + +class GigaChat35MlpPrepNorm(nn.Module): + """Pre-MLP norm wrapper that folds the post-attention (sandwich) norm. + + Replaces the LayerCommunicator's ``post_attention_layernorm`` slot so the + optional ``post_self_attn_layernorm`` (applied to the attention output, + before the residual add) and the pre-MLP ``post_attention_layernorm`` are + both threaded through ``prepare_mlp`` -- matching the per-layer math of the + ``pre_post`` sandwich exactly. + """ + + def __init__( + self, + pre_layernorm: Optional[nn.Module], + post_layernorm: Optional[nn.Module], + ) -> None: + super().__init__() + self.pre_layernorm = pre_layernorm + self.post_layernorm = post_layernorm + + def forward( + self, + x: torch.Tensor, + residual: Optional[torch.Tensor] = None, + post_residual_addition: Optional[torch.Tensor] = None, + ): + if self.post_layernorm is not None: + x = self.post_layernorm(x) + if post_residual_addition is not None: + x = x + post_residual_addition + if residual is None: + merged = x + else: + merged = x + residual + if self.pre_layernorm is None: + return merged, merged + return self.pre_layernorm(merged), merged + + +_GIGACHAT_WEIGHT_NAME_REMAP = ( + (".gate_up_projection.", ".gate_up_lowrank."), + (".gate_down_projection.", ".gate_down_lowrank."), + (".self_attn.gate_proj.", ".self_attn.attn_gate."), +) + + +def _remap_gigachat_weight_names(weights): + for name, loaded_weight in weights: + for src, dst in _GIGACHAT_WEIGHT_NAME_REMAP: + if src in name: + name = name.replace(src, dst) + yield name, loaded_weight + + +class GigaChat35GatedDeltaNet(Qwen3GatedDeltaNet): + def __init__( + self, + config: PretrainedConfig, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + alt_stream: Optional[torch.cuda.Stream] = None, + prefix: str = "", + ) -> None: + super().__init__( + config=config, + layer_id=layer_id, + quant_config=quant_config, + alt_stream=alt_stream, + prefix=prefix, + ) + scale = float(getattr(config, "linear_sigmoid_gate_scale", 1.0)) + zero_centered = ( + "zero_centered" in str(getattr(config, "linear_gating_type", "")).lower() + ) + if (scale != 1.0 or zero_centered) and getattr( + self.norm, "weight", None + ) is not None: + + def _fold_loader(param, loaded_weight, _scale=scale, _zc=zero_centered): + w = loaded_weight.to(param.dtype) + if _zc: + w = 1.0 + w + if _scale != 1.0: + w = _scale * w + param.data.copy_(w) + + self.norm.weight.weight_loader = _fold_loader + + +class GigaChat35AttentionMLA(deepseek_v2.DeepseekV2AttentionMLA): + def __init__( + self, + config: PretrainedConfig, + hidden_size: int, + num_heads: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + q_lora_rank: Optional[int], + kv_lora_rank: int, + rope_theta: float = 10000, + rope_scaling: Optional[Dict[str, Any]] = None, + max_position_embeddings: int = 8192, + quant_config: Optional[QuantizationConfig] = None, + reduce_results: bool = True, + layer_id: Optional[int] = None, + prefix: str = "", + alt_stream: Optional[torch.cuda.Stream] = None, + ) -> None: + super().__init__( + config=config, + hidden_size=hidden_size, + num_heads=num_heads, + qk_nope_head_dim=qk_nope_head_dim, + qk_rope_head_dim=qk_rope_head_dim, + v_head_dim=v_head_dim, + q_lora_rank=q_lora_rank, + kv_lora_rank=kv_lora_rank, + rope_theta=rope_theta, + rope_scaling=rope_scaling, + max_position_embeddings=max_position_embeddings, + quant_config=quant_config, + reduce_results=reduce_results, + layer_id=layer_id, + prefix=prefix, + alt_stream=alt_stream, + ) + + self._norm_type = getattr(config, "norm_type", "LlamaRMSNorm") + use_scaling = bool(getattr(config, "use_mla_scaling_factor", False)) + q_hidden = hidden_size if q_lora_rank is None else q_lora_rank + kv_hidden = kv_lora_rank + alpha_q = (hidden_size / q_hidden) ** 0.5 if use_scaling else 1.0 + alpha_kv = (hidden_size / kv_hidden) ** 0.5 if use_scaling else 1.0 + + if self._norm_type != "LlamaRMSNorm": + if q_lora_rank is not None and hasattr(self, "q_a_layernorm"): + self.q_a_layernorm = build_norm(config, q_lora_rank, scale=alpha_q) + self.kv_a_layernorm = build_norm(config, kv_lora_rank, scale=alpha_kv) + else: + if use_scaling: + if q_lora_rank is not None and hasattr(self, "q_a_layernorm"): + with torch.no_grad(): + self.q_a_layernorm.weight.mul_(alpha_q) + with torch.no_grad(): + self.kv_a_layernorm.weight.mul_(alpha_kv) + + self.gated_attention = bool(getattr(config, "gated_attention", False)) + self._gate_input: Optional[torch.Tensor] = None + if self.gated_attention: + self.attn_gate = ColumnParallelLinear( + hidden_size, + self.num_heads * self.v_head_dim, + bias=False, + quant_config=quant_config, + prefix=add_prefix("attn_gate", prefix), + tp_rank=get_parallel().attn_tp_rank, + tp_size=get_parallel().attn_tp_size, + ) + self.o_proj.register_forward_pre_hook(self._o_proj_gate_hook) + + def _o_proj_gate_hook(self, module, args): + if not self.gated_attention or self._gate_input is None: + return None + attn_output = args[0] + if not isinstance(attn_output, torch.Tensor) or not isinstance( + self._gate_input, torch.Tensor + ): + return None + gate, _ = self.attn_gate(self._gate_input) + attn_output = attn_output * torch.sigmoid(gate) + return (attn_output, *args[1:]) + + def dispatch_attn_forward_method(self, forward_batch: ForwardBatch): + method = super().dispatch_attn_forward_method(forward_batch) + AF = deepseek_v2.AttnForwardMethod + fused = [ + getattr(AF, name, None) + for name in ("MLA_FUSED_ROPE", "MLA_FUSED_ROPE_ROCM", "MLA_FUSED_ROPE_CPU") + ] + if method in [m for m in fused if m is not None]: + return AF.MLA + return method + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + zero_allocator: BumpAllocator, + **kwargs, + ): + if self.gated_attention: + self._gate_input = ( + hidden_states if isinstance(hidden_states, torch.Tensor) else None + ) + try: + return super().forward( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + zero_allocator=zero_allocator, + **kwargs, + ) + finally: + self._gate_input = None + + +class GigaChat35DecoderLayer(deepseek_v2.DeepseekV2DecoderLayer): + def __init__( + self, + config: GigaChat35Config, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + alt_stream: Optional[torch.cuda.Stream] = None, + is_nextn: bool = False, + ) -> None: + super().__init__( + config=config, + layer_id=layer_id, + quant_config=quant_config, + is_nextn=is_nextn, + prefix=prefix, + alt_stream=alt_stream, + ) + + attn_layer_id = config.num_hidden_layers if is_nextn else layer_id + self.use_linear_attn = config.is_linear_attention_layer(attn_layer_id) + + if self.use_linear_attn: + self.self_attn = GigaChat35GatedDeltaNet( + config=config, + layer_id=layer_id, + quant_config=quant_config, + alt_stream=alt_stream, + prefix=add_prefix("self_attn", prefix), + ) + if hasattr(self.self_attn, "out_proj"): + self.self_attn.out_proj.reduce_results = False + self.layer_communicator.qkv_latent_func = None + else: + self.self_attn = GigaChat35AttentionMLA( + config=config, + hidden_size=self.hidden_size, + num_heads=config.num_attention_heads, + qk_nope_head_dim=config.qk_nope_head_dim, + qk_rope_head_dim=config.qk_rope_head_dim, + v_head_dim=config.v_head_dim, + q_lora_rank=getattr(config, "q_lora_rank", None), + kv_lora_rank=config.kv_lora_rank, + rope_theta=config.rope_theta, + rope_scaling=config.rope_scaling, + max_position_embeddings=config.max_position_embeddings, + quant_config=quant_config, + reduce_results=False, + layer_id=layer_id, + prefix=add_prefix("self_attn", prefix), + alt_stream=alt_stream, + ) + self.layer_communicator.qkv_latent_func = self.self_attn.prepare_qkv_latent + + self.is_sparse = self.is_layer_sparse + + swiglu_limit = float(getattr(config, "swiglu_limit", None) or 0.0) + experts = getattr(self.mlp, "experts", None) + runner_config = getattr(experts, "moe_runner_config", None) + if swiglu_limit > 0 and runner_config is not None: + runner_config.gemm1_clamp_limit = swiglu_limit + + layernorm_type = getattr(config, "layernorm_type", "pre") + self._use_pre = layernorm_type in ("pre", "pre_post") + self._use_post = layernorm_type in ("post", "pre_post") + + self.input_layernorm = build_norm(config, config.hidden_size) + self.post_attention_layernorm = build_norm(config, config.hidden_size) + self.post_self_attn_layernorm = ( + build_norm(config, config.hidden_size) if self._use_post else None + ) + self.post_feedforward_layernorm = ( + build_norm(config, config.hidden_size) if self._use_post else None + ) + + attn_prepare_layernorm = ( + self.input_layernorm if self._use_pre else GigaChat35PassthroughNorm() + ) + mlp_prepare_layernorm = GigaChat35MlpPrepNorm( + pre_layernorm=self.post_attention_layernorm if self._use_pre else None, + post_layernorm=self.post_self_attn_layernorm, + ) + self.layer_communicator.input_layernorm = attn_prepare_layernorm + self.layer_communicator.post_attention_layernorm = mlp_prepare_layernorm + + def _is_layer_sparse(self, layer_id: int, is_nextn: bool) -> bool: + if is_nextn: + return bool(getattr(self.config, "nextn_is_sparse", False)) + return ( + self.config.n_routed_experts is not None + and layer_id >= self.config.first_k_dense_replace + and layer_id % self.config.moe_layer_freq == 0 + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + residual: Optional[torch.Tensor], + zero_allocator: BumpAllocator, + **kwargs, + ) -> tuple[torch.Tensor, torch.Tensor]: + hidden_states, residual = self.layer_communicator.prepare_attn( + hidden_states, residual, forward_batch + ) + + if self.use_linear_attn: + hidden_states = self.self_attn( + hidden_states=hidden_states, + forward_batch=forward_batch, + ) + else: + hidden_states = self.self_attn( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + zero_allocator=zero_allocator, + layer_scatter_modes=self.layer_scatter_modes, + ) + get_attn_tp_context().clear_attn_inputs() + + hidden_states, residual = self.layer_communicator.prepare_mlp( + hidden_states, residual, forward_batch + ) + + mlp_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( + forward_batch + ) + # Unlike deepseek_v2, no moe_output_buffer_ctx here: non-inplace MoE + # runners then allocate their output per forward instead of recycling + # the layer-input buffer. The default (inplace) runners are unaffected. + with get_forward().scoped( + fuse_mlp_allreduce=False, + mlp_reduce_scatter=mlp_reduce_scatter, + ): + hidden_states = self.mlp(hidden_states, forward_batch) + + if self.post_feedforward_layernorm is not None: + hidden_states = self.post_feedforward_layernorm(hidden_states) + + hidden_states, residual = self.layer_communicator.postprocess_layer( + hidden_states, residual, forward_batch + ) + return hidden_states, residual + + +class GigaChat35Model(nn.Module): + def __init__( + self, + config: GigaChat35Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + self.padding_idx = getattr(config, "pad_token_id", None) + self.vocab_size = config.vocab_size + self.pp_group = get_pp_group() + + if self.pp_group.is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=add_prefix("embed_tokens", prefix), + ) + else: + self.embed_tokens = None + + self.alt_stream = torch.cuda.Stream() if torch.cuda.is_available() else None + + self.layers, self.start_layer, self.end_layer = make_layers( + config.num_hidden_layers, + lambda idx, prefix: GigaChat35DecoderLayer( + config=config, + layer_id=idx, + quant_config=quant_config, + prefix=prefix, + alt_stream=self.alt_stream, + ), + pp_rank=self.pp_group.rank_in_group, + pp_size=self.pp_group.world_size, + prefix=add_prefix("layers", prefix), + ) + + if self.pp_group.is_last_rank: + self.norm = build_norm(config, config.hidden_size) + else: + self.norm = None + + def forward( + self, + input_ids: Optional[torch.Tensor], + positions: torch.Tensor, + forward_batch: ForwardBatch, + inputs_embeds: Optional[torch.Tensor] = None, + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> torch.Tensor: + if self.pp_group.is_first_rank: + if inputs_embeds is not None: + hidden_states = inputs_embeds + else: + hidden_states = self.embed_tokens(input_ids) + residual = None + else: + assert pp_proxy_tensors is not None + hidden_states = pp_proxy_tensors["hidden_states"] + residual = pp_proxy_tensors["residual"] + + total_num_layers = self.end_layer - self.start_layer + zero_allocator = BumpAllocator( + buffer_size=total_num_layers * 2, + dtype=torch.float32, + device=hidden_states.device, + ) + + for i in range(self.start_layer, self.end_layer): + hidden_states, residual = self.layers[i]( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + residual=residual, + zero_allocator=zero_allocator, + ) + + if not self.pp_group.is_last_rank: + return PPProxyTensors( + {"hidden_states": hidden_states, "residual": residual} + ) + + if hidden_states.shape[0] != 0: + if residual is None: + hidden_states = self.norm(hidden_states) + else: + hidden_states, _ = self.norm(hidden_states, residual) + return hidden_states + + +class GigaChat35ForCausalLM(DeepseekV2WeightLoaderMixin, nn.Module): + def __init__( + self, + config: GigaChat35Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + nn.Module.__init__(self) + self.config = config + self.quant_config = quant_config + self.pp_group = get_pp_group() + self.tp_size = get_parallel().tp_size + self.num_fused_shared_experts = 0 + + self.model = GigaChat35Model( + config, quant_config, prefix=add_prefix("model", prefix) + ) + if self.pp_group.is_last_rank: + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=add_prefix("lm_head", prefix), + use_attn_tp_group=get_parallel().enable_dp_lm_head, + ) + else: + self.lm_head = None + self.logits_processor = LogitsProcessor(config) + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + inputs_embeds: Optional[torch.Tensor] = None, + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> torch.Tensor: + hidden_states = self.model( + input_ids, positions, forward_batch, inputs_embeds, pp_proxy_tensors + ) + if not self.pp_group.is_last_rank: + return hidden_states + return self.logits_processor( + input_ids, hidden_states, self.lm_head, forward_batch + ) + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_head.weight + + def load_weights( + self, weights: Iterable[tuple[str, torch.Tensor]], is_nextn: bool = False + ): + self.do_load_weights(_remap_gigachat_weight_names(weights), is_nextn=is_nextn) + + def post_load_weights( + self, + is_nextn: bool = False, + weight_names: Optional[Iterable[str]] = None, + ) -> None: + if not is_nextn and weight_names is None: + full_ids = set(self.config.full_attention_layer_ids) + weight_names = [ + f"model.layers.{lid}.self_attn.kv_b_proj.weight" + for lid in full_ids + if self.model.start_layer <= lid < self.model.end_layer + ] + return super().post_load_weights(is_nextn=is_nextn, weight_names=weight_names) + + +EntryClass = [GigaChat35ForCausalLM] diff --git a/python/sglang/srt/models/gigachat35_mtp.py b/python/sglang/srt/models/gigachat35_mtp.py new file mode 100644 index 000000000..77924d7cf --- /dev/null +++ b/python/sglang/srt/models/gigachat35_mtp.py @@ -0,0 +1,214 @@ +"""GigaChat 3.5 multi-head MTP. + +Multiple heads are served Step-3.5 style by sglang's multi-layer EAGLE worker +(``MultiLayerEagleDraftWorker``): one ``GigaChat35ForCausalLMNextN`` instance per +speculative step, selected by ``draft_model_idx``, with hidden-state chaining +(``chain_mtp_hidden_states``) so each head consumes the previous head's output +hidden state instead of always reusing the target model's. + +The block mirrors ``DeepseekModelNextN``'s *shape and naming* (so the checkpoint +weight keys map cleanly through ``DeepseekV2WeightLoaderMixin``) but is built +standalone from the GigaChat modules. +""" + +from __future__ import annotations + +from typing import Iterable, Optional + +import torch +from torch import nn + +from sglang.srt.distributed import get_pp_group +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.models.deepseek_common.deepseek_weight_loader import ( + DeepseekV2WeightLoaderMixin, + NextNDisabledConfig, + NextNEnabledConfig, +) +from sglang.srt.models.gigachat35 import ( + GigaChat35Config, + GigaChat35DecoderLayer, + _remap_gigachat_weight_names, + build_norm, +) +from sglang.srt.runtime_context import get_parallel +from sglang.srt.utils import BumpAllocator, add_prefix + + +class GigaChat35ModelNextN(nn.Module): + def __init__( + self, + config: GigaChat35Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + self.vocab_size = config.vocab_size + + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=add_prefix("embed_tokens", prefix), + ) + + self.enorm = build_norm(config, config.hidden_size) + self.hnorm = build_norm(config, config.hidden_size) + self.eh_proj = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False) + + self.alt_stream = torch.cuda.Stream() if torch.cuda.is_available() else None + + self.decoder = GigaChat35DecoderLayer( + config=config, + layer_id=0, + quant_config=quant_config, + prefix=add_prefix("decoder", prefix), + alt_stream=self.alt_stream, + is_nextn=True, + ) + + self.shared_head = nn.Module() + self.shared_head.norm = build_norm(config, config.hidden_size) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + zero_allocator = BumpAllocator( + buffer_size=2, + dtype=torch.float32, + device=( + input_embeds.device if input_embeds is not None else input_ids.device + ), + ) + + if input_embeds is None: + hidden_states = self.embed_tokens(input_ids) + else: + hidden_states = input_embeds + + if hidden_states.shape[0] > 0: + hidden_states = self.eh_proj( + torch.cat( + ( + self.enorm(hidden_states), + self.hnorm(forward_batch.spec_info.hidden_states), + ), + dim=-1, + ) + ) + + residual = None + hidden_states, residual = self.decoder( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + residual=residual, + zero_allocator=zero_allocator, + ) + + hidden_states_before_norm = None + if not forward_batch.forward_mode.is_idle(): + hidden_states_before_norm = ( + hidden_states if residual is None else hidden_states + residual + ) + if residual is not None: + hidden_states, _ = self.shared_head.norm(hidden_states, residual) + else: + hidden_states = self.shared_head.norm(hidden_states) + + return hidden_states, hidden_states_before_norm + + +class GigaChat35ForCausalLMNextN(DeepseekV2WeightLoaderMixin, nn.Module): + def __init__( + self, + config: GigaChat35Config, + quant_config: Optional[QuantizationConfig] = None, + draft_model_idx: Optional[int] = None, + prefix: str = "", + ) -> None: + nn.Module.__init__(self) + self.config = config + self.quant_config = quant_config + self.pp_group = get_pp_group() + self.tp_size = get_parallel().tp_size + self.num_fused_shared_experts = 0 + self.draft_model_idx = draft_model_idx or 0 + + self.model = GigaChat35ModelNextN( + config, quant_config, prefix=add_prefix("model", prefix) + ) + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=add_prefix("model.shared_head.head", prefix), + use_attn_tp_group=get_parallel().enable_dp_lm_head, + ) + self.logits_processor = LogitsProcessor(config) + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + hidden_states, hidden_states_before_norm = self.model( + input_ids, positions, forward_batch + ) + return self.logits_processor( + input_ids, + hidden_states, + self.lm_head, + forward_batch, + hidden_states_before_norm=hidden_states_before_norm, + ) + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_head.weight + + def set_embed_and_head(self, embed, head): + del self.model.embed_tokens.weight + del self.lm_head.weight + self.model.embed_tokens.weight = embed + self.lm_head.weight = head + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.synchronize() + + def _initialize_nextn_conf(self, is_nextn: bool): + if not is_nextn: + return NextNDisabledConfig() + if not hasattr(self.config, "num_nextn_predict_layers"): + raise ValueError("num_nextn_predict_layers is not in the config") + nextn_layer_id = ( + 0 + if self.config.num_hidden_layers == 1 + else self.config.num_hidden_layers + self.draft_model_idx + ) + return NextNEnabledConfig( + num_nextn_layers=1, + nextn_layer_id=nextn_layer_id, + nextn_layer_prefix=f"model.layers.{nextn_layer_id}", + nextn_spec_weight_names=["shared_head.norm", "eh_proj", "enorm", "hnorm"], + ) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): + self.do_load_weights(_remap_gigachat_weight_names(weights), is_nextn=True) + + def post_load_weights(self, is_nextn: bool = True, weight_names=None) -> None: + super().post_load_weights(is_nextn=True, weight_names=weight_names) + + +EntryClass = [GigaChat35ForCausalLMNextN] diff --git a/python/sglang/srt/parser/reasoning_parser.py b/python/sglang/srt/parser/reasoning_parser.py index 05975e490..e4953935d 100644 --- a/python/sglang/srt/parser/reasoning_parser.py +++ b/python/sglang/srt/parser/reasoning_parser.py @@ -2192,6 +2192,7 @@ class ReasoningParser: "granite_thinking_parser": GraniteThinkingDetector, "interns1": Qwen3Detector, "gemma4": Gemma4Detector, + "gigachat35": DeepSeekR1Detector, "inkling": InklingDetector, "cohere_command4": CohereCommand4Detector, } diff --git a/python/sglang/srt/parser/template_detection.py b/python/sglang/srt/parser/template_detection.py index e6bff2258..2ff6c3506 100644 --- a/python/sglang/srt/parser/template_detection.py +++ b/python/sglang/srt/parser/template_detection.py @@ -522,6 +522,10 @@ def _is_deepseek_r1_think_tags(ctx): return not _is_lfm2(ctx) and (ctx.has_text("") or ctx.has_text("")) +def _is_gigachat35(ctx): + return ctx.has_text("<|GCML|tool_calls>") + + # --------------------------------------------------------------------------- # Reasoning parser rules # --------------------------------------------------------------------------- @@ -556,6 +560,7 @@ REASONING_PARSER_RULES = ( ), DetectionRule(name="deepseek_v4", value="deepseek-v4", predicate=_is_deepseek_v4), DetectionRule(name="deepseek_v3", value="deepseek-v3", predicate=_is_deepseek_v3), + DetectionRule(name="gigachat35", value="gigachat35", predicate=_is_gigachat35), DetectionRule( name="deepseek_r1_force", value="deepseek-r1", predicate=_is_deepseek_r1 ), @@ -571,6 +576,7 @@ REASONING_PARSER_RULES = ( # --------------------------------------------------------------------------- TOOL_CALL_PARSER_RULES = ( + DetectionRule(name="gigachat35", value="gigachat35", predicate=_is_gigachat35), DetectionRule(name="k2_horizon", value="k2_horizon", predicate=_is_k2_v3), DetectionRule(name="apertus2509", value="apertus2509", predicate=_is_apertus2509), DetectionRule(name="gemma4", value="gemma4", predicate=_is_gemma4), diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 08ce08d71..2a1888abd 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -191,6 +191,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): self.chain_mtp_hidden_states = draft_arch in [ "Step3p5MTP", "InklingForConditionalGenerationMTP", + "GigaChat35ForCausalLMNextN", ] self.draft_tp_context = ( draft_tp_context if get_parallel().enable_dp_attention else empty_context @@ -268,7 +269,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): return if isinstance(self.req_to_token_pool, HybridReqToTokenPool): conv_state = self.req_to_token_pool.mamba_pool.mamba_cache.conv - self.draft_extend_num_warmup_tokens = conv_state[0].shape[2] + self.draft_extend_num_warmup_tokens = min(conv_state[0].shape[2:]) self.draft_extend_num_front_tokens = ( self.speculative_num_steps - 1 + self.draft_extend_num_warmup_tokens ) @@ -1057,7 +1058,11 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): ) def forward_batch_generation( - self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None + self, + batch: ScheduleBatch, + on_publish=None, + grammar_barrier=None, + pp_proxy_tensors=None, ): self.draft_worker.last_draft_extend_staged = False if batch.forward_mode.is_extend() or batch.is_extend_in_batch: @@ -1068,7 +1073,9 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): else CaptureHiddenMode.FULL ) batch_output = self.target_worker.forward_batch_generation( - batch, capture_hidden_mode=target_capture_mode + batch, + pp_proxy_tensors=pp_proxy_tensors, + capture_hidden_mode=target_capture_mode, ) # Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens. diff --git a/python/sglang/srt/utils/hf_transformers/common.py b/python/sglang/srt/utils/hf_transformers/common.py index 6cd798913..c6f94f3a4 100644 --- a/python/sglang/srt/utils/hf_transformers/common.py +++ b/python/sglang/srt/utils/hf_transformers/common.py @@ -40,6 +40,7 @@ from sglang.srt.configs import ( ExaoneConfig, FalconH1Config, FalconMambaConfig, + GigaChat35Config, Glm5NextConfig, Glm5NextTextConfig, GraniteMoeHybridConfig, @@ -143,6 +144,7 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = { FalconMambaConfig, Mamba2Config, MambaConfig, + GigaChat35Config, GraniteMoeHybridConfig, HYV4Config, DotsVLMConfig, diff --git a/test/registered/unit/parser/test_template_manager.py b/test/registered/unit/parser/test_template_manager.py index ba5282103..4b21912e6 100644 --- a/test/registered/unit/parser/test_template_manager.py +++ b/test/registered/unit/parser/test_template_manager.py @@ -67,6 +67,22 @@ class TestTemplateManagerReasoningDetection(CustomTestCase): ) return force, config, parser + def test_gigachat35_gcml_template_wins_over_generic_think_tags(self): + """The GCML marker must map both parsers to gigachat35; the template also + carries , so the rule has to outrank the generic deepseek-r1 + think-tags fallback.""" + template = """ + {{- 'assistant<|role_sep|>\\n' -}} + <|GCML|tool_calls> + """ + vocab = ["", "", "<|role_sep|>\n"] + force, config, parser = self._detect(template, vocab) + self.assertEqual(parser, "gigachat35") + self.assertEqual( + detect_tool_call_parser(template, _DummyTokenizer(vocab), config, force), + "gigachat35", + ) + def test_qwen3_template_not_misclassified_as_glm45(self): template = """ {% set enable_thinking = enable_thinking if enable_thinking is defined else true %}