diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index 314c897ab..e4c08ab6d 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -2129,11 +2129,9 @@ class NativeSparseAttnBackend( # disable for MTP self.nsa_kv_cache_store_fp8 and self.nsa_prefill_impl == "flashmla_sparse" + and forward_mode == ForwardMode.EXTEND ): topk_transform_method = TopkTransformMethod.RAGGED - - if forward_mode is not None and (forward_mode.is_decode_or_idle()): - topk_transform_method = TopkTransformMethod.PAGED else: topk_transform_method = TopkTransformMethod.PAGED return topk_transform_method diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index d08cd3886..b65922d52 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1491,8 +1491,10 @@ class ServerArgs: self.nsa_decode_backend = "tilelang" elif kv_cache_dtype == "fp8_e4m3": if major >= 10: - self.nsa_prefill_backend = "trtllm" - self.nsa_decode_backend = "trtllm" + if not user_set_prefill: + self.nsa_prefill_backend = "trtllm" + if not user_set_decode: + self.nsa_decode_backend = "trtllm" else: # flashmla_auto dispatches to flashmla_sparse/flashmla_kv based on hardware and heuristics if not user_set_prefill: