import builtins import unittest from unittest.mock import patch import torch from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa_backend import NativeSparseAttnBackend from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=1, suite="stage-a-test-cpu") class TestNSADequantOnlyTopK(unittest.TestCase): def test_env_gate_is_defined_and_default_off(self): self.assertFalse(envs.SGLANG_NSA_DEQUANT_ONLY_TOPK.get()) self.assertFalse(envs.SGLANG_NSA_DEQUANT_ONLY_TOPK_VERIFY.get()) def test_dequant_topk_inplace_falls_back_to_full_dequant_when_kernel_missing(self): backend = object.__new__(NativeSparseAttnBackend) kv_cache = torch.empty((4, 656), dtype=torch.uint8) page_table = torch.tensor([3, 1, 2, 0], dtype=torch.int32) topk_indices = torch.tensor([[0, 2, -1], [1, 3, 2]], dtype=torch.int32) expected = torch.empty((4, 1, 576), dtype=torch.bfloat16) real_import = builtins.__import__ def fake_import(name, *args, **kwargs): if name == "tai_kernel.nsa_prefill": raise ImportError("tai kernel intentionally unavailable") return real_import(name, *args, **kwargs) with patch("builtins.__import__", side_effect=fake_import), patch( "sglang.srt.layers.attention.nsa_backend.dequantize_k_cache_paged", return_value=expected, ) as dequant: out = backend._dequant_topk_inplace(kv_cache, page_table, topk_indices) self.assertIs(out, expected) dequant.assert_called_once_with(kv_cache, page_table) if __name__ == "__main__": unittest.main()