diff --git a/docs/design/attention_backends.md b/docs/design/attention_backends.md index d4a406d080f..34d0a45575b 100644 --- a/docs/design/attention_backends.md +++ b/docs/design/attention_backends.md @@ -170,8 +170,8 @@ Priority is **1 = highest** (tried first). | Backend | Version | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | MM Prefix | DCP | Attention Types | Compute Cap. | | ------- | ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | --------- | --- | --------------- | ------------ | | `CPU_ATTN` | | fp16, bf16, fp32 | `auto`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 112, 128, 160, 192, 224, 256, 512 | ❌ | ❌ | ❌ | ❌ | All | N/A | -| `FLASHINFER` | Native† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64 | 64, 128, 256, 512 | ❌ | ❌ | ❌ | ✅ | Decoder | 7.x-9.x | -| `FLASHINFER` | TRTLLM† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `nvfp4` | 16, 32, 64 | 64, 128, 256, 512 | ✅ | ❌ | ❌ | ✅ | Decoder | 10.x | +| `FLASHINFER` | Native† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ❌ | ❌ | ❌ | ✅ | Decoder | 7.x-9.x | +| `FLASHINFER` | TRTLLM† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `nvfp4` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ✅ | ❌ | ❌ | ✅ | Decoder | 10.x | | `FLASH_ATTN` | FA2* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ❌ | ✅ | All | ≥8.0 | | `FLASH_ATTN` | FA3* | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | 9.x | | `FLASH_ATTN` | FA4* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | ≥10.0 | diff --git a/vllm/v1/attention/backends/flashinfer.py b/vllm/v1/attention/backends/flashinfer.py index a81c5742c1b..93d0606a795 100755 --- a/vllm/v1/attention/backends/flashinfer.py +++ b/vllm/v1/attention/backends/flashinfer.py @@ -337,9 +337,11 @@ class FlashInferBackend(AttentionBackend): @staticmethod def get_supported_kernel_block_sizes() -> list[int | MultipleOf]: - # Note: Not sure for all platforms, but on Blackwell, - # only support a page size of 16, 32, 64. - return [16, 32, 64] + # 16/32/64 are served by trtllm-gen static cubins or the FI native + # wrappers (any head config). Power-of-2 sizes >= 128 are served only by + # the trtllm-gen dynamic kernel, which requires GQA/MQA on Blackwell + # (not MHA); enforced in FlashInferMetadataBuilder.__init__. + return [16, 32, 64, 128, 256, 512, 1024] @staticmethod def get_name() -> str: @@ -642,6 +644,31 @@ class FlashInferMetadataBuilder(AttentionMetadataBuilder[FlashInferMetadata]): # if TRTLLM attention kernel is not used when building attn metadata can_use_trtllm = can_use_trtllm_attention(self.num_qo_heads, self.num_kv_heads) + # Page sizes >= 128 are served only by the trtllm-gen dynamic kernel, + # which requires Blackwell + GQA/MQA (num_qo_heads // num_kv_heads > 1), + # not MHA. We do not fall back to the FI native (FA2) kernels for large + # pages. Fail fast here rather than mid-serving. (Sizes returned by + # get_supported_kernel_block_sizes() are all power-of-2.) + if self.page_size >= 128: + if self.attention_config.use_trtllm_attention is False: + raise ValueError( + f"FlashInfer page size {self.page_size} requires the " + "trtllm-gen backend, but " + "--attention-config.use_trtllm_attention is set to 0." + ) + if not can_use_trtllm: + raise NotImplementedError( + f"FlashInfer page size {self.page_size} requires the " + "trtllm-gen backend (Blackwell with NVIDIA artifactory " + "access and num_qo_heads % num_kv_heads == 0)." + ) + if self.num_qo_heads // self.num_kv_heads <= 1: + raise NotImplementedError( + f"FlashInfer page size {self.page_size} is only supported " + "by the trtllm-gen dynamic kernel, which requires GQA/MQA " + "(num_qo_heads // num_kv_heads > 1), not MHA." + ) + if ( can_use_trtllm and not vllm_config.attention_config.disable_flashinfer_q_quantization @@ -912,6 +939,12 @@ class FlashInferMetadataBuilder(AttentionMetadataBuilder[FlashInferMetadata]): # - Decode (FI native or TRTLLM) use_cascade = common_prefix_len > 0 uses_spec_reorder = self.reorder_batch_threshold > 1 + # Page sizes >= 128 require the trtllm-gen path (the init guard verified + # GQA/MQA on Blackwell); force trtllm for prefill too so it does not + # fall back to the native wrapper. <= 64 keeps auto-detection. + prefill_force_trtllm = ( + True if page_size >= 128 else self.attention_config.use_trtllm_attention + ) prefill_use_trtllm = use_trtllm_attention( self.num_qo_heads, self.num_kv_heads, @@ -921,7 +954,7 @@ class FlashInferMetadataBuilder(AttentionMetadataBuilder[FlashInferMetadata]): self.cache_dtype, self.q_data_type, is_prefill=True, - force_use_trtllm=self.attention_config.use_trtllm_attention, + force_use_trtllm=prefill_force_trtllm, has_sinks=self.has_sinks, has_spec=uses_spec_reorder, )