diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py index 043de321b..865131927 100644 --- a/python/sglang/srt/configs/__init__.py +++ b/python/sglang/srt/configs/__init__.py @@ -24,6 +24,7 @@ from sglang.srt.configs.step3_vl import ( Step3VisionEncoderConfig, Step3VLConfig, ) +from sglang.srt.configs.step3p5 import Step3p5Config __all__ = [ "AfmoeConfig", @@ -50,4 +51,5 @@ __all__ = [ "NemotronH_Nano_VL_V2_Config", "JetNemotronConfig", "JetVLMConfig", + "Step3p5Config", ] diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 43543b655..64b9f6fd3 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -302,6 +302,8 @@ class ModelConfig: and self.hf_config.architectures[0] == "MiMoV2FlashForCausalLM" ): self.hf_config.architectures[0] = "MiMoV2MTP" + if is_draft_model and self.hf_config.architectures[0] == "Step3p5ForCausalLM": + self.hf_config.architectures[0] = "Step3p5MTP" if is_draft_model and self.hf_config.architectures[0] in [ "BailingMoeV2ForCausalLM", "BailingMoeForCausalLM", @@ -606,6 +608,11 @@ class ModelConfig: if hasattr(self.hf_text_config, "swa_num_key_value_heads"): total_num_kv_heads = self.hf_text_config.swa_num_key_value_heads return max(1, total_num_kv_heads // tensor_parallel_size) + elif hasattr(self.hf_text_config, "attention_other_setting"): # For step3p5 + total_num_kv_heads = self.hf_text_config.attention_other_setting.get( + "num_attention_groups" + ) + return max(1, total_num_kv_heads // tensor_parallel_size) else: return self.get_num_kv_heads(tensor_parallel_size) @@ -1268,6 +1275,8 @@ def is_hybrid_swa_model(model_architectures: List[str]): "GptOssForCausalLM", "MiMoV2FlashForCausalLM", "MiMoV2MTP", + "Step3p5ForCausalLM", + "Step3p5MTP", } return any(arch in hybrid_swa_archs for arch in model_architectures) @@ -1303,6 +1312,21 @@ def get_hybrid_layer_ids( elif "MiMoV2MTP" in model_architectures: swa_attention_layer_ids = [0] full_attention_layer_ids = [] + elif "Step3p5ForCausalLM" in model_architectures: + layer_types = hf_text_config.layer_types + swa_attention_layer_ids = [ + i + for i, x in enumerate(layer_types) + if x == "sliding_attention" and i < num_hidden_layers + ] + full_attention_layer_ids = [ + i + for i, x in enumerate(layer_types) + if x == "full_attention" and i < num_hidden_layers + ] + elif "Step3p5MTP" in model_architectures: + swa_attention_layer_ids = [0] + full_attention_layer_ids = [] else: swa_attention_layer_ids = None full_attention_layer_ids = None diff --git a/python/sglang/srt/configs/step3p5.py b/python/sglang/srt/configs/step3p5.py new file mode 100644 index 000000000..eebf137fb --- /dev/null +++ b/python/sglang/srt/configs/step3p5.py @@ -0,0 +1,97 @@ +from typing import Any, Optional + +from transformers.configuration_utils import PretrainedConfig + + +class Step3p5Config(PretrainedConfig): + model_type = "step3p5" + architectures = ["Step3p5ForCausalLM"] + + def __init__( + self, + hidden_size: int = 4096, + intermediate_size: int = 11264, + num_attention_heads: int = 64, + num_attention_groups: int = 8, + num_hidden_layers: int = 45, + max_seq_len: int = 128000, + vocab_size: int = 128815, + rms_norm_eps: float = 1e-5, + moe_intermediate_size: int = 1280, + moe_num_experts: int = 288, + moe_top_k: int = 8, + rope_theta: float = 10000, + rope_scaling: Optional[dict[str, Any]] = None, + max_position_embeddings: int = 128000, + share_expert_dims: int = 1280, + head_dim: int = 128, + norm_expert_weight: bool = True, + layer_types: list[str] = None, + sliding_window: Optional[int] = None, + moe_layers_enum: tuple[int] = ( + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 10, + 11, + 12, + 13, + 14, + 15, + 16, + 17, + 18, + 19, + 20, + 21, + 22, + 23, + 24, + 25, + 26, + 27, + 28, + 29, + 30, + 31, + 32, + 33, + 34, + 35, + 36, + 37, + 38, + 39, + 40, + 41, + 42, + 43, + 44, + ), + **kwargs, + ) -> None: + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.num_attention_heads = num_attention_heads + self.num_attention_groups = num_attention_groups + self.num_hidden_layers = num_hidden_layers + self.max_seq_len = max_seq_len + self.vocab_size = vocab_size + self.rms_norm_eps = rms_norm_eps + self.moe_intermediate_size = moe_intermediate_size + self.moe_num_experts = moe_num_experts + self.moe_top_k = moe_top_k + self.rope_theta = rope_theta + self.rope_scaling = rope_scaling + self.max_position_embeddings = max_position_embeddings + self.share_expert_dim = share_expert_dims + self.head_dim = head_dim + self.norm_expert_weight = norm_expert_weight + self.moe_layers_enum = moe_layers_enum + self.layer_types = layer_types + self.sliding_window = sliding_window + super().__init__(**kwargs) diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 10d14cc43..e7fdd6398 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -62,6 +62,7 @@ class FunctionCallParser: "qwen25": Qwen25Detector, "qwen3_coder": Qwen3CoderDetector, "step3": Step3Detector, + "step3p5": Qwen3CoderDetector, "minimax-m2": MinimaxM2Detector, "trinity": TrinityDetector, "interns1": InternlmDetector, diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py index a1885fade..79f3aa2e9 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py @@ -278,7 +278,18 @@ def moe_sum_reduce_torch_compile(x, out, routed_scaling_factor): @torch.compile -def swiglu_with_alpha_and_limit(x, gemm1_alpha, gemm1_limit): +def _swiglu_silu_clamp_mul(x, gemm1_limit): + gate, up = x.chunk(2, dim=-1) + gate = F.silu(gate) + gate = gate.clamp(min=None, max=gemm1_limit) + up = up.clamp(min=-gemm1_limit, max=gemm1_limit) + return gate * up + + +@torch.compile +def _swiglu_gpt_oss_sigmoid_alpha(x, gemm1_alpha, gemm1_limit): + # NOTE: This variant uses gemm1_alpha, unlike _swiglu_silu_clamp_mul. + # At present, only GPT-OSS uses this variant. gate, up = x[..., ::2], x[..., 1::2] gate = gate.clamp(min=None, max=gemm1_limit) up = up.clamp(min=-gemm1_limit, max=gemm1_limit) @@ -471,12 +482,16 @@ def fused_experts_impl( # Activation function with multiplication if activation == "silu" and is_gated: + # - gemm1_alpha != None: GPT-OSS-style swiglu(alpha, limit) + # - gemm1_alpha == None and gemm1_limit != None: silu+clamp+mul(limit-only) if gemm1_alpha is not None: assert gemm1_limit is not None - intermediate_cache2 = swiglu_with_alpha_and_limit( - intermediate_cache1.view(-1, N), - gemm1_alpha, - gemm1_limit, + intermediate_cache2 = _swiglu_gpt_oss_sigmoid_alpha( + intermediate_cache1.view(-1, N), gemm1_alpha, gemm1_limit + ) + elif gemm1_limit is not None: + intermediate_cache2 = _swiglu_silu_clamp_mul( + intermediate_cache1.view(-1, N), gemm1_limit ) elif _is_cuda or _is_hip: if not filter_expert: diff --git a/python/sglang/srt/layers/moe/moe_runner/triton.py b/python/sglang/srt/layers/moe/moe_runner/triton.py index cdf3e9a47..14fae2623 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton.py @@ -117,10 +117,11 @@ class TritonRunnerCore(MoeRunnerCore): # TODO: move these functions to the triton runner from sglang.srt.layers.moe.fused_moe_triton.fused_moe import ( + _swiglu_gpt_oss_sigmoid_alpha, + _swiglu_silu_clamp_mul, invoke_fused_moe_kernel, moe_sum_reduce_torch_compile, moe_sum_reduce_triton, - swiglu_with_alpha_and_limit, ) hidden_states = runner_input.hidden_states @@ -203,10 +204,12 @@ class TritonRunnerCore(MoeRunnerCore): if activation == "silu": if gemm1_alpha is not None: assert gemm1_limit is not None - intermediate_cache2 = swiglu_with_alpha_and_limit( - intermediate_cache1.view(-1, N), - gemm1_alpha, - gemm1_limit, + intermediate_cache2 = _swiglu_gpt_oss_sigmoid_alpha( + intermediate_cache1.view(-1, N), gemm1_alpha, gemm1_limit + ) + elif gemm1_limit is not None: + intermediate_cache2 = _swiglu_silu_clamp_mul( + intermediate_cache1.view(-1, N), gemm1_limit ) elif _is_cuda or _is_hip: silu_and_mul(intermediate_cache1.view(-1, N), intermediate_cache2) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 01d6c8866..a9a425cde 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -493,6 +493,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) if self.model_config.hf_config.architectures[0] == "MiMoV2MTP": model_num_layers = 1 + elif self.model_config.hf_config.architectures[0] == "Step3p5MTP": + model_num_layers = 1 self.start_layer = getattr(self.model, "start_layer", 0) self.end_layer = getattr(self.model, "end_layer", model_num_layers) self.num_effective_layers = self.end_layer - self.start_layer diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 1b6658c6a..c195e2b26 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -275,6 +275,9 @@ def _initialize_model( kwargs["sparse_head"] = envs.SGLANG_EMBEDDINGS_SPARSE_HEAD.get() kwargs["model_path"] = model_config.model_path + if load_config.draft_model_idx is not None: + kwargs["draft_model_idx"] = load_config.draft_model_idx + return model_class(**kwargs) diff --git a/python/sglang/srt/models/mimo_v2_flash_nextn.py b/python/sglang/srt/models/mimo_v2_flash_nextn.py index 2408f950d..18b545395 100644 --- a/python/sglang/srt/models/mimo_v2_flash_nextn.py +++ b/python/sglang/srt/models/mimo_v2_flash_nextn.py @@ -229,6 +229,7 @@ class MiMoV2MTP(MiMoV2FlashForCausalLM): self, config: PretrainedConfig, quant_config: Optional[QuantizationConfig] = None, + draft_model_idx: Optional[int] = None, prefix: str = "", ) -> None: nn.Module.__init__(self) diff --git a/python/sglang/srt/models/step3p5.py b/python/sglang/srt/models/step3p5.py new file mode 100644 index 000000000..b3f82b916 --- /dev/null +++ b/python/sglang/srt/models/step3p5.py @@ -0,0 +1,1037 @@ +import logging +import os +from typing import Any, Dict, Iterable, Optional, Tuple, Union + +import torch +import torch.nn.functional as F +from torch import nn + +from sglang.srt.distributed import ( + get_moe_expert_parallel_world_size, + get_pp_group, + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, + tensor_model_parallel_all_reduce, +) +from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder +from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo +from sglang.srt.layers.activation import SiluAndMul +from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes +from sglang.srt.layers.dp_attention import ( + get_attention_tp_rank, + get_attention_tp_size, + is_dp_attention_enabled, +) +from sglang.srt.layers.layernorm import GemmaRMSNorm +from sglang.srt.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + QKVParallelLinear, + ReplicatedLinear, + RowParallelLinear, +) +from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.moe import ( + get_moe_a2a_backend, + should_use_flashinfer_cutlass_moe_fp4_allgather, +) +from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class +from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE +from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK +from sglang.srt.layers.moe.utils import ( + RoutingMethodType, + filter_moe_weight_param_global_expert, +) +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 +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.server_args import get_global_server_args +from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers + +Step3p5Config = None + +logger = logging.getLogger(__name__) +_is_cuda = is_cuda() + + +class Step3p5MLP(nn.Module): + def __init__( + self, + hidden_size: int, + intermediate_size: int, + swiglu_limit: Optional[float] = None, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + 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), + ) + self.act_fn = SiluAndMul() + self.limit = swiglu_limit + + def forward(self, x): + if self.limit is not None: + gate_up, _ = self.gate_up_proj(x) + gate, up = gate_up.chunk(2, dim=-1) + gate = F.silu(gate) + gate = gate.clamp(min=None, max=self.limit) + up = up.clamp(min=-self.limit, max=self.limit) + output, _ = self.down_proj(gate * up) + else: + gate_up, _ = self.gate_up_proj(x) + x = self.act_fn(gate_up) + output, _ = self.down_proj(x) + return output + + +class Step3p5MoEMLP(nn.Module): + def __init__( + self, + config, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ): + super().__init__() + self.tp_size = get_tensor_model_parallel_world_size() + self.layer_id = layer_id + + self.need_fp32_gate = config.need_fp32_gate + self.routed_scaling_factor = config.moe_router_scaling_factor + self.use_moe_router_bias = config.use_moe_router_bias + if self.use_moe_router_bias: + self.router_bias = nn.Parameter( + torch.zeros(config.moe_num_experts, dtype=torch.float32), + requires_grad=False, + ) + + if self.tp_size > config.moe_num_experts: + raise ValueError( + f"Tensor parallel size {self.tp_size} is greater than " + f"the number of experts {config.moe_num_experts}." + ) + + self.limit = config.swiglu_limits[layer_id] + self.limit = self.limit if self.limit > 0 else None + + self.topk = TopK( + top_k=config.moe_top_k, + renormalize=True, + use_grouped_topk=False, + scoring_func="sigmoid", + correction_bias=self.router_bias, + apply_routed_scaling_factor_on_output=False, + layer_id=layer_id, + ) + + self.experts = get_moe_impl_class(quant_config)( + num_experts=config.moe_num_experts + + get_global_server_args().ep_num_redundant_experts, + top_k=config.moe_top_k, + layer_id=layer_id, + hidden_size=config.hidden_size, + intermediate_size=config.moe_intermediate_size, + quant_config=quant_config, + prefix=add_prefix("experts", prefix), + routing_method_type=RoutingMethodType.Renormalize, + gemm1_clamp_limit=self.limit, + ) + + self.gate = ReplicatedLinear( + config.hidden_size, + config.moe_num_experts, + bias=False, + quant_config=None, + prefix=add_prefix("gate", prefix), + ) + + if get_moe_a2a_backend().is_deepep(): + # TODO: we will support tp < ep in the future + self.ep_size = get_moe_expert_parallel_world_size() + self.moe_num_experts = ( + config.moe_num_experts + + get_global_server_args().ep_num_redundant_experts + ) + self.top_k = config.moe_top_k + + def forward( + self, + hidden_states: torch.Tensor, + forward_batch: Optional[ForwardBatch] = None, + should_allreduce_fusion: bool = False, + use_reduce_scatter: bool = False, + ) -> torch.Tensor: + + if ( + not get_moe_a2a_backend().is_deepep() + and not get_moe_a2a_backend().is_ascend_fuseep() + ): + return self.forward_normal( + hidden_states, should_allreduce_fusion, use_reduce_scatter + ) + else: + return self.forward_deepep(hidden_states, forward_batch) + + def get_moe_weights(self): + return [ + x.data + for name, x in self.experts.named_parameters() + if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) + ] + + def forward_normal( + self, + hidden_states: torch.Tensor, + should_allreduce_fusion: bool = False, + use_reduce_scatter: bool = False, + ) -> torch.Tensor: + num_tokens, hidden_dim = hidden_states.shape + hidden_states = hidden_states.view(-1, hidden_dim) + # router_logits: (num_tokens, n_experts) + if self.need_fp32_gate: + router_logits = torch.matmul( + hidden_states.to(torch.float32), self.gate.weight.t().to(torch.float32) + ) + else: + # router_logits: (batch * sequence_length, n_experts) + router_logits, _ = self.gate(hidden_states) + topk_output = self.topk(hidden_states, router_logits) + if self.routed_scaling_factor != 1.0: + topk_output = StandardTopKOutput( + topk_weights=topk_output.topk_weights * self.routed_scaling_factor, + topk_ids=topk_output.topk_ids, + router_logits=topk_output.router_logits, + ) + final_hidden_states = self.experts(hidden_states, topk_output) + if ( + self.tp_size > 1 + and not should_allreduce_fusion + and not use_reduce_scatter + and not should_use_flashinfer_cutlass_moe_fp4_allgather() + ): + final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states) + + return final_hidden_states.view(num_tokens, hidden_dim) + + def forward_deepep( + self, hidden_states: torch.Tensor, forward_batch: ForwardBatch + ) -> torch.Tensor: + if hidden_states.shape[0] > 0: + # router_logits: (num_tokens, n_experts) + router_logits, _ = self.gate(hidden_states) + topk_output = self.topk( + hidden_states, + router_logits, + num_token_non_padded=forward_batch.num_token_non_padded, + expert_location_dispatch_info=ExpertLocationDispatchInfo.init_new( + layer_id=self.layer_id, + ), + ) + else: + topk_output = self.topk.empty_topk_output(hidden_states.device) + final_hidden_states = self.experts( + hidden_states=hidden_states, + topk_output=topk_output, + ) + return final_hidden_states + + def op_gate(self, state): + if is_non_idle_and_non_empty( + state.forward_batch.forward_mode, state.hidden_states_mlp_input + ): + # router_logits: (num_tokens, n_experts) + state.router_logits, _ = self.gate(state.hidden_states_mlp_input) + else: + state.router_logits = None + + def op_select_experts(self, state): + router_logits = state.pop("router_logits") + hidden_states = state.hidden_states_mlp_input + if router_logits is not None: + with get_global_expert_distribution_recorder().with_current_layer( + self.layer_id + ): + state.topk_output = self.topk( + hidden_states=hidden_states, + router_logits=router_logits, + num_token_non_padded=state.forward_batch.num_token_non_padded, + expert_location_dispatch_info=ExpertLocationDispatchInfo.init_new( + layer_id=self.layer_id, + ), + ) + else: + state.topk_output = self.topk.empty_topk_output(hidden_states.device) + + def op_dispatch_a(self, state): + if self.ep_size > 1: + self.experts.dispatcher.dispatch_a( + hidden_states=state.pop("hidden_states_mlp_input"), + topk_output=state.pop("topk_output"), + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + + def op_dispatch_b(self, state): + if self.ep_size > 1: + with get_global_expert_distribution_recorder().with_current_layer( + self.layer_id + ): + state.dispatch_output = self.experts.dispatcher.dispatch_b( + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + + def op_experts(self, state): + state.combine_input = self.experts.run_moe_core( + dispatch_output=state.dispatch_output, + ) + + def op_combine_a(self, state): + if self.ep_size > 1: + self.experts.dispatcher.combine_a( + combine_input=state.pop("combine_input"), + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + state.pop("dispatch_output") + + def op_combine_b(self, state): + if self.ep_size > 1: + state.hidden_states_after_combine = self.experts.dispatcher.combine_b( + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + + def op_output(self, state): + state.hidden_states_mlp_output = state.pop("hidden_states_after_combine") + + +class Step3p5Attention(nn.Module): + def __init__( + self, + hidden_size: int, + num_heads: int, + num_kv_heads: int, + layer_id: int = 0, + rope_theta: float = 1000000, + rope_scaling: Optional[Dict[str, Any]] = None, + head_dim: Optional[int] = None, + max_position_embeddings: int = 32768, + quant_config: Optional[QuantizationConfig] = None, + rms_norm_eps: float = None, + partial_rotary_factor: float = 1.0, + use_head_wise_attn_gate: bool = False, + sliding_window_size: int = -1, # if is -1 ,normal attention,else ,window attention + prefix: str = "", + alt_stream: Optional[torch.cuda.Stream] = None, + ) -> None: + super().__init__() + self.hidden_size = hidden_size + self.tp_size = get_tensor_model_parallel_world_size() + self.total_num_heads = num_heads + attn_tp_rank = get_attention_tp_rank() + attn_tp_size = get_attention_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) + self.head_dim = head_dim or 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.tp_rank = get_tensor_model_parallel_rank() + self.q_norm = GemmaRMSNorm(self.head_dim, eps=rms_norm_eps) + self.k_norm = GemmaRMSNorm(self.head_dim, eps=rms_norm_eps) + + self.qkv_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("qkv_proj", prefix), + ) + self.o_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("o_proj", prefix), + ) + + self.use_head_wise_attn_gate = use_head_wise_attn_gate + if self.use_head_wise_attn_gate: + self.g_proj = ColumnParallelLinear( + hidden_size, + self.total_num_heads, + bias=False, + tp_rank=attn_tp_rank, + tp_size=attn_tp_size, + prefix=add_prefix("g_proj", prefix), + ) + + self.rotary_emb = get_rope( + self.head_dim, + rotary_dim=self.head_dim, + max_position=max_position_embeddings, + base=rope_theta, + rope_scaling=rope_scaling, + partial_rotary_factor=partial_rotary_factor, + is_neox_style=True, + ) + self.attn = RadixAttention( + self.num_heads, + self.head_dim, + self.scaling, + num_kv_heads=self.num_kv_heads, + sliding_window_size=sliding_window_size, # if is -1 ,normal attention,else ,window attention + layer_id=layer_id, + prefix=add_prefix("attn", prefix), + ) + self.alt_stream = alt_stream + + def forward_prepare_native(self, positions, hidden_states): + qkv, _ = self.qkv_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + q_shape, k_shape = q.shape, k.shape + q = self.q_norm(q.reshape(-1, self.head_dim)).reshape(q_shape) + k = self.k_norm(k.reshape(-1, self.head_dim)).reshape(k_shape) + q, k = self.rotary_emb(positions, q, k) + return q, k, v + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + + q, k, v = self.forward_prepare_native( + positions=positions, + hidden_states=hidden_states, + ) + if self.use_head_wise_attn_gate: + gate_states, _ = self.g_proj(hidden_states) + attn_output = self.attn(q, k, v, forward_batch) + if self.use_head_wise_attn_gate: + output = ( + attn_output.view( + attn_output.shape[0], + self.num_heads, # TODO: check if this is correct + self.head_dim, + ) + * gate_states.unsqueeze(-1).sigmoid() + ) + attn_output = output.view(*attn_output.shape) + output, _ = self.o_proj(attn_output) + return output + + +class Step3p5DecoderLayer(nn.Module): + def __init__( + self, + config: Step3p5Config, + layer_id: int = 0, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + alt_stream: Optional[torch.cuda.Stream] = None, + ) -> None: + super().__init__() + self.hidden_size = config.hidden_size + layer_types = config.layer_types + yarn_only_types = config.yarn_only_types + if layer_types[layer_id] not in yarn_only_types: + rope_scaling = None + else: + rope_scaling = config.rope_scaling + rope_theta = config.rope_theta + max_position_embeddings = config.max_position_embeddings + head_dim = config.head_dim + moe_layers_list = [int(x) for x in config.moe_layers_enum.split(",")] + self.num_attention_heads = config.num_attention_heads + self.num_key_value_heads = config.num_attention_groups + self.is_moe_layer = layer_id in moe_layers_list + num_hidden_layers = config.num_hidden_layers + + if ( + config.swiglu_limits_shared + and config.swiglu_limits_shared[layer_id] is not None + and config.swiglu_limits_shared[layer_id] != 0 + ): + swiglu_limit_shared = config.swiglu_limits_shared[layer_id] + else: + swiglu_limit_shared = None + + self.sliding_window = -1 + + enable_sliding_window = layer_types[layer_id] == "sliding_attention" + + if enable_sliding_window: + self.sliding_window = config.sliding_window + self.num_attention_heads = config.attention_other_setting[ + "num_attention_heads" + ] + self.num_key_value_heads = config.attention_other_setting[ + "num_attention_groups" + ] + + self.self_attn = Step3p5Attention( + hidden_size=self.hidden_size, + num_heads=self.num_attention_heads, + num_kv_heads=self.num_key_value_heads, + layer_id=( + layer_id + if layer_id < num_hidden_layers + else layer_id - num_hidden_layers + ), + rope_theta=rope_theta[layer_id], + rope_scaling=rope_scaling, + head_dim=head_dim, + max_position_embeddings=max_position_embeddings, + sliding_window_size=self.sliding_window, + partial_rotary_factor=config.partial_rotary_factors[layer_id], + quant_config=quant_config, + rms_norm_eps=config.rms_norm_eps, + use_head_wise_attn_gate=config.use_head_wise_attn_gate, + prefix=add_prefix("self_attn", prefix), + alt_stream=alt_stream, + ) + self.use_moe = False + if self.is_moe_layer: + self.moe = Step3p5MoEMLP( + config, + layer_id=layer_id, + quant_config=quant_config, + prefix=add_prefix("mlp", prefix), + ) + self.share_expert = Step3p5MLP( + hidden_size=self.hidden_size, + intermediate_size=config.share_expert_dim, + swiglu_limit=swiglu_limit_shared, + quant_config=quant_config, + prefix=add_prefix("share_expert", prefix), + ) + self.use_moe = True + else: + self.mlp = Step3p5MLP( + hidden_size=self.hidden_size, + intermediate_size=config.intermediate_size, + swiglu_limit=swiglu_limit_shared, + quant_config=quant_config, + prefix=add_prefix("mlp", prefix), + ) + + self.input_layernorm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = GemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + self.layer_scatter_modes = LayerScatterModes.init_new( + layer_id=layer_id, + num_layers=( + config.num_hidden_layers if layer_id < config.num_hidden_layers else 1 + ), # 1 is for mtp + is_layer_sparse=False, + is_previous_layer_sparse=False, + is_next_layer_sparse=False, + ) + self.layer_communicator = LayerCommunicator( + layer_scatter_modes=self.layer_scatter_modes, + input_layernorm=self.input_layernorm, + post_attention_layernorm=self.post_attention_layernorm, + ) + + self.layer_id = layer_id + self.dump_intermediate = ( + os.environ.get("SGLANG_DUMP_STEP3P5_INTERMEDIATE") == "1" + ) + self._dump_step = 0 + + def _dump_tensor( + self, + name: str, + tensor: Optional[torch.Tensor], + step_id: Optional[int] = None, + ) -> None: + if not self.dump_intermediate or tensor is None or not torch.is_tensor(tensor): + return + dump_dir = "/sgl-workspace/sgl" + try: + os.makedirs(dump_dir, exist_ok=True) + tp_rank = get_tensor_model_parallel_rank() + step_part = f"_step{step_id}" if step_id is not None else "" + path = os.path.join( + dump_dir, + f"step3p5_layer{self.layer_id}{step_part}_{name}_tp{tp_rank}.pt", + ) + torch.save(tensor.detach().cpu(), path) + except Exception: + logger.exception( + "Failed to dump tensor %s for layer %s", name, self.layer_id + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + residual: Optional[torch.Tensor], + post_residual_addition: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + # Self Attention + hidden_states, residual = self.layer_communicator.prepare_attn( + hidden_states, + residual, + forward_batch, + post_residual_addition=post_residual_addition, + ) + dump_step = None + if self.dump_intermediate: + dump_step = self._dump_step + self._dump_step += 1 + self._dump_tensor("attn_input", hidden_states, dump_step) + if hidden_states.shape[0] != 0: + hidden_states = self.self_attn( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + ) + self._dump_tensor("attn_output", hidden_states, dump_step) + # Fully Connected + # hidden_states, residual = self.layer_communicator.prepare_mlp( + # hidden_states, + # residual, + # forward_batch, + # ) + hidden_states = residual + hidden_states + residual = hidden_states + self._dump_tensor("post_attn_residual", hidden_states, dump_step) + hidden_states = self.post_attention_layernorm(hidden_states) + self._dump_tensor("mlp_input", hidden_states, dump_step) + if self.use_moe: + share_output = self.share_expert(hidden_states) + moe_output = self.moe(hidden_states) + hidden_states = moe_output + share_output + else: + hidden_states = self.mlp(hidden_states) + self._dump_tensor("mlp_output", hidden_states, dump_step) + hidden_states, residual = self.layer_communicator.postprocess_layer( + hidden_states, residual, forward_batch + ) + self._dump_tensor("layer_output", hidden_states, dump_step) + return hidden_states, residual + + +class Step3p5Model(nn.Module): + def __init__( + self, + config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + self.padding_idx = config.pad_token_id + self.vocab_size = config.vocab_size + self.pp_group = get_pp_group() + + alt_stream = torch.cuda.Stream() if _is_cuda else None + + if self.pp_group.is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + enable_tp=not is_dp_attention_enabled(), + prefix=add_prefix("embed_tokens", prefix), + params_dtype=( + torch.float32 + if get_global_server_args().rl_on_policy_target is not None + else None + ), + ) + else: + self.embed_tokens = PPMissingLayer() + + self.layers, self.start_layer, self.end_layer = make_layers( + config.num_hidden_layers, + # 1, + lambda idx, prefix: Step3p5DecoderLayer( + layer_id=idx, + config=config, + quant_config=quant_config, + prefix=prefix, + alt_stream=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 = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + else: + self.norm = PPMissingLayer(return_tuple=True) + + def get_input_embedding(self, input_ids: torch.Tensor) -> torch.Tensor: + if hasattr(self.config, "scale_emb"): + return self.get_input_embeddings()(input_ids) * self.config.scale_emb + else: + return self.get_input_embeddings()(input_ids) + + 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, + ) + # break + if not self.pp_group.is_last_rank: + return PPProxyTensors( + { + "hidden_states": hidden_states, + "residual": residual, + } + ) + else: + hidden_states_before_norm = None + if not self.pp_group.is_last_rank: + return PPProxyTensors( + { + "hidden_states": hidden_states, + "residual": residual, + } + ) + else: + if hidden_states.shape[0] > 0: + # if forward_batch.return_hidden_states_before_norm: + hidden_states_before_norm = ( + hidden_states if residual is None else hidden_states + residual + ) + if residual is None: + hidden_states = self.norm(hidden_states) + else: + hidden_states, _ = self.norm(hidden_states, residual) + return hidden_states, hidden_states_before_norm + + +class Step3p5ForCausalLM(nn.Module): + # BitandBytes specific attributes + default_bitsandbytes_target_modules = [ + ".gate_proj.", + ".down_proj.", + ".up_proj.", + ".q_proj.", + ".k_proj.", + ".v_proj.", + ".o_proj.", + ] + bitsandbytes_stacked_params_mapping = { + # shard_name, weight_name, index + "q_proj": ("qkv_proj", 0), + "k_proj": ("qkv_proj", 1), + "v_proj": ("qkv_proj", 2), + "gate_proj": ("gate_up_proj", 0), + "up_proj": ("gate_up_proj", 1), + } + + def __init__( + self, + config: Step3p5Config, + 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 = Step3p5Model( + config, quant_config=quant_config, prefix=add_prefix("model", prefix) + ) + + self.tie_word_embeddings = False + self.num_fused_shared_experts = 0 + + # handle the lm head on different pp ranks + if self.pp_group.is_last_rank: + if self.pp_group.world_size == 1 and self.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, + use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + prefix=add_prefix("lm_head", prefix), + ) + else: + # ranks other than the last rank will have a placeholder layer + self.lm_head = PPMissingLayer() + + # perform weight tying for PP + if self.pp_group.world_size > 1 and self.tie_word_embeddings: + if self.pp_group.is_first_rank: + self.pp_group.send( + self.model.embed_tokens.weight, dst=self.pp_group.world_size - 1 + ) + elif self.pp_group.is_last_rank: + emb_token_weight = self.pp_group.recv( + size=self.lm_head.weight.shape, + dtype=next(self.model.parameters()).dtype, + src=0, + ) + self.lm_head.weight.copy_(emb_token_weight) + + self.logits_processor = LogitsProcessor(config) + + def get_input_embeddings(self) -> nn.Embedding: + return self.model.get_input_embeddings() + + @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, hidden_states_before_norm = 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, + hidden_states_before_norm=hidden_states_before_norm, + ) + else: + return hidden_states + + @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]], is_nextn=False): + # NOTE: + # Step3p5 HF checkpoints (e.g. MTP/nextn variants) may include an extra + # "nextn predict layer" appended after the main decoder layers, such as: + # model.layers..(eh_proj|enorm|hnorm|transformer.shared_head.*) + # This implementation currently does NOT instantiate those nextn modules, + # so we must safely skip them (or load them only when a corresponding + # nextn model is implemented). + + def _get_layer_id_from_weight_name(weight_name: str) -> Optional[int]: + # Expected format: "model.layers....." + parts = weight_name.split(".") + if len(parts) >= 3 and parts[0] == "model" and parts[1] == "layers": + try: + return int(parts[2]) + except ValueError: + return None + return None + + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + (".qkv_proj", ".q_proj", "q"), + (".qkv_proj", ".k_proj", "k"), + (".qkv_proj", ".v_proj", "v"), + (".gate_up_proj", ".gate_proj", 0), + (".gate_up_proj", ".up_proj", 1), + ] + + if self.num_fused_shared_experts > 0: + assert self.num_fused_shared_experts == 1 + + expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.moe_num_experts + self.num_fused_shared_experts, + ) + + params_dict = dict(self.named_parameters()) + loaded_params = set() + + def match_expert_and_shard_ids(name_path: str, weight_path: str) -> bool: + name_parts = name_path.split(".") + weight_parts = weight_path.split(".") + # Be defensive: some unexpected weight names may not match the shape. + if len(name_parts) <= 4 or len(weight_parts) <= 2: + return False + shard_id_matches = name_parts[4] == weight_parts[2] + return shard_id_matches + + for name, loaded_weight in weights: + # Filter nextn layer weights. + if hasattr(self.config, "num_nextn_predict_layers"): + num_nextn_layers = getattr(self.config, "num_nextn_predict_layers", 0) + if num_nextn_layers and name.startswith("model.layers."): + layer_id = _get_layer_id_from_weight_name(name) + if layer_id is not None: + if not is_nextn: + # Normal load: skip layers appended after the main decoder. + if layer_id >= self.config.num_hidden_layers: + continue + else: + # nextn load: only keep the appended nextn layer. + # (Only 1 nextn layer is supported by current checkpoints.) + if num_nextn_layers != 1: + raise ValueError( + "Only 1 nextn layer is supported for Step3p5 checkpoints." + ) + nextn_layer_id = ( + 0 + if self.config.num_hidden_layers == 1 + else self.config.num_hidden_layers + ) + if layer_id != nextn_layer_id: + # # nextn/MTP load: only keep the appended nextn layers. + # # Expected layer ids: [num_hidden_layers, num_hidden_layers + num_nextn_layers). + # start = self.config.num_hidden_layers + # end = self.config.num_hidden_layers + num_nextn_layers + # if not (start <= layer_id < end): + continue + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + if "gate." not in name and "moe" in name: + continue + name = name.replace(weight_name, param_name) + if name not in params_dict: + # Extra / unsupported weights (e.g. nextn) should not crash loading. + continue + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + loaded_params.add(name) + break + else: + if "moe" not in name or "router_bias" in name: + if name not in params_dict: + continue + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight) + loaded_params.add(name) + else: + if "gate." in name: + if name not in params_dict: + continue + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight) + loaded_params.add(name) + continue + + for mapping in expert_params_mapping: + param_name, weight_name, expert_id, shard_id = mapping + if expert_id == self.config.moe_num_experts: + continue + if not match_expert_and_shard_ids(name, weight_name): + continue + part_name = weight_name.split(".")[-2] + fake_weight_name = name.replace(part_name, weight_name[:-1]) + actual_param_name = name.replace(part_name + ".", param_name) + if actual_param_name not in params_dict: + continue + param = params_dict[actual_param_name] + weight_loader = param.weight_loader + weight_loader( + param, + loaded_weight[expert_id], + name, + shard_id=shard_id, + expert_id=expert_id, + ) + loaded_params.add(actual_param_name) + + print_params = set(params_dict.keys()) - loaded_params + assert len(print_params) == 0, f"Some parameters are not loaded: {print_params}" + + 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 + torch.cuda.empty_cache() + torch.cuda.synchronize() + + +EntryClass = Step3p5ForCausalLM diff --git a/python/sglang/srt/models/step3p5_mtp.py b/python/sglang/srt/models/step3p5_mtp.py new file mode 100644 index 000000000..c52c83847 --- /dev/null +++ b/python/sglang/srt/models/step3p5_mtp.py @@ -0,0 +1,336 @@ +import logging +from collections.abc import Iterable +from typing import Optional + +import torch +import torch.nn as nn +from transformers import PretrainedConfig + +from sglang.srt.distributed import get_tensor_model_parallel_world_size +from sglang.srt.layers.layernorm import GemmaRMSNorm +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.model_loader.weight_utils import default_weight_loader +from sglang.srt.models.step3p5 import Step3p5DecoderLayer, Step3p5ForCausalLM +from sglang.srt.utils import add_prefix + +logger = logging.getLogger(__name__) + + +def get_spec_layer_idx_from_weight_name( + config: PretrainedConfig, weight_name: str +) -> Optional[int]: + """Return MTP/nextn layer index if this weight belongs to spec layers. + + Step3p5 MTP/nextn checkpoints append extra layers after the main decoder: + model.layers.[num_hidden_layers ... num_hidden_layers + num_nextn_predict_layers) + """ + if hasattr(config, "num_nextn_predict_layers") and ( + getattr(config, "num_nextn_predict_layers", 0) > 0 + ): + base = config.num_hidden_layers + for i in range(config.num_nextn_predict_layers): + if weight_name.startswith(f"model.layers.{base + i}."): + return base + i + return None + + +class SharedHead(nn.Module): + + def __init__( + self, + config, + quant_config=None, + ) -> None: + super().__init__() + self.norm = GemmaRMSNorm(config.hidden_size, config.rms_norm_eps) + self.head = ParallelLMHead( + config.vocab_size, config.hidden_size, quant_config=quant_config + ) + self.lm_head = self.head + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.norm(hidden_states) + + +class Step3p5AMultiTokenPredictor(nn.Module): + def __init__( + self, + config: PretrainedConfig, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + ) + self.mtp_start_layer_idx = config.num_hidden_layers + self.num_mtp_layers = config.num_nextn_predict_layers + + layer_id = 45 # FIXME + + self.enorm = GemmaRMSNorm(config.hidden_size, config.rms_norm_eps) + self.hnorm = GemmaRMSNorm(config.hidden_size, config.rms_norm_eps) + self.eh_proj = nn.Linear(config.hidden_size * 2, config.hidden_size, bias=False) + self.shared_head = SharedHead(config=config, quant_config=quant_config) + self.mtp_block = Step3p5DecoderLayer( + config=config, layer_id=layer_id, prefix=f"{prefix}.mtp_block" + ) + self.lm_head = self.shared_head.head + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: torch.Tensor = None, + ) -> torch.Tensor: + 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, + ) + ) + hidden_states, residual = self.mtp_block( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + residual=None, + ) + hidden_states_before_norm = None + if not forward_batch.forward_mode.is_idle(): + # if forward_batch.return_hidden_states_before_norm: + 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 + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) + + +# The current implementation differs slightly from the standard MTP implementation in Step3.5 Flash. +# In the standard multi-layer MTP design of Step3.5 Flash, +# the hidden states of each MTP layer are passed from the preceding MTP layer +# (the hidden states of the initial (layer-0) MTP still being provided by the target model). +# In contrast, the current SGL implementation obtains hidden states directly from the target model for all MTP layers. +# Empirical evaluations indicate that the overall performance remains strong; +# however, this design choice may lead to a slight reduction in acceptance rate in certain scenarios. +# This behavior will be corrected shortly, and we expect to implement the standard multi-layer MTP design of Step3.5 Flash in the near future. +# FIXME(yhyang201) +class Step3p5MTP(Step3p5ForCausalLM): + def __init__( + self, + config: PretrainedConfig, + quant_config: Optional[QuantizationConfig] = None, + draft_model_idx: Optional[int] = None, + prefix: str = "", + ) -> None: + nn.Module.__init__(self) + self.config = config + self.tp_size = get_tensor_model_parallel_world_size() + self.quant_config = quant_config + self.draft_model_idx = draft_model_idx + + self.model = Step3p5AMultiTokenPredictor( + config=config, quant_config=quant_config, prefix=add_prefix("model", prefix) + ) + self.logits_processor = LogitsProcessor(config) + self.lm_head = self.model.lm_head + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_input_ids(input_ids) + + 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.model.shared_head.head, + forward_batch, + hidden_states_before_norm=hidden_states_before_norm, + ) + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.model.shared_head.head.weight + + def set_embed_and_head(self, embed, head): + return + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + + expert_params_mapping = [ + (".moe.experts.w13_weight", ".moe.gate_proj.weight", "w1"), + (".moe.experts.w13_weight", ".moe.up_proj.weight", "w3"), + (".moe.experts.w2_weight", ".moe.down_proj.weight", "w2"), + ] + + params_dict = dict(self.named_parameters()) + loaded_params: set[str] = set() + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + spec_layer = get_spec_layer_idx_from_weight_name(self.config, name) + if spec_layer is not None and spec_layer != ( + self.config.num_hidden_layers + self.draft_model_idx + ): + continue + if "embed_tokens" not in name and spec_layer is None: + continue + name = self._rewrite_spec_layer_name(spec_layer, name) + for param_name, weight_name, shard_id in stacked_params_mapping: + # Skip non-stacked layers and experts (experts handled below). + if weight_name not in name: + continue + # We have mlp.experts[0].gate_proj in the checkpoint. + # Since we handle the experts below in expert_params_mapping, + # we need to skip here BEFORE we update the name, otherwise + # name will be updated to mlp.experts[0].gate_up_proj, which + # will then be updated below in expert_params_mapping + # for mlp.experts[0].gate_gate_up_proj, which breaks load. + if ("mlp.experts." in name) and name not in params_dict: + continue + if "experts" in name or "moe" in name: + continue + name = name.replace(weight_name, param_name) + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue + + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + break + else: + for mapping in expert_params_mapping: + param_name, weight_name, shard_id = mapping + if weight_name not in name: + continue + name = name.replace(weight_name, param_name) + # Skip loading extra bias for GPTQ models. + if ( + name.endswith(".bias") or name.endswith("_bias") + ) and name not in params_dict: + continue + param = params_dict[name] + weight_loader = param.weight_loader + for expert_id in range(loaded_weight.shape[0]): + loaded_weight_expert = loaded_weight[expert_id] + weight_loader( + param, + loaded_weight_expert, + name, + shard_id=shard_id, + expert_id=expert_id, + ) + loaded_params.add(name) + break + else: + # Skip loading extra bias for GPTQ models. + if ( + name.endswith(".bias") + and name not in params_dict + or "tok_embeddings" in name + ): + continue + + if "shared_head" in name: + name = name.replace("shared_head.output", "shared_head.head") + if "embed_tokens" in name: + assert ( + hasattr(self.config, "num_nextn_predict_layers") + and self.config.num_nextn_predict_layers > 0 + ) + name = "model.embed_tokens.weight" + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight) + loaded_params.add(name) + params_need_to_load = set(params_dict.keys()) + if params_need_to_load != loaded_params: + missing_params = list(params_need_to_load - loaded_params) + param_name_example = missing_params[0] + raise RuntimeError( + f"Some parameters like {param_name_example} are not in the checkpoint and will falsely use random initialization" + ) + return loaded_params + + def _rewrite_spec_layer_name(self, spec_layer: Optional[int], name: str) -> str: + """ + Rewrite the weight name to match the format of the original model. + Add .mtp_block for modules in transformer layer block for spec layer + """ + if spec_layer is None: + return name + + # Some checkpoints place MTP weights under "model.layers..transformer.*". + # Our modules use "model.layers..*", so drop the ".transformer." segment. + transformer_prefix = f"model.layers.{spec_layer}.transformer." + if name.startswith(transformer_prefix): + name = name.replace(".transformer.", ".", 1) + + spec_layer_weight_names = [ + "embed_tokens", + "enorm", + "hnorm", + "eh_proj", + "shared_head", + ] + spec_layer_weight = False + for weight_name in spec_layer_weight_names: + if weight_name in name: + spec_layer_weight = True + break + if not spec_layer_weight: + # treat rest weights as weights for transformer layer block + name = name.replace( + f"model.layers.{spec_layer}.", f"model.layers.{spec_layer}.mtp_block." + ) + + # NEW: drop "layers.." from the rewritten name (minimal change). + layers_prefix = f"model.layers.{spec_layer}." + if name.startswith(layers_prefix): + name = name.replace(layers_prefix, "model.", 1) + + return name + + +EntryClass = [Step3p5MTP] diff --git a/python/sglang/srt/parser/reasoning_parser.py b/python/sglang/srt/parser/reasoning_parser.py index 455101b5c..cd346625a 100644 --- a/python/sglang/srt/parser/reasoning_parser.py +++ b/python/sglang/srt/parser/reasoning_parser.py @@ -384,6 +384,7 @@ class ReasoningParser: "minimax": Qwen3Detector, "minimax-append-think": MiniMaxAppendThinkDetector, "step3": DeepSeekR1Detector, + "step3p5": DeepSeekR1Detector, "nano_v3": NanoV3Detector, "interns1": Qwen3Detector, } diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index dedb90c6c..6a1eb50d6 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1406,6 +1406,26 @@ class ServerArgs: logger.warning( "Disable hybrid SWA memory for MiMoV2FlashForCausalLM model with hierarchical cache" ) + elif "Step3p5ForCausalLM" in model_arch: + if self.speculative_algorithm == "EAGLE": + self.enable_multi_layer_eagle = True + logger.info( + "Enable multi-layer EAGLE speculative decoding for Step3p5ForCausalLM model." + ) + if not envs.SGLANG_ENABLE_SPEC_V2.get(): + envs.SGLANG_ENABLE_SPEC_V2.set(True) + logger.warning( + "Spec v2 is enabled for multi-layer EAGLE speculative decoding." + ) + if self.enable_hierarchical_cache: + self.swa_full_tokens_ratio = 1.0 + logger.warning( + "Reset swa_full_tokens_ratio to 1.0 for Step3p5ForCausalLM model with hierarchical cache" + ) + self.disable_hybrid_swa_memory = True + logger.warning( + "Disable hybrid SWA memory for Step3p5ForCausalLM model with hierarchical cache" + ) elif "Llama4" in model_arch and self.device != "cpu": # Auto-select attention backend for Llama4 if not specified if self.attention_backend is None: diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 44bb2f0de..8320eaa7f 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -466,11 +466,12 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker): draft_logits_output.topk_index, ) else: - draft_logits_output, _ = self.draft_runner_list[step].forward( + draft_logits_output = self.draft_runner_list[step].forward( forward_batch, skip_attn_backend_init=True ) probs = torch.softmax( - draft_logits_output.next_token_logits[select_index], dim=-1 + draft_logits_output.logits_output.next_token_logits[select_index], + dim=-1, ) ret_topk_p, ret_topk_index = fast_topk(probs, self.topk, dim=-1) if forward_batch.extend_seq_lens is not None: diff --git a/python/sglang/srt/utils/hf_transformers_utils.py b/python/sglang/srt/utils/hf_transformers_utils.py index cc1534c56..2efb1f7b5 100644 --- a/python/sglang/srt/utils/hf_transformers_utils.py +++ b/python/sglang/srt/utils/hf_transformers_utils.py @@ -63,6 +63,7 @@ from sglang.srt.configs import ( NemotronHConfig, Olmo3Config, Qwen3NextConfig, + Step3p5Config, Step3VLConfig, ) from sglang.srt.configs.deepseek_ocr import DeepseekVLV2Config @@ -95,6 +96,7 @@ _CONFIG_REGISTRY: List[Type[PretrainedConfig]] = [ JetNemotronConfig, JetVLMConfig, KimiK25Config, + Step3p5Config, ] _CONFIG_REGISTRY = {