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
This commit is contained in:
laoyao0822
2026-06-10 07:00:43 +08:00
parent c8c736e75c
commit f78c414b64
3 changed files with 101 additions and 3 deletions

View File

@@ -422,6 +422,11 @@ class Envs:
# NSA Backend
SGLANG_NSA_FUSE_TOPK = EnvBool(True)
# Dequantize only the KV rows referenced by sparse-attention topk_indices
# instead of full materialized FP8 K cache. Default OFF; VERIFY runs the full
# dequant reference and compares gathered rows, so it is for validation only.
SGLANG_NSA_DEQUANT_ONLY_TOPK = EnvBool(False)
SGLANG_NSA_DEQUANT_ONLY_TOPK_VERIFY = EnvBool(False)
SGLANG_NSA_ENABLE_MTP_PRECOMPUTE_METADATA = EnvBool(True)
SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True)
SGLANG_NSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD = EnvInt(2048)

View File

@@ -2641,9 +2641,17 @@ class NativeSparseAttnBackend(
nvtx_source="mla.ragged_partial_current_sync",
)
)
kv_cache = dequantize_k_cache_paged(
kv_cache, page_table_1_flattened
)
if envs.SGLANG_NSA_DEQUANT_ONLY_TOPK.get():
kv_cache = self._dequant_topk_inplace(
kv_cache,
page_table_1_flattened,
topk_indices,
verify=envs.SGLANG_NSA_DEQUANT_ONLY_TOPK_VERIFY.get(),
)
else:
kv_cache = dequantize_k_cache_paged(
kv_cache, page_table_1_flattened
)
else:
kv_cache = _cat([k, k_rope], dim=-1)
page_table_1 = topk_indices
@@ -2720,6 +2728,47 @@ class NativeSparseAttnBackend(
index_prefetcher.wait_attention_window()
return attn_output
def _dequant_topk_inplace(
self,
kv_cache_fp8: torch.Tensor,
page_table_1_flattened: torch.Tensor,
topk_indices: torch.Tensor,
verify: bool = False,
) -> torch.Tensor:
"""Dequantize only topk-referenced FP8 KV rows at their original row ids.
The output keeps the full logical row count ``[num_kv_rows, 1, 576]`` so
sparse attention can keep using the unchanged ``topk_indices``. This is
the final syh in-place contract, not the earlier compact/remap variant:
no data-dependent output shape, no topk remap, and no CPU synchronization.
"""
try:
from tai_kernel.nsa_prefill import nsa_dequant_topk_inplace
except ImportError:
nsa_dequant_topk_inplace = None
if nsa_dequant_topk_inplace is not None:
out = nsa_dequant_topk_inplace(
kv_cache_fp8,
page_table_1_flattened,
topk_indices,
)
else:
out = dequantize_k_cache_paged(kv_cache_fp8, page_table_1_flattened)
if verify:
kv_ref = dequantize_k_cache_paged(kv_cache_fp8, page_table_1_flattened)
flat = topk_indices.reshape(-1)
valid = (flat >= 0) & (flat < int(page_table_1_flattened.shape[0]))
rows = flat[valid].long()
ref_rows = kv_ref.reshape(-1, kv_ref.shape[-1])[rows]
new_rows = out.reshape(-1, out.shape[-1])[rows]
assert torch.equal(ref_rows, new_rows), (
"[SGLANG_NSA_DEQUANT_ONLY_TOPK_VERIFY] in-place dequant != full "
f"dequant (N={int(page_table_1_flattened.shape[0])})"
)
return out
def forward_decode(
self,
q: torch.Tensor,