From 495ae9aaa609d39fd4e294d4f64faf52e29f50a6 Mon Sep 17 00:00:00 2001 From: Elizaveta Martirosian Date: Wed, 15 Jul 2026 15:23:04 +0300 Subject: [PATCH] Fix Ministral3 accuracy issue by aligning YaRN RoPE scaling with Transformers implementation (#31232) Co-authored-by: Elizaveta Martirosian Co-authored-by: ronnie_zheng --- .../srt/layers/rotary_embedding/factory.py | 9 ++++++++- .../sglang/srt/layers/rotary_embedding/yarn.py | 16 ++++++++++++++-- 2 files changed, 22 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/layers/rotary_embedding/factory.py b/python/sglang/srt/layers/rotary_embedding/factory.py index d058ea08a..a5f4d9bb3 100644 --- a/python/sglang/srt/layers/rotary_embedding/factory.py +++ b/python/sglang/srt/layers/rotary_embedding/factory.py @@ -243,7 +243,14 @@ def get_rope( k: v for k, v in rope_scaling.items() if k - in ("extrapolation_factor", "attn_factor", "beta_fast", "beta_slow") + in ( + "extrapolation_factor", + "attn_factor", + "beta_fast", + "beta_slow", + "mscale", + "mscale_all_dim", + ) } extra_kwargs["truncate"] = rope_scaling.get("truncate", True) if "mrope_section" in rope_scaling: diff --git a/python/sglang/srt/layers/rotary_embedding/yarn.py b/python/sglang/srt/layers/rotary_embedding/yarn.py index 61f648ac0..e2ccb82f5 100644 --- a/python/sglang/srt/layers/rotary_embedding/yarn.py +++ b/python/sglang/srt/layers/rotary_embedding/yarn.py @@ -83,6 +83,8 @@ class YaRNScalingRotaryEmbedding(RotaryEmbedding): beta_fast: int = 32, beta_slow: int = 1, truncate: bool = True, + mscale: float = None, + mscale_all_dim: float = None, ) -> None: self.scaling_factor = scaling_factor self.extrapolation_factor = extrapolation_factor @@ -90,8 +92,18 @@ class YaRNScalingRotaryEmbedding(RotaryEmbedding): self.beta_fast = beta_fast self.beta_slow = beta_slow self.truncate = truncate - # Get n-d magnitude scaling corrected for interpolation - self.mscale = float(yarn_get_mscale_simple(self.scaling_factor) * attn_factor) + + if mscale is not None and mscale_all_dim is not None: + # Match Hugging Face's YaRN RoPE scaling (supports mscale/mscale_all_dim) + self.mscale = float( + yarn_get_mscale(self.scaling_factor, mscale) + / yarn_get_mscale(self.scaling_factor, mscale_all_dim) + ) + else: + # Get n-d magnitude scaling corrected for interpolation + self.mscale = float(yarn_get_mscale_simple(self.scaling_factor)) + self.mscale *= attn_factor + super().__init__( head_size, rotary_dim, max_position_embeddings, base, is_neox_style, dtype )