diff --git a/docs/basic_usage/deepseek_v32.md b/docs/basic_usage/deepseek_v32.md index db1d0e71e..095060a7f 100644 --- a/docs/basic_usage/deepseek_v32.md +++ b/docs/basic_usage/deepseek_v32.md @@ -66,14 +66,14 @@ To serve GLM-5, just replace the `--model` argument with `zai-org/GLM-5-FP8`. - **Choices of Attention Kernels**: The attention backend is automatically set to `nsa` attention backend for DeepSeek V3.2 model. In this backend, different kernels for sparse prefilling/decoding are implemented, which can be specified by `--nsa-prefill-backend` and `--nsa-decode-backend` server arguments. The choices of nsa prefill/decode attention kernels include: - `flashmla_sparse`: `flash_mla_sparse_fwd` kernel from `flash_mla` library. Can run on both Hopper and Blackwell GPUs. It requires bf16 q, kv inputs. - `flashmla_kv`: `flash_mla_with_kvcache` kernel from `flash_mla` library. Can run on both Hopper and Blackwell GPUs. It requires bf16 q, fp8 k_cache inputs. - - `flashmla_auto`: enables automatic selection of either `flashmla_sparse` or `flashmla_kv` kernel for prefill based on KV cache dtype, hardware, and heuristics. When FP8 KV cache is enabled and `total_kv_tokens < total_q_tokens * 512`, it uses the `flashmla_sparse` kernel; otherwise, it falls back to the `flashmla_kv` kernel. The heuristics may need to be tuned if the performance of either the `flashmla_sparse` or `flashmla_kv` kernel changes significantly. + - `flashmla_auto`: enables automatic selection of either `flashmla_sparse` or `flashmla_kv` kernel for prefill based on KV cache dtype, hardware, and heuristics. With BF16 KV cache, `flashmla_sparse` is always used on both Hopper and Blackwell. With FP8 KV cache: On Hopper (SM90), it unconditionally uses `flashmla_kv`; On Blackwell (SM100), it uses `flashmla_sparse` when `total_kv_tokens < total_q_tokens * 512`, otherwise falls back to `flashmla_kv`. The heuristics may need to be tuned if the performance of either kernel changes significantly. - `fa3`: `flash_attn_with_kvcache` kernel from `flash_attn` library. Can only run on Hopper GPUs. It requires bf16 q, kv inputs. - `tilelang`: `tilelang` implementation that can run on GPU, HPU and NPU. - `aiter`: Aiter kernel on AMD HPUs. Can only be used as decode kernel. - `trtllm`: `trtllm-mla` sparse kernel from flashinfer library. Only run on blackwell GPUs. It requires q,k,v to be uniformly bf16 or fp8_e4m3 format. - On the basis of performance benchmarks, the default configuration of DSA kernels on Hopper and Blackwell are set as follows : - Bfloat 16 kv cache: On Hopper, `flashmla_sparse` prefill attention, `fa3` decode attention; On Blackwell, `flashmla_sparse` prefill attention, `trtllm` decode attention - - Float8_e4m3fn KV cache: On Hopper, `flashmla_auto` prefill attention, `flashmla_kv` decode attention; On Blackwell, `trtllm` prefill attention and `trtllm` decode attention. + - Float8_e4m3fn KV cache: On Hopper, `flashmla_kv` prefill attention, `flashmla_kv` decode attention; On Blackwell, `trtllm` prefill attention and `trtllm` decode attention. - **Index Cache**: Introduce in [this paper](https://arxiv.org/abs/2603.12201), IndexCache improves speed by reusing the result of indexer across different layers, only at cost of negligible accuracy loss. For **GLM-5** model, we recommend appending `--json-model-override-args '{"index_topk_pattern": "FFSFSSSFSSFFFSSSFFFSFSSSSSSFFSFFSFFSSFFFFFFSFFFFFSFFSSSSSSFSFFFSFSSSFSFFSFFSSS"}'` to command for better tradeoff between speedup and performance. ## Multi-token Prediction diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index 6cbb8cee0..397932f38 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -326,6 +326,13 @@ class NativeSparseAttnBackend( model_runner.server_args.nsa_prefill_backend ) self.nsa_decode_impl: _NSA_IMPL_T = model_runner.server_args.nsa_decode_backend + if self.num_q_heads <= 64: + self.flashmla_kv_num_q_heads = 64 + elif self.num_q_heads <= 128: + self.flashmla_kv_num_q_heads = 128 + else: + # Keep original head count if it exceeds current padded variants. + self.flashmla_kv_num_q_heads = self.num_q_heads self.enable_auto_select_prefill_impl = self.nsa_prefill_impl == "flashmla_auto" self._arange_buf = torch.arange(16384, device=self.device, dtype=torch.int32) @@ -416,6 +423,16 @@ class NativeSparseAttnBackend( # Centralized dispatch: decide all strategies for this batch self.set_nsa_prefill_impl(forward_batch) + nsa_impl_for_batch = ( + self.nsa_decode_impl + if ( + forward_batch.forward_mode.is_decode_or_idle() + or forward_batch.forward_mode.is_target_verify() + or forward_batch.forward_mode.is_draft_extend(include_v2=True) + ) + else self.nsa_prefill_impl + ) + use_flashmla_kv = (not self.use_mha) and nsa_impl_for_batch == "flashmla_kv" topk_transform_method = self.get_topk_transform_method( forward_batch.forward_mode ) @@ -651,7 +668,7 @@ class NativeSparseAttnBackend( cache_seqlens=nsa_cache_seqlens_int32, seq_len_q=1, ) - if self.nsa_decode_impl == "flashmla_kv" + if use_flashmla_kv else None ), paged_mqa_schedule_metadata=paged_mqa_schedule_metadata, @@ -1719,9 +1736,21 @@ class NativeSparseAttnBackend( from sgl_kernel.flash_mla import flash_mla_with_kvcache cache_seqlens = metadata.nsa_cache_seqlens_int32 + assert metadata.flashmla_metadata is not None # TODO the 2nd dim is seq_len_q, need to be >1 when MTP q_all = q_all.view(-1, 1, layer.tp_q_head_num, layer.head_dim) + num_q_heads = q_all.shape[2] + target_q_heads = self.flashmla_kv_num_q_heads + if target_q_heads != num_q_heads: + # Pad q heads to match FlashMLA decode supported head-count variants. + q_input = q_all.new_zeros( + q_all.shape[0], q_all.shape[1], target_q_heads, q_all.shape[3] + ) + q_input[:, :, :num_q_heads, :] = q_all + else: + q_input = q_all + kv_cache = kv_cache.view(-1, self.real_page_size, 1, self.kv_cache_dim) assert self.real_page_size == 64, "only page size 64 is supported" @@ -1735,7 +1764,7 @@ class NativeSparseAttnBackend( ) # requirement of FlashMLA decode kernel o, _ = flash_mla_with_kvcache( - q=q_all, + q=q_input, k_cache=kv_cache, cache_seqlens=cache_seqlens, head_dim_v=v_head_dim, @@ -1749,6 +1778,10 @@ class NativeSparseAttnBackend( ), is_fp8_kvcache=True, ) + + if target_q_heads != num_q_heads: + o = o[:, :, :num_q_heads, :] + return o def _forward_standard_mha( @@ -2198,13 +2231,15 @@ class NativeSparseAttnBackend( def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int): from sgl_kernel.flash_mla import get_mla_metadata + num_heads_q = self.flashmla_kv_num_q_heads + flashmla_metadata, num_splits = get_mla_metadata( cache_seqlens=cache_seqlens, # TODO doc says `num_q_tokens_per_q_seq * num_heads_q // num_heads_k` # but the name looks like need seq_len_q? - num_q_tokens_per_head_k=seq_len_q * self.num_q_heads // 1, + num_q_tokens_per_head_k=seq_len_q * num_heads_q // 1, num_heads_k=1, - num_heads_q=self.num_q_heads, + num_heads_q=num_heads_q, is_fp8_kvcache=True, topk=self.nsa_index_topk, ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 5f92dc531..cb549bfcd 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1534,9 +1534,9 @@ class ServerArgs: if not user_set_decode: self.nsa_decode_backend = "trtllm" else: - # flashmla_auto dispatches to flashmla_sparse/flashmla_kv based on hardware and heuristics + # Hopper FP8 defaults to flashmla_kv for both prefill and decode. if not user_set_prefill: - self.nsa_prefill_backend = "flashmla_auto" + self.nsa_prefill_backend = "flashmla_kv" if not user_set_decode: self.nsa_decode_backend = "flashmla_kv" else: