[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:
co-authored by
Stanislav Petrov
Viacheslav Barinov
Viacheslav
Xinyuan Tong
Xinyuan Tong
parent
b54d5b7c7b
commit
b63f8416b3
@@ -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
|
||||
@@ -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",
|
||||
|
||||
@@ -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 = "<|GCML|tool_calls>"
|
||||
_GCML_CLOSE = "</|GCML|tool_calls>"
|
||||
_GCML_INVOKE_RE = re.compile(
|
||||
r"<|GCML|invoke\s+name=\"(?P<name>[^\"]*)\"\s*>"
|
||||
r"(?P<body>.*?)"
|
||||
r"</|GCML|invoke>",
|
||||
re.DOTALL,
|
||||
)
|
||||
_GCML_PARAM_RE = re.compile(
|
||||
r"<|GCML|parameter\s+name=\"(?P<name>[^\"]*)\"\s+"
|
||||
r"string=\"(?P<is_string>true|false)\"\s*>"
|
||||
r"(?P<value>.*?)"
|
||||
r"</|GCML|parameter>",
|
||||
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 <|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."
|
||||
)
|
||||
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
@@ -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("<|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),
|
||||
|
||||
@@ -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>' -}}
|
||||
<|GCML|tool_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 %}
|
||||
|
||||
Reference in New Issue
Block a user