[Fix] Restore layer-level DSV4 RoPE policy (#34788)
This commit is contained in:
@@ -607,15 +607,23 @@ class MqaAttentionBase(nn.Module):
|
|||||||
from sglang.kernels.ops.attention.deepseek_v4_rope import precompute_freqs_cis
|
from sglang.kernels.ops.attention.deepseek_v4_rope import precompute_freqs_cis
|
||||||
|
|
||||||
rope_theta, rope_scaling = get_rope_config(config)
|
rope_theta, rope_scaling = get_rope_config(config)
|
||||||
self.rope_scaling = rope_scaling
|
self.rope_scaling = dict(rope_scaling) if rope_scaling else None
|
||||||
scaling = rope_scaling or {}
|
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 = (
|
self.rope_base = (
|
||||||
config.compress_rope_theta if self.compress_ratio else rope_theta
|
config.compress_rope_theta if self.compress_ratio else rope_theta
|
||||||
)
|
)
|
||||||
original_seq_len: int = (
|
original_seq_len: int = (
|
||||||
rope_original_seq_len
|
rope_original_seq_len
|
||||||
if rope_original_seq_len is not None
|
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(
|
freqs_cis = precompute_freqs_cis(
|
||||||
dim=self.qk_rope_head_dim,
|
dim=self.qk_rope_head_dim,
|
||||||
@@ -693,14 +701,16 @@ class MQALayer(MqaAttentionBase):
|
|||||||
compress_ratio=compress_ratio_override,
|
compress_ratio=compress_ratio_override,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.rope_scaling:
|
active_rope_scaling = None
|
||||||
self.rope_scaling["rope_type"] = "deepseek_yarn"
|
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(
|
self.rotary_emb = get_rope_wrapper(
|
||||||
head_size=self.rope_head_dim,
|
head_size=self.rope_head_dim,
|
||||||
rotary_dim=self.rope_head_dim,
|
rotary_dim=self.rope_head_dim,
|
||||||
max_position=config.max_position_embeddings,
|
max_position=config.max_position_embeddings,
|
||||||
base=self.rope_base,
|
base=self.rope_base,
|
||||||
rope_scaling=self.rope_scaling,
|
rope_scaling=active_rope_scaling,
|
||||||
is_neox_style=False,
|
is_neox_style=False,
|
||||||
device=get_device().device,
|
device=get_device().device,
|
||||||
)
|
)
|
||||||
@@ -743,7 +753,7 @@ class MQALayer(MqaAttentionBase):
|
|||||||
head_dim=self.head_dim,
|
head_dim=self.head_dim,
|
||||||
rotate=False,
|
rotate=False,
|
||||||
prefix=add_prefix("compressor", prefix),
|
prefix=add_prefix("compressor", prefix),
|
||||||
rotary_emb=getattr(self, "rotary_emb", None),
|
rotary_emb=self.rotary_emb,
|
||||||
)
|
)
|
||||||
if self.compress_ratio == 4:
|
if self.compress_ratio == 4:
|
||||||
self.indexer = C4Indexer(
|
self.indexer = C4Indexer(
|
||||||
@@ -753,7 +763,7 @@ class MQALayer(MqaAttentionBase):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("indexer", prefix),
|
prefix=add_prefix("indexer", prefix),
|
||||||
alt_streams=self.alt_streams_indexer,
|
alt_streams=self.alt_streams_indexer,
|
||||||
rotary_emb=getattr(self, "rotary_emb", None),
|
rotary_emb=self.rotary_emb,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.attn_mqa = RadixAttention(
|
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()
|
||||||
Reference in New Issue
Block a user