diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index bdd091e66..5406522c1 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -438,23 +438,22 @@ class ModelConfig: if is_deepseek_nsa(self.hf_text_config) else None ) - - if "Glm4MoeLiteForCausalLM" in self.hf_config.architectures: - self.scaling = 1 - self.hf_config.rope_scaling = None - else: - # Handle rope scaling with yarn - self.scaling = 1 / math.sqrt( - self.qk_nope_head_dim + self.qk_rope_head_dim + # Handle rope scaling + self.scaling = 1 / math.sqrt(self.qk_nope_head_dim + self.qk_rope_head_dim) + # in transformers v5, rope_scaling is just rope_parameters for backward compatibility + rope_scaling = self.hf_text_config.rope_scaling + if rope_scaling: + # v5 uses "rope_type", v4 uses "type" + rope_type = ( + rope_scaling.get("rope_type") + or rope_scaling.get("type") + or "default" ) - if self.hf_text_config.rope_scaling: - mscale_all_dim = self.hf_text_config.rope_scaling.get( - "mscale_all_dim", False - ) - scaling_factor = self.hf_text_config.rope_scaling["factor"] + if rope_type != "default": + mscale_all_dim = rope_scaling.get("mscale_all_dim", False) + scaling_factor = rope_scaling["factor"] mscale = yarn_get_mscale(scaling_factor, float(mscale_all_dim)) self.scaling = self.scaling * mscale * mscale - elif "MiniCPM3ForCausalLM" in self.hf_config.architectures: self.head_dim = 128 self.attention_arch = AttentionArch.MLA diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 7dbe7e726..7de4f472f 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -353,6 +353,7 @@ class Envs: # NSA Backend SGLANG_NSA_FUSE_TOPK = EnvBool(True) SGLANG_NSA_ENABLE_MTP_PRECOMPUTE_METADATA = EnvBool(True) + SGLANG_NSA_FORCE_MLA = EnvBool(False) # sgl-kernel SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index cf362e712..c270b1e4e 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -145,6 +145,7 @@ class Indexer(MultiPlatformOp): scale_fmt: Optional[str], block_size: int = 128, rope_scaling: Optional[Dict[str, Any]] = None, + is_neox_style: bool = True, prefix: str = "", quant_config: Optional[QuantizationConfig] = None, alt_stream: Optional[torch.cuda.Stream] = None, @@ -202,7 +203,7 @@ class Indexer(MultiPlatformOp): max_position=max_position_embeddings, base=rope_theta, # type: ignore rope_scaling=rope_scaling, - is_neox_style=True, + is_neox_style=is_neox_style, device=get_global_server_args().device, ) self.block_size = block_size diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index b6bb7d08b..2de045c82 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -300,6 +300,8 @@ class NativeSparseAttnBackend( self.req_to_token = model_runner.req_to_token_pool.req_to_token self.use_mha: bool = False + # Force NSA prefill to use MLA (i.e. disable MHA_ONE_SHOT), controlled by env var. + self._force_attn_forward_mla: bool = envs.SGLANG_NSA_FORCE_MLA.get() self.nsa_prefill_impl: _NSA_IMPL_T = ( model_runner.server_args.nsa_prefill_backend ) @@ -1837,6 +1839,8 @@ class NativeSparseAttnBackend( ) else: self.use_mha = False # Decode/verify always use MLA + if self._force_attn_forward_mla: + self.use_mha = False # Set MLA implementation only if not using MHA if not self.use_mha and self.enable_auto_select_prefill_impl: diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 988ed91e8..e0583d308 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -1150,6 +1150,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): ) if self.use_nsa: + is_neox_style = not getattr(config, "indexer_rope_interleave", False) self.indexer = Indexer( hidden_size=hidden_size, index_n_heads=get_nsa_index_n_heads(config), @@ -1162,6 +1163,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): scale_fmt="ue8m0", block_size=128, rope_scaling=rope_scaling, + is_neox_style=is_neox_style, prefix=add_prefix("indexer", prefix), quant_config=quant_config, layer_id=layer_id, @@ -1191,13 +1193,14 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): self.kv_a_layernorm = RMSNorm(self.kv_lora_rank, eps=config.rms_norm_eps) if not skip_rope: + is_neox_style = not getattr(config, "rope_interleave", True) self.rotary_emb = get_rope_wrapper( qk_rope_head_dim, rotary_dim=qk_rope_head_dim, max_position=max_position_embeddings, base=rope_theta, rope_scaling=rope_scaling, - is_neox_style=False, + is_neox_style=is_neox_style, device=get_global_server_args().device, ) @@ -2227,9 +2230,15 @@ class DeepseekV2DecoderLayer(nn.Module): super().__init__() self.hidden_size = config.hidden_size self.config = config - rope_theta = getattr(config, "rope_theta", 10000) - rope_scaling = getattr(config, "rope_scaling", None) - max_position_embeddings = getattr(config, "max_position_embeddings", 8192) + if hasattr(config, "rope_parameters"): + rope_theta = config.rope_parameters.get("rope_theta") + assert rope_theta is not None, f"rope_theta not found in config: {config}" + rope_type = config.rope_parameters.get("rope_type") + rope_scaling = config.rope_parameters if rope_type != "default" else None + else: + rope_theta = config.rope_theta + rope_scaling = config.rope_scaling + max_position_embeddings = config.max_position_embeddings self.speculative_algorithm = SpeculativeAlgorithm.from_string( get_global_server_args().speculative_algorithm )