High cache-hit NSA sparse attention only gathers rows selected by topk_indices, but the FP8 RAGGED path dequantized the whole materialized K buffer first. Wire the syh in-place topk-only dequant path behind an explicit env gate so the sparse path can dequant referenced rows at their original row ids while keeping topk_indices unchanged. The old full-dequant path remains the default and the verification gate stays available for small correctness checks only. Constraint: Production runs should enable only SGLANG_NSA_DEQUANT_ONLY_TOPK=1; VERIFY also runs the full-dequant reference and is too expensive for perf tests. Rejected: Compact/remap topk buffer | data-dependent shape and remap add synchronization/aliasing risk; syh final in-place contract avoids both. Confidence: medium Scope-risk: moderate Directive: Do not enable SGLANG_NSA_DEQUANT_ONLY_TOPK_VERIFY in throughput runs; use it only for small correctness probes. Tested: python -m py_compile python/sglang/srt/environ.py python/sglang/srt/layers/attention/nsa_backend.py test/registered/unit/layers/test_nsa_dequant_only_topk.py Tested: git diff --check Tested: remote g0034 container PYTHONPATH=python python -m pytest -q test/registered/unit/layers/test_nsa_dequant_only_topk.py -> 2 passed Tested: remote g0034 CUDA smoke for tai_kernel.nsa_prefill.nsa_dequant_topk_inplace topk rows matched torch reference Not-tested: Full GSM8K or replay ETE with SGLANG_NSA_DEQUANT_ONLY_TOPK=1
45 lines
1.6 KiB
Python
45 lines
1.6 KiB
Python
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()
|