[DSA] Hopper FP8 FlashMLA KV padding (#22372)
This commit is contained in:
@@ -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:
|
- **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_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_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.
|
- `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.
|
- `tilelang`: `tilelang` implementation that can run on GPU, HPU and NPU.
|
||||||
- `aiter`: Aiter kernel on AMD HPUs. Can only be used as decode kernel.
|
- `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.
|
- `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 :
|
- 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
|
- 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.
|
- **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
|
## Multi-token Prediction
|
||||||
|
|||||||
@@ -326,6 +326,13 @@ class NativeSparseAttnBackend(
|
|||||||
model_runner.server_args.nsa_prefill_backend
|
model_runner.server_args.nsa_prefill_backend
|
||||||
)
|
)
|
||||||
self.nsa_decode_impl: _NSA_IMPL_T = model_runner.server_args.nsa_decode_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.enable_auto_select_prefill_impl = self.nsa_prefill_impl == "flashmla_auto"
|
||||||
|
|
||||||
self._arange_buf = torch.arange(16384, device=self.device, dtype=torch.int32)
|
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
|
# Centralized dispatch: decide all strategies for this batch
|
||||||
self.set_nsa_prefill_impl(forward_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(
|
topk_transform_method = self.get_topk_transform_method(
|
||||||
forward_batch.forward_mode
|
forward_batch.forward_mode
|
||||||
)
|
)
|
||||||
@@ -651,7 +668,7 @@ class NativeSparseAttnBackend(
|
|||||||
cache_seqlens=nsa_cache_seqlens_int32,
|
cache_seqlens=nsa_cache_seqlens_int32,
|
||||||
seq_len_q=1,
|
seq_len_q=1,
|
||||||
)
|
)
|
||||||
if self.nsa_decode_impl == "flashmla_kv"
|
if use_flashmla_kv
|
||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
||||||
@@ -1719,9 +1736,21 @@ class NativeSparseAttnBackend(
|
|||||||
from sgl_kernel.flash_mla import flash_mla_with_kvcache
|
from sgl_kernel.flash_mla import flash_mla_with_kvcache
|
||||||
|
|
||||||
cache_seqlens = metadata.nsa_cache_seqlens_int32
|
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
|
# 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)
|
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)
|
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"
|
assert self.real_page_size == 64, "only page size 64 is supported"
|
||||||
|
|
||||||
@@ -1735,7 +1764,7 @@ class NativeSparseAttnBackend(
|
|||||||
) # requirement of FlashMLA decode kernel
|
) # requirement of FlashMLA decode kernel
|
||||||
|
|
||||||
o, _ = flash_mla_with_kvcache(
|
o, _ = flash_mla_with_kvcache(
|
||||||
q=q_all,
|
q=q_input,
|
||||||
k_cache=kv_cache,
|
k_cache=kv_cache,
|
||||||
cache_seqlens=cache_seqlens,
|
cache_seqlens=cache_seqlens,
|
||||||
head_dim_v=v_head_dim,
|
head_dim_v=v_head_dim,
|
||||||
@@ -1749,6 +1778,10 @@ class NativeSparseAttnBackend(
|
|||||||
),
|
),
|
||||||
is_fp8_kvcache=True,
|
is_fp8_kvcache=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if target_q_heads != num_q_heads:
|
||||||
|
o = o[:, :, :num_q_heads, :]
|
||||||
|
|
||||||
return o
|
return o
|
||||||
|
|
||||||
def _forward_standard_mha(
|
def _forward_standard_mha(
|
||||||
@@ -2198,13 +2231,15 @@ class NativeSparseAttnBackend(
|
|||||||
def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int):
|
def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int):
|
||||||
from sgl_kernel.flash_mla import get_mla_metadata
|
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(
|
flashmla_metadata, num_splits = get_mla_metadata(
|
||||||
cache_seqlens=cache_seqlens,
|
cache_seqlens=cache_seqlens,
|
||||||
# TODO doc says `num_q_tokens_per_q_seq * num_heads_q // num_heads_k`
|
# TODO doc says `num_q_tokens_per_q_seq * num_heads_q // num_heads_k`
|
||||||
# but the name looks like need seq_len_q?
|
# 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_k=1,
|
||||||
num_heads_q=self.num_q_heads,
|
num_heads_q=num_heads_q,
|
||||||
is_fp8_kvcache=True,
|
is_fp8_kvcache=True,
|
||||||
topk=self.nsa_index_topk,
|
topk=self.nsa_index_topk,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1534,9 +1534,9 @@ class ServerArgs:
|
|||||||
if not user_set_decode:
|
if not user_set_decode:
|
||||||
self.nsa_decode_backend = "trtllm"
|
self.nsa_decode_backend = "trtllm"
|
||||||
else:
|
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:
|
if not user_set_prefill:
|
||||||
self.nsa_prefill_backend = "flashmla_auto"
|
self.nsa_prefill_backend = "flashmla_kv"
|
||||||
if not user_set_decode:
|
if not user_set_decode:
|
||||||
self.nsa_decode_backend = "flashmla_kv"
|
self.nsa_decode_backend = "flashmla_kv"
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user