Add fmha_v2 attention backend for SM90/120 (#23112)

Co-authored-by: Yangmin Li <yangminl@nvidia.com>
This commit is contained in:
akhilg-nv
2026-08-18 18:41:03 -07:00
committed by GitHub
co-authored by Yangmin Li
parent baa2251847
commit 5d12280ae7
2 changed files with 92 additions and 21 deletions
@@ -2,7 +2,14 @@ from __future__ import annotations
""" """
Support attention backend for TRTLLM MHA kernels from flashinfer. Support attention backend for TRTLLM MHA kernels from flashinfer.
The kernel supports sm100 only, with sliding window and attention sink features.
Prefill dispatch:
- SM90 / SM120: trtllm_fmha_v2_prefill (Q_PAGED_KV_NHD layout)
- SM100: trtllm_batch_context_with_kv_cache (HND layout)
Decode: uses XQA on SM90 and SM120, TRTLLM-GEN on SM100.
Sliding window and attention sink features are supported.
""" """
import logging import logging
@@ -217,6 +224,9 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
# KV fp8: q_type = fp8, out_type=model_runner.dtype # KV fp8: q_type = fp8, out_type=model_runner.dtype
self.is_xqa_impl = is_sm90_supported() or is_sm120_supported() self.is_xqa_impl = is_sm90_supported() or is_sm120_supported()
# fmha_v2 prefill kernel supports SM90 and SM120
self.use_fmha_v2 = is_sm90_supported() or is_sm120_supported()
# trtllm-gen serves page_size >= 128 only through its dynamic # trtllm-gen serves page_size >= 128 only through its dynamic
# tokens-per-page kernels, which exist solely for GQA with equal QK/V # tokens-per-page kernels, which exist solely for GQA with equal QK/V
# head dims (power-of-2 pages). Mirror that precondition here so an # head dims (power-of-2 pages). Mirror that precondition here so an
@@ -1247,33 +1257,41 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
and not use_fused_qkv and not use_fused_qkv
): ):
q = q.to(torch.float8_e4m3fn) q = q.to(torch.float8_e4m3fn)
q = q.reshape(-1, layer.tp_q_head_num, layer.head_dim)
# [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, page_size, head_dim]
k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
k_cache = k_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
).permute(0, 2, 1, 3)
v_cache = v_cache.view(
-1, self.page_size, layer.tp_v_head_num, layer.head_dim
).permute(0, 2, 1, 3)
if layer.tp_k_head_num == 1: if self.use_fmha_v2:
k_cache = canonicalize_stride(k_cache) q = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim)
if layer.tp_v_head_num == 1: else:
v_cache = canonicalize_stride(v_cache) q = q.reshape(-1, layer.tp_q_head_num, layer.head_dim)
# NHD layout (native pool format): [num_pages, page_size, num_kv_heads, head_dim]
k_cache_raw, v_cache_raw = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
is_decode_mode = (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
)
if not self.use_fmha_v2 or is_decode_mode:
# Decode and SM100 batch_context kernels require HND layout.
k_cache, v_cache = self._reshape_paged_kv_cache(
k_cache_raw, v_cache_raw, layer, layer.head_dim
)
else:
k_cache = k_cache_raw.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
)
v_cache = v_cache_raw.view(
-1, self.page_size, layer.tp_v_head_num, layer.head_dim
)
kv_cache = (k_cache, v_cache) kv_cache = (k_cache, v_cache)
# sink: additional value per head in the denominator of the softmax. # sink: additional value per head in the denominator of the softmax.
attention_sink = kwargs.get("sinks", None) attention_sink = kwargs.get("sinks", None)
bmm1_scale, bmm2_scale = self._get_bmm_scales(layer, q_scale) bmm1_scale, bmm2_scale = self._get_bmm_scales(layer, q_scale)
page_table = self._get_layer_page_table(layer, forward_batch) page_table = self._get_layer_page_table(layer, forward_batch)
if ( if is_decode_mode:
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
):
if ( if (
forward_batch.forward_mode.is_target_verify() forward_batch.forward_mode.is_target_verify()
and layer.attn_type == AttentionType.ENCODER_ONLY and layer.attn_type == AttentionType.ENCODER_ONLY
@@ -1342,6 +1360,30 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
q_len_per_req=self.forward_metadata.max_seq_len_q, q_len_per_req=self.forward_metadata.max_seq_len_q,
multi_ctas_kv_counter_buffer=self._multi_ctas_kv_counter_buffer, multi_ctas_kv_counter_buffer=self._multi_ctas_kv_counter_buffer,
) )
elif self.use_fmha_v2 and not cp_v2_active:
# CP-v2 must go through cp_strategy.run_attention (per-shard
# masking); the plain-causal fmha_v2 call below would be wrong.
paged_kv = torch.stack([k_cache, v_cache], dim=1)
o = flashinfer.prefill.trtllm_fmha_v2_prefill(
(q, paged_kv),
input_layout="Q_PAGED_KV_NHD",
workspace_buffer=self.workspace_buffer,
seq_lens=self.forward_metadata.cache_seqlens_int32,
max_q_len=self.forward_metadata.max_seq_len_q,
max_kv_len=self.max_context_len,
bmm1_scale=bmm1_scale,
bmm2_scale=bmm2_scale,
batch_size=forward_batch.batch_size,
cum_seq_lens_q=self.forward_metadata.cu_seqlens_q,
cum_seq_lens_kv=self.forward_metadata.cu_seqlens_k,
block_tables=page_table,
out_dtype=self.q_data_type,
mask_mode="causal",
window_left=layer.sliding_window_size,
sinks=attention_sink,
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get()
or 0.0,
)
else: else:
def _trtllm_context_attn( def _trtllm_context_attn(
+31 -2
View File
@@ -6013,9 +6013,28 @@ class ServerArgs:
prefill_backend, decode_backend = self._resolved_attention_backends() prefill_backend, decode_backend = self._resolved_attention_backends()
if "trtllm_mha" in (prefill_backend, decode_backend): if "trtllm_mha" in (prefill_backend, decode_backend):
if prefill_backend == "trtllm_mha" and not is_sm100_supported(): if prefill_backend == "trtllm_mha" and not (
is_sm90_supported() or is_sm100_supported() or is_sm120_supported()
):
raise ValueError( raise ValueError(
"TRTLLM MHA backend for prefill is only supported on Blackwell GPUs (SM100). Please use a different prefill backend." "TRTLLM MHA backend for prefill requires Hopper (SM90), Blackwell (SM100), or SM120 GPUs. "
"Please use a different prefill backend."
)
if (
prefill_backend == "trtllm_mha"
and is_sm120_supported()
and (
self.kv_cache_dtype == "fp8_e4m3"
or (
envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get()
or 0.0
)
> 0
)
):
raise ValueError(
"TRTLLM FMHAv2 prefill on SM120 does not support "
"fp8_e4m3 KV cache or skip-softmax."
) )
if decode_backend == "trtllm_mha" and not ( if decode_backend == "trtllm_mha" and not (
is_sm90_supported() or is_sm100_supported() or is_sm120_supported() is_sm90_supported() or is_sm100_supported() or is_sm120_supported()
@@ -6023,6 +6042,16 @@ class ServerArgs:
raise ValueError( raise ValueError(
"TRTLLM MHA backend for decode is only supported on Hopper (SM90), Blackwell (SM100) and (SM120) GPUs. Please use a different decode backend." "TRTLLM MHA backend for decode is only supported on Hopper (SM90), Blackwell (SM100) and (SM120) GPUs. Please use a different decode backend."
) )
if (
prefill_backend == "trtllm_mha"
and not is_sm100_supported()
and (self.enable_prefill_context_parallel or self.attn_cp_size > 1)
):
raise ValueError(
"Prefill context parallelism with the TRTLLM MHA prefill backend "
"requires SM100 (trtllm-gen context kernel): the SM90/SM120 "
"fmha_v2 prefill path does not implement CP shard masking."
)
run_post_process_pass(self, _attention_backend_fa3_fp8_fallback) run_post_process_pass(self, _attention_backend_fa3_fp8_fallback)