[Refactor] Rename NSA → DSA: user-facing aliases, file/class/import rename (#25821)

Co-authored-by: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-05-20 00:18:04 -07:00
committed by GitHub
co-authored by Claude Sonnet 4.6
parent da6d549ab2
commit 8131641bc6
162 changed files with 11298 additions and 10740 deletions
@@ -11,15 +11,15 @@ from sglang.test.ci.ci_register import register_cuda_ci
_dp_attn.get_attention_tp_size = lambda: 1 # TP size = 1 for unit test
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.nsa.nsa_indexer import (
from sglang.srt.layers.attention.dsa.dsa_indexer import (
BaseIndexerMetadata,
Indexer,
rotate_activation,
)
from sglang.srt.layers.attention.nsa_backend import NativeSparseAttnBackend
from sglang.srt.layers.attention.dsa_backend import DeepseekSparseAttnBackend
from sglang.srt.layers.layernorm import LayerNorm
from sglang.srt.layers.linear import LinearBase
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.test_utils import CustomTestCase
@@ -136,7 +136,7 @@ class MockIndexerMetadata(BaseIndexerMetadata):
"""Return: seq lens for each batch."""
return torch.tensor(self.seq_lens, dtype=torch.int32, device=self.device)
def get_nsa_extend_len_cpu(self) -> List[int]:
def get_dsa_extend_len_cpu(self) -> List[int]:
"""
Return: extend seq lens for each batch.
"""
@@ -179,7 +179,7 @@ class MockModelRunner:
max_context_len = self.config["context_len"]
max_batch_size = self.config["max_bs"]
# Create mock hf_config for NSA - instantiate it as an object, not a type
# Create mock hf_config for DSA - instantiate it as an object, not a type
hf_config = type(
"HfConfig",
(),
@@ -224,9 +224,9 @@ class MockModelRunner:
},
)()
# Create NSATokenToKVPool
# Create DSATokenToKVPool
max_total_num_tokens = max_batch_size * max_context_len
self.token_to_kv_pool = NSATokenToKVPool(
self.token_to_kv_pool = DSATokenToKVPool(
size=max_total_num_tokens,
page_size=self.config["page_size"],
dtype=self.config["kv_cache_dtype"],
@@ -239,7 +239,7 @@ class MockModelRunner:
kv_cache_dim=self.config["kv_lora_rank"] + self.config["qk_rope_head_dim"],
)
# Required by backend with NSA-specific attributes
# Required by backend with DSA-specific attributes
self.server_args = type(
"ServerArgs",
(),
@@ -248,21 +248,21 @@ class MockModelRunner:
"speculative_eagle_topk": None,
"speculative_num_draft_tokens": 0,
"enable_deterministic_inference": False,
"nsa_prefill_backend": "flashmla_sparse",
"nsa_decode_backend": "fa3",
"dsa_prefill_backend": "flashmla_sparse",
"dsa_decode_backend": "fa3",
},
)()
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
class TestNSAIndexer(CustomTestCase):
class TestDSAIndexer(CustomTestCase):
@classmethod
def setUpClass(cls):
"""Set up global server args for testing."""
server_args = ServerArgs(model_path="dummy")
server_args.enable_dp_attention = False
server_args.nsa_prefill_backend = "flashmla_sparse"
server_args.nsa_decode_backend = "flashmla_sparse"
server_args.dsa_prefill_backend = "flashmla_sparse"
server_args.dsa_decode_backend = "flashmla_sparse"
set_global_server_args_for_scheduler(server_args)
# Check GPU capability for FP8
@@ -289,7 +289,7 @@ class TestNSAIndexer(CustomTestCase):
if config_override:
config.update(config_override)
self.model_runner = MockModelRunner(config)
self.backend = NativeSparseAttnBackend(self.model_runner)
self.backend = DeepseekSparseAttnBackend(self.model_runner)
def _create_indexer(self, **kwargs):
"""Create an Indexer instance with default parameters."""
@@ -417,7 +417,7 @@ class TestNSAIndexer(CustomTestCase):
"Output should have padding or exact topk size",
)
@patch("sglang.srt.layers.attention.nsa.nsa_indexer.deep_gemm")
@patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm")
def test_indexer_basic_creation(self, mock_deep_gemm):
"""Test basic indexer creation and initialization."""
mock_deep_gemm.get_num_sms.return_value = 132
@@ -431,8 +431,8 @@ class TestNSAIndexer(CustomTestCase):
self.assertEqual(indexer.index_topk, self.config["index_topk"])
self.assertEqual(indexer.layer_id, self.config["layer_id"])
@patch("sglang.srt.layers.attention.nsa.nsa_indexer.deep_gemm")
@patch("sglang.srt.layers.attention.nsa.triton_kernel.act_quant")
@patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm")
@patch("sglang.srt.layers.attention.dsa.triton_kernel.act_quant")
def test_forward_extend_mode(self, mock_act_quant, mock_deep_gemm):
"""Test indexer forward pass in extend mode."""
if not self.supports_fp8:
@@ -513,8 +513,8 @@ class TestNSAIndexer(CustomTestCase):
topk_indices, self.batch_size, self.seq_len, self.config["index_topk"]
)
@patch("sglang.srt.layers.attention.nsa.nsa_indexer.deep_gemm")
@patch("sglang.srt.layers.attention.nsa.triton_kernel.act_quant")
@patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm")
@patch("sglang.srt.layers.attention.dsa.triton_kernel.act_quant")
def test_forward_decode_mode(self, mock_act_quant, mock_deep_gemm):
"""Test indexer forward pass in decode mode."""
if not self.supports_fp8:
@@ -627,7 +627,7 @@ class TestNSAIndexer(CustomTestCase):
self.assertEqual(topk_indices.shape, (batch_size, topk))
# TODO: enable this test after indexer accuracy aligned
# @patch("sglang.srt.layers.attention.nsa.nsa_indexer.deep_gemm")
# @patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm")
# def test_indexer_with_different_topk(self, mock_deep_gemm):
# """Test indexer with different topk values."""
# mock_deep_gemm.get_num_sms.return_value = 132
@@ -637,7 +637,7 @@ class TestNSAIndexer(CustomTestCase):
# indexer = self._create_indexer(index_topk=topk)
# self.assertEqual(indexer.index_topk, topk)
@patch("sglang.srt.layers.attention.nsa.nsa_indexer.deep_gemm")
@patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm")
def test_indexer_with_fused_wk(self, mock_deep_gemm):
"""Test indexer creation with fused wk and weights projection."""
mock_deep_gemm.get_num_sms.return_value = 132
@@ -647,7 +647,7 @@ class TestNSAIndexer(CustomTestCase):
indexer = self._create_indexer()
self.assertIsNotNone(indexer)
@patch("sglang.srt.layers.attention.nsa.nsa_indexer.deep_gemm")
@patch("sglang.srt.layers.attention.dsa.dsa_indexer.deep_gemm")
def test_indexer_with_alt_stream(self, mock_deep_gemm):
"""Test indexer creation with alternative CUDA stream."""
mock_deep_gemm.get_num_sms.return_value = 132