diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py index 53fecdb38..8f86dcbef 100644 --- a/python/sglang/srt/configs/__init__.py +++ b/python/sglang/srt/configs/__init__.py @@ -57,6 +57,7 @@ from sglang.srt.configs.qwen3_5 import ( ) from sglang.srt.configs.qwen3_asr import Qwen3ASRConfig from sglang.srt.configs.qwen3_next import Qwen3NextConfig +from sglang.srt.configs.spark3 import Spark3Config from sglang.srt.configs.step3_vl import ( Step3TextConfig, Step3VisionEncoderConfig, @@ -117,6 +118,7 @@ __all__ = [ "MiniCPMHybridConfig", "Step3p5Config", "MiniMaxM3VLConfig", + "Spark3Config", "Step3p7Config", "Qwen3ASRConfig", "InklingAudioConfig", diff --git a/python/sglang/srt/configs/spark3.py b/python/sglang/srt/configs/spark3.py new file mode 100644 index 000000000..65842dc9f --- /dev/null +++ b/python/sglang/srt/configs/spark3.py @@ -0,0 +1,63 @@ +from typing import Any, Optional + +from transformers.configuration_utils import PretrainedConfig + + +class Spark3Config(PretrainedConfig): + model_type = "spark3" + architectures = ["Spark3ForCausalLM"] + + def __init__( + self, + hidden_size: int = 2048, + intermediate_size: int = 6656, + num_attention_heads: int = 8, + num_key_value_heads: int = 2, + num_hidden_layers: int = 28, + head_dim: int = 256, + headwise_attn_output_gate: bool = True, + sliding_window: int = 512, + vocab_size: int = 133120, + rms_norm_eps: float = 1e-6, + max_position_embeddings: int = 8192, + rope_parameters: Optional[dict[str, Any]] = None, + layer_types: list[str] = None, + tie_word_embeddings: Optional[bool] = None, + **kwargs, + ) -> None: + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.num_attention_heads = num_attention_heads + self.num_key_value_heads = num_key_value_heads + self.num_hidden_layers = num_hidden_layers + self.head_dim = head_dim + self.headwise_attn_output_gate = headwise_attn_output_gate + self.sliding_window = sliding_window + self.vocab_size = vocab_size + self.rms_norm_eps = rms_norm_eps + self.max_position_embeddings = max_position_embeddings + + if layer_types is not None: + layer_types = layer_types[: self.num_hidden_layers] + else: + layer_types = [ + "sliding_attention" if bool((i + 1) % 4) else "full_attention" + for i in range(self.num_hidden_layers) + ] + self.layer_types = layer_types + + if rope_parameters is not None: + self.rope_parameters = rope_parameters + else: + self.rope_parameters = { + "full_attention": { + "rope_theta": 5000000, + "partial_rotary_factor": 0.25, + }, + "sliding_attention": { + "rope_theta": 10000, + "partial_rotary_factor": 1.0, + }, + } + + super().__init__(**kwargs, tie_word_embeddings=tie_word_embeddings) diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 1530227b8..529c6e876 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -43,6 +43,7 @@ from sglang.srt.function_call.poolside_v1_detector import PoolsideV1Detector from sglang.srt.function_call.pythonic_detector import PythonicDetector from sglang.srt.function_call.qwen3_coder_detector import Qwen3CoderDetector from sglang.srt.function_call.qwen25_detector import Qwen25Detector +from sglang.srt.function_call.spark3_detector import Spark3Detector from sglang.srt.function_call.step3_detector import Step3Detector from sglang.srt.function_call.trinity_detector import TrinityDetector from sglang.srt.function_call.utils import ( @@ -87,6 +88,7 @@ class FunctionCallParser: "qwen": Qwen25Detector, "qwen25": Qwen25Detector, "qwen3_coder": Qwen3CoderDetector, + "spark": Spark3Detector, "step3": Step3Detector, "step3p5": Qwen3CoderDetector, "minimax-m2": MinimaxM2Detector, diff --git a/python/sglang/srt/function_call/spark3_detector.py b/python/sglang/srt/function_call/spark3_detector.py new file mode 100644 index 000000000..17fea5b3c --- /dev/null +++ b/python/sglang/srt/function_call/spark3_detector.py @@ -0,0 +1,263 @@ +import json +import re +from dataclasses import dataclass +from typing import Any + +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, +) + +TOOL_CALL_BEGIN = "" +TOOL_CALL_END = "" +ARG_KEY_BEGIN = "" +ARG_KEY_END = "" +ARG_VALUE_BEGIN = "" +ARG_VALUE_END = "" + +ARG_PAIR_PATTERN = re.compile( + rf"{re.escape(ARG_KEY_BEGIN)}(.*?){re.escape(ARG_KEY_END)}" + rf"{re.escape(ARG_VALUE_BEGIN)}(.*?){re.escape(ARG_VALUE_END)}", + re.DOTALL, +) + + +@dataclass(frozen=True) +class _Spark3ToolCall: + name: str + arguments: dict[str, Any] + + def arguments_json(self) -> str: + return json.dumps( + self.arguments, + ensure_ascii=False, + separators=(",", ":"), + ) + + +def _get_param_type(tools: list[Tool], function_name: str, param_name: str) -> str: + """Return a parameter's declared JSON Schema type, or ``string``.""" + for tool in tools: + function = getattr(tool, "function", None) + if function is None or function.name != function_name: + continue + parameters = getattr(function, "parameters", None) + if not isinstance(parameters, dict): + continue + properties = parameters.get("properties") + if not isinstance(properties, dict): + continue + definition = properties.get(param_name) + if isinstance(definition, dict) and isinstance(definition.get("type"), str): + return definition["type"] + return "string" + + +def _convert_value(value: str, param_type: str) -> Any: + """Convert Spark3 XML text according to the model's tool protocol.""" + if value.lower() == "null": + return None + + normalized_type = param_type.lower() + try: + if normalized_type in {"string", "str", "text"}: + return value + if normalized_type in {"integer", "int"}: + return int(value) + if normalized_type in {"number", "float"}: + number = float(value) + return int(number) if number.is_integer() else number + if normalized_type in {"boolean", "bool"}: + normalized_value = value.strip().lower() + if normalized_value not in {"true", "1", "false", "0"}: + raise ValueError(f"invalid boolean: {value}") + return normalized_value in {"true", "1"} + return json.loads(value) + except (TypeError, ValueError, json.JSONDecodeError): + try: + return json.loads(value) + except (TypeError, ValueError, json.JSONDecodeError): + return value + + +def _parse_tool_call_xml(tool_xml: str, tools: list[Tool]) -> _Spark3ToolCall | None: + if not tool_xml.startswith(TOOL_CALL_BEGIN) or not tool_xml.endswith(TOOL_CALL_END): + return None + + body = tool_xml[len(TOOL_CALL_BEGIN) : -len(TOOL_CALL_END)] + first_arg = body.find(ARG_KEY_BEGIN) + function_name = (body if first_arg < 0 else body[:first_arg]).strip() + if not function_name: + return None + + arguments: dict[str, Any] = {} + for match in ARG_PAIR_PATTERN.finditer(body): + key, raw_value = match.group(1), match.group(2) + if not key: + continue + arguments[key] = _convert_value( + raw_value, + _get_param_type(tools, function_name, key), + ) + return _Spark3ToolCall(name=function_name, arguments=arguments) + + +def _partial_marker_suffix_length(text: str, marker: str) -> int: + """Length of the suffix that may become ``marker`` in the next chunk.""" + for size in range(min(len(text), len(marker) - 1), 0, -1): + if text.endswith(marker[:size]): + return size + return 0 + + +class Spark3Detector(BaseFormatDetector): + """Detector for Spark3's XML-KV tool-call format. + + Wire format:: + + function_name + keyvalue + + + Values are converted with the parameter's JSON Schema type. A complete + block is emitted atomically in streaming mode so XML fragments are never + exposed as JSON argument deltas. + """ + + def __init__(self): + super().__init__() + self.bot_token = TOOL_CALL_BEGIN + self.eot_token = TOOL_CALL_END + + def has_tool_call(self, text: str) -> bool: + return TOOL_CALL_BEGIN in text + + def _build_item( + self, + parsed: _Spark3ToolCall, + tools: list[Tool], + tool_index: int, + ) -> ToolCallItem | None: + validated = self.parse_base_json( + {"name": parsed.name, "arguments": parsed.arguments}, tools + ) + if not validated: + return None + return ToolCallItem( + tool_index=tool_index, + name=parsed.name, + parameters=parsed.arguments_json(), + ) + + def _record_streamed_item( + self, parsed: _Spark3ToolCall, item: ToolCallItem + ) -> None: + self.prev_tool_call_arr.append( + {"name": parsed.name, "arguments": parsed.arguments} + ) + self.streamed_args_for_tool.append(item.parameters) + + def detect_and_parse(self, text: str, tools: list[Tool]) -> StreamingParseResult: + calls: list[ToolCallItem] = [] + normal_parts: list[str] = [] + cursor = 0 + + while cursor < len(text): + start = text.find(TOOL_CALL_BEGIN, cursor) + if start < 0: + normal_parts.append(text[cursor:]) + break + + normal_parts.append(text[cursor:start]) + end = text.find(TOOL_CALL_END, start + len(TOOL_CALL_BEGIN)) + if end < 0: + normal_parts.append(text[start:]) + break + + end += len(TOOL_CALL_END) + raw_tool_call = text[start:end] + parsed = _parse_tool_call_xml(raw_tool_call, tools) + if parsed is None: + normal_parts.append(raw_tool_call) + else: + item = self._build_item(parsed, tools, len(calls)) + if item is not None: + calls.append(item) + cursor = end + + return StreamingParseResult( + normal_text="".join(normal_parts), + calls=calls, + ) + + def parse_streaming_increment( + self, new_text: str, tools: list[Tool] + ) -> StreamingParseResult: + self._buffer += new_text + calls: list[ToolCallItem] = [] + normal_parts: list[str] = [] + + while self._buffer: + start = self._buffer.find(TOOL_CALL_BEGIN) + if start < 0: + keep = _partial_marker_suffix_length(self._buffer, TOOL_CALL_BEGIN) + if keep: + normal_parts.append(self._buffer[:-keep]) + self._buffer = self._buffer[-keep:] + else: + normal_parts.append(self._buffer) + self._buffer = "" + break + + if start > 0: + normal_parts.append(self._buffer[:start]) + self._buffer = self._buffer[start:] + + end = self._buffer.find(TOOL_CALL_END, len(TOOL_CALL_BEGIN)) + if end < 0: + break + + end += len(TOOL_CALL_END) + raw_tool_call = self._buffer[:end] + self._buffer = self._buffer[end:] + parsed = _parse_tool_call_xml(raw_tool_call, tools) + if parsed is None: + normal_parts.append(raw_tool_call) + continue + + item = self._build_item(parsed, tools, self.current_tool_id + 1) + if item is not None: + self.current_tool_id += 1 + self._record_streamed_item(parsed, item) + calls.append(item) + + return StreamingParseResult( + normal_text="".join(normal_parts), + calls=calls, + ) + + def finish(self, tools: list[Tool]) -> StreamingParseResult: + del tools + pending = self._buffer + self._buffer = "" + if TOOL_CALL_BEGIN in pending: + # A complete opening marker means this is a truncated protocol + # block, not user-visible text. Partial marker prefixes are still + # released because the stream has ended and they cannot become a + # tool call anymore. + pending = pending[: pending.find(TOOL_CALL_BEGIN)] + return StreamingParseResult(normal_text=pending) + + def supports_structural_tag(self) -> bool: + return False + + def parses_required_natively(self) -> bool: + return True + + def structure_info(self) -> _GetInfoFunc: + raise NotImplementedError( + "Spark3 XML arguments cannot be represented by legacy structural tags" + ) diff --git a/python/sglang/srt/models/spark3.py b/python/sglang/srt/models/spark3.py new file mode 100644 index 000000000..f5fd231b3 --- /dev/null +++ b/python/sglang/srt/models/spark3.py @@ -0,0 +1,564 @@ +import logging +from typing import Iterable, Optional, Tuple, Union + +import torch +from torch import nn + +from sglang.srt.distributed import get_pp_group +from sglang.srt.layers.activation import GeluAndMul +from sglang.srt.layers.dp_attention import is_dp_attention_enabled +from sglang.srt.layers.layernorm import RMSNorm +from sglang.srt.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + QKVParallelLinear, + RowParallelLinear, +) +from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.quantization.base_config import QuantizationConfig +from sglang.srt.layers.radix_attention import RadixAttention +from sglang.srt.layers.rotary_embedding import get_rope +from sglang.srt.layers.utils import PPMissingLayer, get_layer_id +from sglang.srt.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) +from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors +from sglang.srt.model_loader.weight_utils import ( + default_weight_loader, +) +from sglang.srt.platforms import current_platform +from sglang.srt.runtime_context import get_parallel +from sglang.srt.utils import add_prefix, make_layers + +Spark3Config = None + +logger = logging.getLogger(__name__) + + +# Aligned with HF's implementation, using sliding window inclusive with the last token +# SGLang assumes exclusive +def _get_attention_sliding_window_size(config): + return config.sliding_window - 1 + + +class Spark3MLP(nn.Module): + def __init__( + self, + hidden_size: int, + intermediate_size: int, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + reduce_results: bool = True, + ) -> None: + super().__init__() + self.gate_up_proj = MergedColumnParallelLinear( + hidden_size, + [intermediate_size] * 2, + bias=False, + quant_config=quant_config, + prefix=add_prefix("gate_up_proj", prefix), + ) + self.down_proj = RowParallelLinear( + intermediate_size, + hidden_size, + bias=False, + quant_config=quant_config, + prefix=add_prefix("down_proj", prefix), + reduce_results=reduce_results, + ) + self.act_fn = GeluAndMul() + + def forward( + self, + x, + forward_batch=None, + ): + gate_up, _ = self.gate_up_proj(x) + x = self.act_fn(gate_up) + x, _ = self.down_proj(x) + return x + + +class Spark3Attention(nn.Module): + def __init__( + self, + hidden_size: int, + num_heads: int, + num_kv_heads: int, + head_dim: Optional[int] = None, + layer_id: int = 0, + rope_theta: float = 10000, + partial_rotary_factor: float = 1.0, + max_position_embeddings: int = 8192, + quant_config: Optional[QuantizationConfig] = None, + sliding_window: int = 512, + layer_type: str = "sliding_attention", + headwise_attn_output_gate: bool = True, + prefix: str = "", + ) -> None: + super().__init__() + self.hidden_size = hidden_size + self.total_num_heads = num_heads + attn_tp_rank = get_parallel().attn_tp_rank + attn_tp_size = get_parallel().attn_tp_size + + assert self.total_num_heads % attn_tp_size == 0 + self.num_heads = self.total_num_heads // attn_tp_size + self.total_num_kv_heads = num_kv_heads + if self.total_num_kv_heads >= attn_tp_size: + # Number of KV heads is greater than TP size, so we partition + # the KV heads across multiple tensor parallel GPUs. + assert self.total_num_kv_heads % attn_tp_size == 0 + else: + # Number of KV heads is less than TP size, so we replicate + # the KV heads across multiple tensor parallel GPUs. + assert attn_tp_size % self.total_num_kv_heads == 0 + self.num_kv_heads = max(1, self.total_num_kv_heads // attn_tp_size) + if head_dim is not None: + self.head_dim = head_dim + else: + self.head_dim = hidden_size // self.total_num_heads + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_kv_heads * self.head_dim + self.scaling = self.head_dim**-0.5 + self.rope_theta = rope_theta + self.max_position_embeddings = max_position_embeddings + self.partial_rotary_factor = partial_rotary_factor + self.headwise_attn_output_gate = headwise_attn_output_gate + + self.q_k_v_proj = QKVParallelLinear( + hidden_size, + self.head_dim, + self.total_num_heads, + self.total_num_kv_heads, + bias=False, + quant_config=quant_config, + tp_rank=attn_tp_rank, + tp_size=attn_tp_size, + prefix=add_prefix("q_k_v_proj", prefix), + ) + if self.headwise_attn_output_gate: + self.g_proj = ColumnParallelLinear( + hidden_size, + self.total_num_heads, + bias=False, + quant_config=None, # g_proj keeps bf16. + tp_rank=attn_tp_rank, + tp_size=attn_tp_size, + prefix=add_prefix("g_proj", prefix), + ) + + self.out_proj = RowParallelLinear( + self.total_num_heads * self.head_dim, + hidden_size, + bias=False, + quant_config=quant_config, + tp_rank=attn_tp_rank, + tp_size=attn_tp_size, + prefix=add_prefix("out_proj", prefix), + ) + + self.rotary_emb = get_rope( + self.head_dim, + rotary_dim=self.head_dim, + max_position=max_position_embeddings, + base=rope_theta, + rope_scaling=None, + partial_rotary_factor=partial_rotary_factor, + is_neox_style=True, + ) + if layer_type not in ("sliding_attention", "full_attention"): + raise ValueError(f"Unsupported Spark3 layer_type: {layer_type}") + sliding_window_size = ( + sliding_window if layer_type == "sliding_attention" else -1 + ) + self.attn = RadixAttention( + self.num_heads, + self.head_dim, + self.scaling, + num_kv_heads=self.num_kv_heads, + sliding_window_size=sliding_window_size, + layer_id=layer_id, + quant_config=quant_config, + prefix=add_prefix("attn", prefix), + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + qkv, _ = self.q_k_v_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + q, k = self.rotary_emb(positions, q, k) + attn_output = self.attn(q, k, v, forward_batch) + + if self.headwise_attn_output_gate: + g, _ = self.g_proj(hidden_states) + g = torch.sigmoid(g.float()).to(attn_output.dtype) + gate_output = attn_output.view( + attn_output.shape[0], + self.num_heads, + self.head_dim, + ) * g.unsqueeze(-1) + attn_output = gate_output.view(*attn_output.shape) + + output, _ = self.out_proj(attn_output) + return output + + +class Spark3DecoderLayer(nn.Module): + """A single transformer layer. + + Transformer layer takes input with size [s, b, h] and returns an + output of the same size. + """ + + def __init__( + self, + config: Spark3Config, + layer_id: int = 0, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + + max_position_embeddings = getattr(config, "max_position_embeddings", 8192) + head_dim = getattr(config, "head_dim", None) + layer_type = config.layer_types[layer_id] + rope_params = config.rope_parameters[layer_type] + rope_theta = rope_params.get("rope_theta") + partial_rotary_factor = rope_params.get("partial_rotary_factor") + if partial_rotary_factor is None: + partial_rotary_factor = 1.0 + + self.self_attn = Spark3Attention( + hidden_size=config.hidden_size, + num_heads=config.num_attention_heads, + num_kv_heads=config.num_key_value_heads, + head_dim=head_dim, + layer_id=layer_id, + rope_theta=rope_theta, + partial_rotary_factor=partial_rotary_factor, + max_position_embeddings=max_position_embeddings, + quant_config=quant_config, + sliding_window=_get_attention_sliding_window_size(config), + layer_type=layer_type, + headwise_attn_output_gate=getattr( + config, "headwise_attn_output_gate", True + ), + prefix=add_prefix("self_attn", prefix), + ) + + # MLP + self.mlp = Spark3MLP( + config.hidden_size, + intermediate_size=config.intermediate_size, + quant_config=quant_config, + prefix=add_prefix("mlp", prefix), + ) + + self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + residual: Optional[torch.Tensor], + ) -> Tuple[torch.Tensor, torch.Tensor]: + # Self Attention + if residual is None: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + else: + hidden_states, residual = self.input_layernorm(hidden_states, residual) + + hidden_states = self.self_attn( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + ) + hidden_states, residual = self.post_attention_layernorm(hidden_states, residual) + hidden_states = self.mlp(hidden_states) + + return hidden_states, residual + + +class Spark3Model(nn.Module): + def __init__( + self, + config: Spark3Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + 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, + quant_config=quant_config, + use_attn_tp_group=is_dp_attention_enabled(), + prefix=add_prefix("embed_tokens", prefix), + ) + else: + self.embed_tokens = PPMissingLayer() + + self.layers, self.start_layer, self.end_layer = make_layers( + config.num_hidden_layers, + lambda idx, prefix: Spark3DecoderLayer( + layer_id=idx, + config=config, + quant_config=quant_config, + prefix=prefix, + ), + 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 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + else: + self.norm = PPMissingLayer(return_tuple=True) + + def get_input_embeddings(self) -> nn.Embedding: + return self.embed_tokens + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: torch.Tensor = None, + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> Union[torch.Tensor, PPProxyTensors]: + if self.pp_group.is_first_rank: + if input_embeds is None: + hidden_states = self.embed_tokens(input_ids) + else: + hidden_states = input_embeds + residual = None + else: + assert pp_proxy_tensors is not None + hidden_states = pp_proxy_tensors["hidden_states"] + residual = pp_proxy_tensors["residual"] + + for i in range(self.start_layer, self.end_layer): + layer = self.layers[i] + hidden_states, residual = layer( + positions, + hidden_states, + forward_batch, + residual, + ) + if not self.pp_group.is_last_rank: + return PPProxyTensors( + { + "hidden_states": hidden_states, + "residual": residual, + } + ) + else: + 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 Spark3ForCausalLM(nn.Module): + def __init__( + self, + config: Spark3Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.pp_group = get_pp_group() + self.config = config + self.quant_config = quant_config + self.model = Spark3Model( + config, quant_config=quant_config, prefix=add_prefix("model", prefix) + ) + + # handle the lm head on different pp ranks + if self.pp_group.is_last_rank: + if self.pp_group.world_size == 1 and config.tie_word_embeddings: + self.lm_head = self.model.embed_tokens + else: + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=add_prefix("lm_head", prefix), + ) + else: + # ranks other than the last rank will have a placeholder layer + self.lm_head = PPMissingLayer() + + self.logits_processor = LogitsProcessor(config) + + def get_input_embeddings(self) -> nn.Embedding: + return self.model.embed_tokens + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: torch.Tensor = None, + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> torch.Tensor: + hidden_states = self.model( + input_ids, + positions, + forward_batch, + input_embeds, + pp_proxy_tensors=pp_proxy_tensors, + ) + + if self.pp_group.is_last_rank: + return self.logits_processor( + input_ids, + hidden_states, + self.lm_head, + forward_batch, + ) + else: + return hidden_states + + @torch.no_grad() + def forward_split_prefill( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + split_interval: Tuple[int, int], # [start, end) 0-based + input_embeds: torch.Tensor = None, + ): + start, end = split_interval + # embed + if start == 0: + if input_embeds is None: + forward_batch.hidden_states = self.model.embed_tokens(input_ids) + else: + forward_batch.hidden_states = input_embeds + # decoder layer + for i in range(start, end): + layer = self.model.layers[i] + forward_batch.hidden_states, forward_batch.residual = layer( + positions, + forward_batch.hidden_states, + forward_batch, + forward_batch.residual, + ) + + if end == self.model.config.num_hidden_layers: + # norm + hidden_states, _ = self.model.norm( + forward_batch.hidden_states, forward_batch.residual + ) + forward_batch.hidden_states = hidden_states + # logits process + result = self.logits_processor( + input_ids, forward_batch.hidden_states, self.lm_head, forward_batch + ) + else: + result = None + + return result + + @property + def start_layer(self): + return self.model.start_layer + + @property + def end_layer(self): + return self.model.end_layer + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + stacked_params_mapping = [ + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + params_dict = dict(self.named_parameters()) + for name, loaded_weight in weights: + original_name = name + layer_id = get_layer_id(name) + if ( + layer_id is not None + and hasattr(self.model, "start_layer") + and ( + layer_id < self.model.start_layer + or layer_id >= self.model.end_layer + ) + ): + continue + + if self.config.tie_word_embeddings and "lm_head.weight" in name: + continue + + if name in ("model.embedding.weight",): + name = "model.embed_tokens.weight" + if ( + name == "model.embed_tokens.weight" + and self.config.tie_word_embeddings + and self.pp_group.world_size > 1 + ): + if self.pp_group.is_last_rank: + name = "lm_head.weight" + elif not self.pp_group.is_first_rank: + continue + + loaded = False + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + mapped_name = name.replace(weight_name, param_name) + if mapped_name not in params_dict: + continue + param = params_dict[mapped_name] + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight, shard_id) + loaded = True + break + if loaded: + continue + + if name in params_dict: + param = params_dict[name] + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight) + elif original_name in ("model.embedding.weight",): + continue + else: + logger.warning(f"Parameter {name} not found in params_dict") + + 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 + current_platform.empty_cache() + current_platform.synchronize() + + def get_attention_sliding_window_size(self) -> int: + return _get_attention_sliding_window_size(self.config) + + +EntryClass = [Spark3ForCausalLM] diff --git a/python/sglang/srt/utils/hf_transformers/common.py b/python/sglang/srt/utils/hf_transformers/common.py index c5f346786..4606733fb 100644 --- a/python/sglang/srt/utils/hf_transformers/common.py +++ b/python/sglang/srt/utils/hf_transformers/common.py @@ -66,6 +66,7 @@ from sglang.srt.configs import ( Qwen3_5MoeTextConfig, Qwen3_5TextConfig, Qwen3NextConfig, + Spark3Config, Step3p5Config, Step3p7Config, Step3VLConfig, @@ -102,6 +103,7 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = { LocateAnythingConfig, InternVLChatConfig, LagunaConfig, + Spark3Config, Step3VLConfig, LongcatFlashConfig, Olmo3Config, diff --git a/test/registered/unit/function_call/test_spark3_detector.py b/test/registered/unit/function_call/test_spark3_detector.py new file mode 100644 index 000000000..950e2d7fe --- /dev/null +++ b/test/registered/unit/function_call/test_spark3_detector.py @@ -0,0 +1,202 @@ +"""Unit tests for Spark3Detector - no server, no model loading.""" + +import json +import unittest + +from sglang.srt.entrypoints.openai.protocol import Function, Tool +from sglang.srt.environ import envs +from sglang.srt.function_call.function_call_parser import FunctionCallParser +from sglang.srt.function_call.spark3_detector import Spark3Detector +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +def _xml(name: str, arguments: list[tuple[str, str]]) -> str: + pairs = "".join( + f"{key}{value}" + for key, value in arguments + ) + return f"{name}{pairs}" + + +def _tools(): + return [ + Tool( + type="function", + function=Function( + name="set_state", + parameters={ + "type": "object", + "properties": { + "name": {"type": "string"}, + "count": {"type": "integer"}, + "ratio": {"type": "number"}, + "active": {"type": "boolean"}, + "items": {"type": "array"}, + "metadata": {"type": "object"}, + }, + }, + ), + ), + Tool( + type="function", + function=Function( + name="now", + parameters={"type": "object", "properties": {}}, + ), + ), + ] + + +class TestSpark3DetectorDetectAndParse(CustomTestCase): + def setUp(self): + self.tools = _tools() + self.detector = Spark3Detector() + + def test_spark3_parser_is_registered(self): + self.assertIs(FunctionCallParser.ToolCallParserEnum["spark"], Spark3Detector) + + def test_nonstream_parses_multiple_calls_and_preserves_normal_text(self): + text = ( + "before" + + _xml( + "set_state", + [ + ("name", "上海"), + ("count", "42"), + ("ratio", "2.5"), + ("active", "1"), + ("items", '["a", "b"]'), + ("metadata", '{"source":"spark"}'), + ], + ) + + "middle" + + _xml("now", []) + + "after" + ) + + result = self.detector.detect_and_parse(text, self.tools) + + self.assertEqual(result.normal_text, "beforemiddleafter") + self.assertEqual([call.tool_index for call in result.calls], [0, 1]) + self.assertEqual([call.name for call in result.calls], ["set_state", "now"]) + self.assertEqual( + json.loads(result.calls[0].parameters), + { + "name": "上海", + "count": 42, + "ratio": 2.5, + "active": True, + "items": ["a", "b"], + "metadata": {"source": "spark"}, + }, + ) + self.assertEqual(json.loads(result.calls[1].parameters), {}) + + def test_null_and_conversion_fallbacks_match_spark3_protocol(self): + text = _xml( + "set_state", + [ + ("name", "null"), + ("count", "not-an-int"), + ("active", "false"), + ("undeclared", "42"), + ], + ) + + result = self.detector.detect_and_parse(text, self.tools) + + self.assertEqual( + json.loads(result.calls[0].parameters), + { + "name": None, + "count": "not-an-int", + "active": False, + "undeclared": "42", + }, + ) + + def test_malformed_block_is_text_and_unknown_tool_honors_policy(self): + malformed = "xy" + result = self.detector.detect_and_parse(malformed, self.tools) + self.assertEqual(result.normal_text, malformed) + self.assertEqual(result.calls, []) + + unknown = _xml("missing", [("value", "1")]) + with envs.SGLANG_FORWARD_UNKNOWN_TOOLS.override(False): + result = self.detector.detect_and_parse(unknown, self.tools) + self.assertEqual(result.normal_text, "") + self.assertEqual(result.calls, []) + with envs.SGLANG_FORWARD_UNKNOWN_TOOLS.override(True): + result = self.detector.detect_and_parse(unknown, self.tools) + self.assertEqual(result.calls[0].name, "missing") + self.assertEqual(json.loads(result.calls[0].parameters), {"value": "1"}) + + def test_stream_end_flushes_partial_marker_and_required_stays_native(self): + result = self.detector.parse_streaming_increment("plainset_statecount", self.tools + ) + self.assertEqual(result.normal_text, "plain") + self.assertEqual(truncated.finish(self.tools).normal_text, "") + + self.assertFalse(self.detector.supports_structural_tag()) + self.assertTrue(self.detector.parses_required_natively()) + self.assertIs( + FunctionCallParser(self.tools, "spark").get_structure_constraint( + "required" + ), + None, + ) + + +class TestSpark3DetectorStreaming(CustomTestCase): + def setUp(self): + self.tools = _tools() + + def test_streaming_character_chunks_match_nonstream_result(self): + text = ( + "answer:" + + _xml("set_state", [("count", "42"), ("active", "0")]) + + _xml("now", []) + + "done" + ) + detector = Spark3Detector() + normal_parts = [] + calls = [] + + for character in text: + result = detector.parse_streaming_increment(character, self.tools) + normal_parts.append(result.normal_text) + calls.extend(result.calls) + end = detector.finish(self.tools) + normal_parts.append(end.normal_text) + calls.extend(end.calls) + + self.assertEqual("".join(normal_parts), "answer:done") + self.assertEqual([call.tool_index for call in calls], [0, 1]) + self.assertEqual( + json.loads(calls[0].parameters), {"count": 42, "active": False} + ) + self.assertEqual(json.loads(calls[1].parameters), {}) + self.assertEqual( + detector.prev_tool_call_arr, + [ + {"name": "set_state", "arguments": {"count": 42, "active": False}}, + {"name": "now", "arguments": {}}, + ], + ) + self.assertEqual( + detector.streamed_args_for_tool, + ['{"count":42,"active":false}', "{}"], + ) + + +if __name__ == "__main__": + unittest.main()