Files
sglang/test/registered/unit/layers/test_nsa_dequant_only_topk.py
laoyao0822 f78c414b64 Skip unused FP8 NSA K-cache dequant work
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
2026-06-10 07:00:43 +08:00

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()