[Fix] Restore layer-level DSV4 RoPE policy (#34788)

This commit is contained in:
Baizhou Zhang
2026-08-13 21:26:45 -07:00
committed by GitHub
parent e1c4db9621
commit 0a6bbbe128
2 changed files with 163 additions and 8 deletions
+18 -8
View File
@@ -607,15 +607,23 @@ class MqaAttentionBase(nn.Module):
from sglang.kernels.ops.attention.deepseek_v4_rope import precompute_freqs_cis
rope_theta, rope_scaling = get_rope_config(config)
self.rope_scaling = rope_scaling
scaling = rope_scaling or {}
self.rope_scaling = dict(rope_scaling) if rope_scaling else None
scaling = self.rope_scaling or {}
# RoPE is selected at layer granularity in the reference model. Pure
# SWA layers use the main unscaled RoPE, while C4/C128 layers use the
# compressed YaRN RoPE for Q, their SWA branch, and compressed KV.
self.rope_base = (
config.compress_rope_theta if self.compress_ratio else rope_theta
)
original_seq_len: int = (
rope_original_seq_len
if rope_original_seq_len is not None
else scaling["original_max_position_embeddings"]
else (
scaling["original_max_position_embeddings"]
if self.compress_ratio
else 0
)
)
freqs_cis = precompute_freqs_cis(
dim=self.qk_rope_head_dim,
@@ -693,14 +701,16 @@ class MQALayer(MqaAttentionBase):
compress_ratio=compress_ratio_override,
)
if self.rope_scaling:
self.rope_scaling["rope_type"] = "deepseek_yarn"
active_rope_scaling = None
if self.compress_ratio in (4, 128):
active_rope_scaling = dict(self.rope_scaling or {})
active_rope_scaling["rope_type"] = "deepseek_yarn"
self.rotary_emb = get_rope_wrapper(
head_size=self.rope_head_dim,
rotary_dim=self.rope_head_dim,
max_position=config.max_position_embeddings,
base=self.rope_base,
rope_scaling=self.rope_scaling,
rope_scaling=active_rope_scaling,
is_neox_style=False,
device=get_device().device,
)
@@ -743,7 +753,7 @@ class MQALayer(MqaAttentionBase):
head_dim=self.head_dim,
rotate=False,
prefix=add_prefix("compressor", prefix),
rotary_emb=getattr(self, "rotary_emb", None),
rotary_emb=self.rotary_emb,
)
if self.compress_ratio == 4:
self.indexer = C4Indexer(
@@ -753,7 +763,7 @@ class MQALayer(MqaAttentionBase):
quant_config=quant_config,
prefix=add_prefix("indexer", prefix),
alt_streams=self.alt_streams_indexer,
rotary_emb=getattr(self, "rotary_emb", None),
rotary_emb=self.rotary_emb,
)
self.attn_mqa = RadixAttention(
@@ -0,0 +1,145 @@
"""Unit tests for DeepSeek-V4 layer-level RoPE selection."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from torch import nn
import sglang.srt.models.deepseek_v4 as deepseek_v4
from sglang.kernels.ops.attention.deepseek_v4_rope import precompute_freqs_cis
from sglang.srt.environ import envs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class _ModuleStub(nn.Module):
def __init__(self, *args, **kwargs):
super().__init__()
class _RMSNormStub(_ModuleStub):
def __init__(self, hidden_size, *args, **kwargs):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
class _RoPEConsumerStub(_ModuleStub):
def __init__(self, *args, freqs_cis, rotary_emb=None, **kwargs):
super().__init__()
self.freqs_cis = freqs_cis
self.rotary_emb = rotary_emb
class _RotaryEmbeddingStub(_ModuleStub):
def __init__(self, *args, **kwargs):
super().__init__()
self.rope_scaling = kwargs["rope_scaling"]
class TestDeepseekV4RoPEPolicy(CustomTestCase):
@staticmethod
def _config(compress_ratio):
return SimpleNamespace(
hidden_size=16,
qk_rope_head_dim=64,
qk_nope_head_dim=64,
head_dim=128,
num_attention_heads=2,
num_key_value_heads=1,
o_groups=2,
q_lora_rank=8,
o_lora_rank=8,
rms_norm_eps=1e-6,
compress_ratios=[compress_ratio],
rope_theta=10_000,
compress_rope_theta=160_000,
max_position_embeddings=128,
rope_scaling={
"original_max_position_embeddings": 65_536,
"factor": 16.0,
"beta_fast": 32,
"beta_slow": 1,
"type": "yarn",
},
)
def _make_layer(self, compress_ratio):
parallel = SimpleNamespace(attn_tp_rank=0, attn_tp_size=1, tp_size=1)
device = SimpleNamespace(device=torch.device("cpu"))
with (
envs.SGLANG_OPT_FUSE_WQA_WKV.override(False),
envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.override(False),
patch.object(deepseek_v4, "_FP8_WO_A_GEMM", False),
patch.object(deepseek_v4, "_is_hip", False),
patch.object(deepseek_v4, "_is_npu", False),
patch.object(deepseek_v4, "is_dsa_enable_prefill_cp", return_value=False),
patch.object(deepseek_v4, "get_parallel", return_value=parallel),
patch.object(deepseek_v4, "get_device", return_value=device),
patch.object(deepseek_v4, "ReplicatedLinear", _ModuleStub),
patch.object(deepseek_v4, "ColumnParallelLinear", _ModuleStub),
patch.object(deepseek_v4, "RowParallelLinear", _ModuleStub),
patch.object(deepseek_v4, "RMSNorm", _RMSNormStub),
patch.object(
deepseek_v4,
"get_rope_wrapper",
side_effect=lambda *args, **kwargs: _RotaryEmbeddingStub(
*args, **kwargs
),
),
patch.object(deepseek_v4, "Compressor", _RoPEConsumerStub),
patch.object(deepseek_v4, "C4Indexer", _RoPEConsumerStub),
patch.object(deepseek_v4, "RadixAttention", _ModuleStub),
):
return deepseek_v4.MQALayer(
config=self._config(compress_ratio),
layer_id=0,
)
def test_pure_swa_layer_uses_unscaled_main_rope(self):
layer = self._make_layer(0)
expected = precompute_freqs_cis(
dim=64,
seqlen=128,
original_seq_len=0,
base=10_000,
factor=16.0,
beta_fast=32,
beta_slow=1,
)
torch.testing.assert_close(layer.freqs_cis, expected)
self.assertIsNone(layer.rotary_emb.rope_scaling)
self.assertIsNone(layer.compressor)
self.assertIsNone(layer.indexer)
def test_c4_and_c128_layers_share_yarn_compress_rope(self):
for compress_ratio in (4, 128):
with self.subTest(compress_ratio=compress_ratio):
layer = self._make_layer(compress_ratio)
expected_compressed = precompute_freqs_cis(
dim=64,
seqlen=128,
original_seq_len=65_536,
base=160_000,
factor=16.0,
beta_fast=32,
beta_slow=1,
)
torch.testing.assert_close(layer.freqs_cis, expected_compressed)
self.assertIs(layer.compressor.freqs_cis, layer.freqs_cis)
self.assertIs(layer.compressor.rotary_emb, layer.rotary_emb)
self.assertEqual(
layer.rotary_emb.rope_scaling["rope_type"], "deepseek_yarn"
)
if compress_ratio == 4:
self.assertIs(layer.indexer.freqs_cis, layer.compressor.freqs_cis)
self.assertIs(layer.indexer.rotary_emb, layer.compressor.rotary_emb)
if __name__ == "__main__":
unittest.main()