[Feature] Gigachat 3.5 support (#29189)

Co-authored-by: Stanislav Petrov <stapetrov@sberbank.ru>
Co-authored-by: Viacheslav Barinov <vvadbarinov@sberbank.ru>
Co-authored-by: Viacheslav <viacheslav.teh@gmail.com>
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
Stanislav Petrov
2026-09-21 14:37:55 +08:00
committed by GitHub
co-authored by Stanislav Petrov Viacheslav Barinov Viacheslav Xinyuan Tong Xinyuan Tong
parent b54d5b7c7b
commit b63f8416b3
17 changed files with 1366 additions and 7 deletions
@@ -1154,13 +1154,13 @@ Combining `--enable-response-store` with `--disaggregation-mode=prefill` or `dec
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--reasoning-parser`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Specify the parser for reasoning models. Use `auto` to detect the parser from the model's chat template.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`None`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>apertus2509</code>, <code>deepseek-r1</code>, <code>deepseek-v3</code>, <code>deepseek-v4</code>, <code>dots</code>, <code>glm45</code>, <code>ling3</code>, <code>hunyuan</code>, <code>gpt-oss</code>, <code>k2_horizon</code>, <code>kimi</code>, <code>kimi_k2</code>, <code>kimi_k3</code>, <code>mimo</code>, <code>muse</code>, <code>poolside_v1</code>, <code>qwen3</code>, <code>qwen3-thinking</code>, <code>minimax</code>, <code>minimax-append-think</code>, <code>minimax-m3</code>, <code>step3</code>, <code>step3p5</code>, <code>mistral</code>, <code>nemotron_3</code>, <code>interns1</code>, <code>gemma4</code>, <code>inkling</code>, <code>cohere_command4</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>apertus2509</code>, <code>deepseek-r1</code>, <code>deepseek-v3</code>, <code>deepseek-v4</code>, <code>dots</code>, <code>glm45</code>, <code>ling3</code>, <code>hunyuan</code>, <code>gpt-oss</code>, <code>k2_horizon</code>, <code>kimi</code>, <code>kimi_k2</code>, <code>kimi_k3</code>, <code>mimo</code>, <code>muse</code>, <code>poolside_v1</code>, <code>qwen3</code>, <code>qwen3-thinking</code>, <code>minimax</code>, <code>minimax-append-think</code>, <code>minimax-m3</code>, <code>step3</code>, <code>step3p5</code>, <code>mistral</code>, <code>nemotron_3</code>, <code>interns1</code>, <code>gemma4</code>, <code>gigachat35</code>, <code>inkling</code>, <code>cohere_command4</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--tool-call-parser`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Specify the parser for handling tool-call interactions. Use `auto` to detect the parser from the model's chat template.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`None`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>apertus2509</code>, <code>cohere_command4</code>, <code>deepseekv3</code>, <code>deepseekv31</code>, <code>deepseekv32</code>, <code>deepseekv4</code>, <code>dots</code>, <code>glm</code>, <code>glm45</code>, <code>glm47</code>, <code>gpt-oss</code>, <code>k2_horizon</code>, <code>kimi_k2</code>, <code>kimi_k3</code>, <code>lfm2</code>, <code>ling3</code>, <code>llama3</code>, <code>mimo</code>, <code>minicpm5</code>, <code>mistral</code>, <code>muse</code>, <code>poolside_v1</code>, <code>pythonic</code>, <code>qwen</code>, <code>qwen25</code>, <code>qwen3_coder</code>, <code>spark25</code>, <code>step3</code>, <code>step3p5</code>, <code>minimax-m2</code>, <code>minimax-m3</code>, <code>trinity</code>, <code>interns1</code>, <code>hermes</code>, <code>hunyuan</code>, <code>gigachat3</code>, <code>gemma4</code>, <code>inkling</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>apertus2509</code>, <code>cohere_command4</code>, <code>deepseekv3</code>, <code>deepseekv31</code>, <code>deepseekv32</code>, <code>deepseekv4</code>, <code>dots</code>, <code>glm</code>, <code>glm45</code>, <code>glm47</code>, <code>gpt-oss</code>, <code>k2_horizon</code>, <code>kimi_k2</code>, <code>kimi_k3</code>, <code>lfm2</code>, <code>ling3</code>, <code>llama3</code>, <code>mimo</code>, <code>minicpm5</code>, <code>mistral</code>, <code>muse</code>, <code>poolside_v1</code>, <code>pythonic</code>, <code>qwen</code>, <code>qwen25</code>, <code>qwen3_coder</code>, <code>spark25</code>, <code>step3</code>, <code>step3p5</code>, <code>minimax-m2</code>, <code>minimax-m3</code>, <code>trinity</code>, <code>interns1</code>, <code>hermes</code>, <code>hunyuan</code>, <code>gigachat3</code>, <code>gigachat35</code>, <code>gemma4</code>, <code>inkling</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--tool-server`</td>
@@ -308,5 +308,10 @@ in the GitHub search bar.
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>JetBrains/Mellum2-12B-A2.5B-Base</code>, <code>JetBrains/Mellum2-12B-A2.5B-Thinking</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>JetBrains' Qwen3-MoE-based code generation model with interleaved sliding-window/full attention, per-layer-type RoPE, and per-layer dense/sparse MLP routing.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong>GigaChat 3.5</strong></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>ai-sage/GigaChat3.5-432B-A28B</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>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.</td>
</tr>
</tbody>
</table>
@@ -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
@@ -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
+2
View File
@@ -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",
+252
View File
@@ -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,
)
)
@@ -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
@@ -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,
}
@@ -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 = "<GCMLtool_calls>"
_GCML_CLOSE = "</GCMLtool_calls>"
_GCML_INVOKE_RE = re.compile(
r"<GCMLinvoke\s+name=\"(?P<name>[^\"]*)\"\s*>"
r"(?P<body>.*?)"
r"</GCMLinvoke>",
re.DOTALL,
)
_GCML_PARAM_RE = re.compile(
r"<GCMLparameter\s+name=\"(?P<name>[^\"]*)\"\s+"
r"string=\"(?P<is_string>true|false)\"\s*>"
r"(?P<value>.*?)"
r"</GCMLparameter>",
re.DOTALL,
)
_TRAILING_MARKER_RE = re.compile(r"(?:<\|message_sep\|>|</s>)+\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 <GCMLparameter ...> 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."
)
@@ -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:
+668
View File
@@ -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]
+214
View File
@@ -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]
@@ -2192,6 +2192,7 @@ class ReasoningParser:
"granite_thinking_parser": GraniteThinkingDetector,
"interns1": Qwen3Detector,
"gemma4": Gemma4Detector,
"gigachat35": DeepSeekR1Detector,
"inkling": InklingDetector,
"cohere_command4": CohereCommand4Detector,
}
@@ -522,6 +522,10 @@ def _is_deepseek_r1_think_tags(ctx):
return not _is_lfm2(ctx) and (ctx.has_text("<think>") or ctx.has_text("</think>"))
def _is_gigachat35(ctx):
return ctx.has_text("<GCMLtool_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),
@@ -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.
@@ -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,
@@ -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 <think>, so the rule has to outrank the generic deepseek-r1
think-tags fallback."""
template = """
{{- 'assistant<|role_sep|>\\n<think>' -}}
<GCMLtool_calls>
"""
vocab = ["<think>", "</think>", "<|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 %}