From 2911b782883d61c75aa7eb3570e8b9fdea6ffab6 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 10 Dec 2025 15:26:16 -0500 Subject: [PATCH] feat: Separate Flash Attention 2/3 and SageAttention 2/3 backends MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Rename attention modes: flash_attn→flash_attn_2/3, sa2/sa3→sageattn_2/3 - Add separate detection and wrappers for FA2, FA3, SA2, SA3 in compatibility.py - FA3: Filter unsupported params (dropout_p, window_size), return tuple[0] - SA2/SA3: Add half-precision dtype handling (convert fp32/fp8→bf16) - SA3: Add varlen-to-batched conversion with SA2 fallback for non-uniform seqs - Add fallback chains: FA3→FA2→SDPA, SA3→SA2→SDPA - Update debug.py to show granular availability: FlashAttn / SageAttn - Update all references: README, CLI, ComfyUI nodes, docstrings --- README.md | 9 +- inference_cli.py | 4 +- src/core/generation_utils.py | 2 +- src/core/model_configuration.py | 15 +- src/interfaces/dit_model_loader.py | 11 +- src/models/dit_3b/attention.py | 39 ++- src/models/dit_7b/attention.py | 39 ++- src/optimization/compatibility.py | 406 ++++++++++++++++++++++++----- src/utils/debug.py | 28 +- 9 files changed, 450 insertions(+), 103 deletions(-) diff --git a/README.md b/README.md index ec659a6..c6045e9 100644 --- a/README.md +++ b/README.md @@ -424,8 +424,11 @@ Configure the DiT (Diffusion Transformer) model for video upscaling. - Requires offload_device to be set and different from device - **attention_mode**: Attention computation backend - - `sdpa`: PyTorch scaled_dot_product_attention (default, stable, always available) - - `flash_attn`: Flash Attention 2 (faster on supported hardware, requires flash-attn package) + - `sdpa`: PyTorch scaled_dot_product_attention (default, always available) + - `flash_attn_2`: Flash Attention 2 (Ampere+, requires flash-attn package) + - `flash_attn_3`: Flash Attention 3 (Hopper+, requires flash-attn with FA3 support) + - `sageattn_2`: SageAttention 2 (requires sageattention package) + - `sageattn_3`: SageAttention 3 (Blackwell/RTX 50xx, requires sageattn3 package) - **torch_compile_args**: Connect to SeedVR2 Torch Compile Settings node for 20-40% speedup @@ -882,7 +885,7 @@ python inference_cli.py media_folder/ \ **Performance Optimization:** - `--allow_vram_overflow`: Allow VRAM overflow to system RAM. Prevents OOM but may cause severe slowdown -- `--attention_mode`: Attention backend: 'sdpa' (default, stable) or 'flash_attn' (faster, requires package) +- `--attention_mode`: Attention backend: 'sdpa' (default), 'flash_attn_2' (Ampere+), 'flash_attn_3' (Hopper+), 'sageattn_2', or 'sageattn_3' (Blackwell) - `--compile_dit`: Enable torch.compile for DiT model (20-40% speedup, requires PyTorch 2.0+ and Triton) - `--compile_vae`: Enable torch.compile for VAE model (15-25% speedup, requires PyTorch 2.0+ and Triton) - `--compile_backend`: Compilation backend: 'inductor' (full optimization) or 'cudagraphs' (lightweight) (default: inductor) diff --git a/inference_cli.py b/inference_cli.py index f19b52d..eae5699 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -1351,8 +1351,8 @@ Examples: # Performance perf_group = parser.add_argument_group('Performance optimization') perf_group.add_argument("--attention_mode", type=str, default="sdpa", - choices=["sdpa", "flash_attn", "sa2", "sa3"], - help="Attention backend: 'sdpa' (default), 'flash_attn', 'sa2', or 'sa3'") + choices=["sdpa", "flash_attn_2", "flash_attn_3", "sageattn_2", "sageattn_3"], + help="Attention backend: 'sdpa' (default), 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3' (Blackwell GPUs)") perf_group.add_argument("--compile_dit", action="store_true", help="Enable torch.compile for DiT model (20-40%% speedup, requires PyTorch 2.0+ and Triton)") perf_group.add_argument("--compile_vae", action="store_true", diff --git a/src/core/generation_utils.py b/src/core/generation_utils.py index 6b21d96..9869c48 100644 --- a/src/core/generation_utils.py +++ b/src/core/generation_utils.py @@ -457,7 +457,7 @@ def prepare_runner( decode_tile_size: Tile size for decoding (height, width) decode_tile_overlap: Tile overlap for decoding (height, width) tile_debug: Tile visualization mode (false/encode/decode) - attention_mode: Attention computation backend ('sdpa' or 'flash_attn') + attention_mode: Attention computation backend ('sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3') torch_compile_args_dit: Optional torch.compile configuration for DiT model torch_compile_args_vae: Optional torch.compile configuration for VAE model diff --git a/src/core/model_configuration.py b/src/core/model_configuration.py index 4190861..684f9d2 100644 --- a/src/core/model_configuration.py +++ b/src/core/model_configuration.py @@ -171,7 +171,7 @@ def _describe_attention_mode(attention_mode: Optional[str]) -> str: Generate human-readable description of attention mode configuration. Args: - attention_mode: Attention mode string ('sdpa' or 'flash_attn' or 'sa2' or 'sa3') + attention_mode: Attention mode string ('sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3') Returns: Human-readable description string @@ -181,9 +181,10 @@ def _describe_attention_mode(attention_mode: Optional[str]) -> str: mode_descriptions = { 'sdpa': 'PyTorch SDPA', - 'flash_attn': 'Flash Attention 2', - 'sa2': 'SageAttention v2', - 'sa3': 'SageAttention v3' + 'flash_attn_2': 'Flash Attention 2', + 'flash_attn_3': 'Flash Attention 3', + 'sageattn_2': 'SageAttention 2', + 'sageattn_3': 'SageAttention 3 (Blackwell)' } return mode_descriptions.get(attention_mode, attention_mode) @@ -438,7 +439,7 @@ def _update_dit_config( - dynamic: bool - Enable dynamic shapes - dynamo_cache_size_limit: int - Cache size limit - dynamo_recompile_limit: int - Recompilation limit - attention_mode: Attention computation backend ('sdpa' or 'flash_attn') + attention_mode: Attention computation backend ('sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3') debug: Debug instance for logging Returns: @@ -773,7 +774,7 @@ def configure_runner( decode_tile_size: Tile size for decoding (height, width) decode_tile_overlap: Tile overlap for decoding (height, width) tile_debug: Tile visualization mode (false/encode/decode) - attention_mode: Attention computation backend ('sdpa' or 'flash_attn') + attention_mode: Attention computation backend ('sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3') torch_compile_args_dit: Optional torch.compile configuration for DiT model torch_compile_args_vae: Optional torch.compile configuration for VAE model @@ -859,7 +860,7 @@ def _configure_runner_settings( decode_tile_size: Tile dimensions (height, width) for decoding in pixels decode_tile_overlap: Overlap dimensions (height, width) between decoding tiles tile_debug: Tile visualization mode (false/encode/decode) - attention_mode: Attention computation backend ('sdpa' or 'flash_attn') + attention_mode: Attention computation backend ('sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3') torch_compile_args_dit: torch.compile configuration for DiT model or None torch_compile_args_vae: torch.compile configuration for VAE model or None block_swap_config: BlockSwap configuration for DiT model or None diff --git a/src/interfaces/dit_model_loader.py b/src/interfaces/dit_model_loader.py index 526754d..e9b3548 100644 --- a/src/interfaces/dit_model_loader.py +++ b/src/interfaces/dit_model_loader.py @@ -100,15 +100,16 @@ class SeedVR2LoadDiTModel(io.ComfyNode): ) ), io.Combo.Input("attention_mode", - options=["sdpa", "flash_attn", "sa2", "sa3"], + options=["sdpa", "flash_attn_2", "flash_attn_3", "sageattn_2", "sageattn_3"], default="sdpa", optional=True, tooltip=( "Attention computation backend:\n" "• sdpa: PyTorch scaled_dot_product_attention (default, stable, always available)\n" - "• flash_attn: Flash Attention 2 (faster on supported hardware, requires flash-attn package)\n" - "• sa2: SageAttention v2 (requires sageattention package)\n" - "• sa3: SageAttention v3 (requires sageattention package)\n" + "• flash_attn_2: Flash Attention 2 (Ampere+, requires flash-attn package)\n" + "• flash_attn_3: Flash Attention 3 (Hopper+, requires flash-attn with FA3 support)\n" + "• sageattn_2: SageAttention 2 (requires sageattention package)\n" + "• sageattn_3: SageAttention 3 (Blackwell/RTX 50xx only, requires sageattn3 package)\n" "\n" "SDPA is recommended - stable and works everywhere.\n" "Flash Attention and SageAttention provide speedup through optimized CUDA kernels on compatible GPUs." @@ -144,7 +145,7 @@ class SeedVR2LoadDiTModel(io.ComfyNode): cache_model: Whether to keep model loaded between runs blocks_to_swap: Number of transformer blocks to swap (requires offload_device != device) swap_io_components: Whether to offload I/O components (requires offload_device != device) - attention_mode: Attention computation backend ('sdpa' or 'flash_attn') + attention_mode: Attention computation backend ('sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3') torch_compile_args: Optional torch.compile configuration from settings node Returns: diff --git a/src/models/dit_3b/attention.py b/src/models/dit_3b/attention.py index d6ace0d..03e3c63 100644 --- a/src/models/dit_3b/attention.py +++ b/src/models/dit_3b/attention.py @@ -16,7 +16,10 @@ import torch import torch.nn.functional as F # Import flash/sage attn with automatic fallback from compatibility layer -from ...optimization.compatibility import call_flash_attn_varlen, call_sage_attn_varlen +from ...optimization.compatibility import ( + call_flash_attn_2_varlen, call_flash_attn_3_varlen, + call_sage_attn_2_varlen, call_sage_attn_3_varlen +) from torch import nn @@ -76,12 +79,16 @@ class TorchAttention(nn.Module): class FlashAttentionVarlen(nn.Module): """ - Variable-length attention with configurable backend (Flash Attention or PyTorch SDPA). + Variable-length attention with configurable backend. - Backend selection is validated during model configuration. - Compilation behavior: - - SDPA: Fully compilable, optimal performance - - Flash Attention: Uses @torch._dynamo.disable wrapper (C++ extension) + Supported backends: + - sdpa: PyTorch SDPA (fully compilable, always available) + - flash_attn_2: Flash Attention 2 (Ampere+) + - flash_attn_3: Flash Attention 3 (Hopper+) + - sageattn_2: SageAttention 2 + - sageattn_3: SageAttention 3 (Blackwell/RTX 50xx) + + All non-SDPA backends use @torch._dynamo.disable wrapper (C++ extensions). """ def __init__(self, attention_mode: str = 'sdpa', compute_dtype: torch.dtype = None): @@ -89,7 +96,7 @@ class FlashAttentionVarlen(nn.Module): Initialize with specified attention backend. Args: - attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_attention_mode) + attention_mode: 'sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3' compute_dtype: Compute dtype for attention (set by pipeline, defaults to None for auto-detection) """ super().__init__() @@ -113,13 +120,23 @@ class FlashAttentionVarlen(nn.Module): k = k.to(self.compute_dtype) v = v.to(self.compute_dtype) - if self.attention_mode == 'flash_attn': - return call_flash_attn_varlen( + if self.attention_mode == 'flash_attn_3': + return call_flash_attn_3_varlen( q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs ) - elif self.attention_mode in ('sa2', 'sa3'): - return call_sage_attn_varlen( + elif self.attention_mode == 'flash_attn_2': + return call_flash_attn_2_varlen( + q, k, v, cu_seqlens_q, cu_seqlens_k, + max_seqlen_q, max_seqlen_k, **kwargs + ) + elif self.attention_mode == 'sageattn_3': + return call_sage_attn_3_varlen( + q, k, v, cu_seqlens_q, cu_seqlens_k, + max_seqlen_q, max_seqlen_k, **kwargs + ) + elif self.attention_mode == 'sageattn_2': + return call_sage_attn_2_varlen( q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs ) diff --git a/src/models/dit_7b/attention.py b/src/models/dit_7b/attention.py index d6ace0d..03e3c63 100644 --- a/src/models/dit_7b/attention.py +++ b/src/models/dit_7b/attention.py @@ -16,7 +16,10 @@ import torch import torch.nn.functional as F # Import flash/sage attn with automatic fallback from compatibility layer -from ...optimization.compatibility import call_flash_attn_varlen, call_sage_attn_varlen +from ...optimization.compatibility import ( + call_flash_attn_2_varlen, call_flash_attn_3_varlen, + call_sage_attn_2_varlen, call_sage_attn_3_varlen +) from torch import nn @@ -76,12 +79,16 @@ class TorchAttention(nn.Module): class FlashAttentionVarlen(nn.Module): """ - Variable-length attention with configurable backend (Flash Attention or PyTorch SDPA). + Variable-length attention with configurable backend. - Backend selection is validated during model configuration. - Compilation behavior: - - SDPA: Fully compilable, optimal performance - - Flash Attention: Uses @torch._dynamo.disable wrapper (C++ extension) + Supported backends: + - sdpa: PyTorch SDPA (fully compilable, always available) + - flash_attn_2: Flash Attention 2 (Ampere+) + - flash_attn_3: Flash Attention 3 (Hopper+) + - sageattn_2: SageAttention 2 + - sageattn_3: SageAttention 3 (Blackwell/RTX 50xx) + + All non-SDPA backends use @torch._dynamo.disable wrapper (C++ extensions). """ def __init__(self, attention_mode: str = 'sdpa', compute_dtype: torch.dtype = None): @@ -89,7 +96,7 @@ class FlashAttentionVarlen(nn.Module): Initialize with specified attention backend. Args: - attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_attention_mode) + attention_mode: 'sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3' compute_dtype: Compute dtype for attention (set by pipeline, defaults to None for auto-detection) """ super().__init__() @@ -113,13 +120,23 @@ class FlashAttentionVarlen(nn.Module): k = k.to(self.compute_dtype) v = v.to(self.compute_dtype) - if self.attention_mode == 'flash_attn': - return call_flash_attn_varlen( + if self.attention_mode == 'flash_attn_3': + return call_flash_attn_3_varlen( q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs ) - elif self.attention_mode in ('sa2', 'sa3'): - return call_sage_attn_varlen( + elif self.attention_mode == 'flash_attn_2': + return call_flash_attn_2_varlen( + q, k, v, cu_seqlens_q, cu_seqlens_k, + max_seqlen_q, max_seqlen_k, **kwargs + ) + elif self.attention_mode == 'sageattn_3': + return call_sage_attn_3_varlen( + q, k, v, cu_seqlens_q, cu_seqlens_k, + max_seqlen_q, max_seqlen_k, **kwargs + ) + elif self.attention_mode == 'sageattn_2': + return call_sage_attn_2_varlen( q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs ) diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index 5b79853..bcce384 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -92,71 +92,161 @@ import os # Flash/Sage Attention & Triton Compatibility Layer -# 1. Flash Attention - speedup for attention operations -try: - from flash_attn import flash_attn_varlen_func - # Force load the CUDA extension to verify it's not corrupted - import flash_attn_2_cuda # noqa: F401 - FLASH_ATTN_AVAILABLE = True -except (ImportError, AttributeError, OSError): - flash_attn_varlen_func = None - FLASH_ATTN_AVAILABLE = False -# 2. SageAttention - speedup for attention operations +# 1. Flash Attention 3 (Hopper+, faster, no dropout/window support) +flash_attn_3_varlen_func = None +FLASH_ATTN_3_AVAILABLE = False try: - from sageattention import sageattn_varlen - SAGE_ATTN_AVAILABLE = True + import flash_attn_interface + flash_attn_3_varlen_func = flash_attn_interface.flash_attn_varlen_func + FLASH_ATTN_3_AVAILABLE = True except (ImportError, AttributeError, OSError): - sageattn_varlen = None - SAGE_ATTN_AVAILABLE = False + pass + +# 2. Flash Attention 2 (wider compatibility, supports dropout/window) +flash_attn_2_varlen_func = None +FLASH_ATTN_2_AVAILABLE = False +try: + from flash_attn import flash_attn_varlen_func as _fa2_varlen + import flash_attn_2_cuda # noqa: F401 + flash_attn_2_varlen_func = _fa2_varlen + FLASH_ATTN_2_AVAILABLE = True +except (ImportError, AttributeError, OSError): + pass + +FLASH_ATTN_AVAILABLE = FLASH_ATTN_2_AVAILABLE or FLASH_ATTN_3_AVAILABLE + +# 3. SageAttention 2 (varlen support) +sageattn_varlen = None +SAGE_ATTN_2_AVAILABLE = False +try: + from sageattention import sageattn_varlen as _sa2_varlen + sageattn_varlen = _sa2_varlen + SAGE_ATTN_2_AVAILABLE = True +except (ImportError, AttributeError, OSError): + pass + +# 4. SageAttention 3 / Blackwell (RTX 50xx only, batched attention) +sageattn_blackwell = None +SAGE_ATTN_3_AVAILABLE = False +try: + from sageattn3 import sageattn3_blackwell as _sa3_blackwell + sageattn_blackwell = _sa3_blackwell + SAGE_ATTN_3_AVAILABLE = True +except (ImportError, AttributeError, OSError): + try: + from sageattention import sageattn_blackwell as _sa3_blackwell + sageattn_blackwell = _sa3_blackwell + SAGE_ATTN_3_AVAILABLE = True + except (ImportError, AttributeError, OSError): + pass + +SAGE_ATTN_AVAILABLE = SAGE_ATTN_2_AVAILABLE or SAGE_ATTN_3_AVAILABLE def validate_attention_mode(requested_mode: str, debug=None) -> str: """ - Validate attention mode availability with automatic fallback to sdpa. + Validate attention mode availability with automatic fallback. Args: - requested_mode: 'sdpa', 'flash_attn', 'sa2', or 'sa3' + requested_mode: 'sdpa', 'flash_attn_2', 'flash_attn_3', 'sageattn_2', or 'sageattn_3' debug: Optional debug instance for logging Returns: Validated mode that is available """ - # SageAttention modes - if requested_mode in ('sa2', 'sa3'): - if SAGE_ATTN_AVAILABLE: + # Flash Attention 3 + if requested_mode == 'flash_attn_3': + if FLASH_ATTN_3_AVAILABLE: return requested_mode + if FLASH_ATTN_2_AVAILABLE: + if debug: + debug.log( + "Flash Attention 3 not available (requires Hopper+ GPU and flash-attn with FA3 support).\n" + "Falling back to Flash Attention 2.", + level="WARNING", category="setup", force=True + ) + return 'flash_attn_2' error_msg = ( - f"Cannot use '{requested_mode}' attention mode: SageAttention is not installed.\n" - f"\n" - f"SageAttention provides speedup on some hardware through optimized CUDA kernels.\n" - f"Falling back to PyTorch SDPA (scaled dot-product attention).\n" - f"\n" - f"To fix this issue:\n" - f" 1. Install SageAttention: pip install sageattention\n" - f" 2. OR change attention_mode to 'flash_attn' or 'sdpa'\n" - f"\n" - f"For more info: https://github.com/thu-ml/SageAttention" + "Cannot use 'flash_attn_3' attention mode: Flash Attention is not installed.\n" + "\n" + "Flash Attention 3 provides maximum speedup on Hopper+ GPUs through optimized CUDA kernels.\n" + "Falling back to PyTorch SDPA (scaled dot-product attention).\n" + "\n" + "To fix this issue:\n" + " 1. Install Flash Attention: pip install flash-attn\n" + " 2. OR change attention_mode to 'sdpa' (default, always available)\n" + "\n" + "For more info: https://github.com/Dao-AILab/flash-attention" ) if debug: debug.log(error_msg, level="WARNING", category="setup", force=True) return 'sdpa' - # Flash Attention - if requested_mode == 'flash_attn': - if FLASH_ATTN_AVAILABLE: + # Flash Attention 2 + if requested_mode == 'flash_attn_2': + if FLASH_ATTN_2_AVAILABLE: return requested_mode error_msg = ( - f"Cannot use 'flash_attn' attention mode: Flash Attention is not installed.\n" - f"\n" - f"Flash Attention provides speedup on some hardware through optimized CUDA kernels.\n" - f"Falling back to PyTorch SDPA (scaled dot-product attention).\n" - f"\n" - f"To fix this issue:\n" - f" 1. Install Flash Attention: pip install flash-attn\n" - f" 2. OR change attention_mode to 'sdpa' (default, always available)\n" - f"\n" - f"For more info: https://github.com/Dao-AILab/flash-attention" + "Cannot use 'flash_attn_2' attention mode: Flash Attention 2 is not installed.\n" + "\n" + "Flash Attention 2 provides speedup on Ampere+ GPUs through optimized CUDA kernels.\n" + "Falling back to PyTorch SDPA (scaled dot-product attention).\n" + "\n" + "To fix this issue:\n" + " 1. Install Flash Attention: pip install flash-attn\n" + " 2. OR change attention_mode to 'sdpa' (default, always available)\n" + "\n" + "For more info: https://github.com/Dao-AILab/flash-attention" + ) + if debug: + debug.log(error_msg, level="WARNING", category="setup", force=True) + return 'sdpa' + + # SageAttention 3 (Blackwell) + if requested_mode == 'sageattn_3': + if SAGE_ATTN_3_AVAILABLE: + return requested_mode + if SAGE_ATTN_2_AVAILABLE: + if debug: + debug.log( + "SageAttention 3 (Blackwell) not available (requires RTX 50xx GPU and sageattn3 package).\n" + "Falling back to SageAttention 2.", + level="WARNING", category="setup", force=True + ) + return 'sageattn_2' + error_msg = ( + "Cannot use 'sageattn_3' attention mode: SageAttention is not installed.\n" + "\n" + "SageAttention 3 provides maximum speedup on Blackwell (RTX 50xx) GPUs.\n" + "Falling back to PyTorch SDPA (scaled dot-product attention).\n" + "\n" + "To fix this issue:\n" + " 1. Install SageAttention: pip install sageattention\n" + " 2. For SA3 Blackwell support: pip install sageattn3\n" + " 3. OR change attention_mode to 'flash_attn_2' or 'sdpa'\n" + "\n" + "For more info: https://github.com/thu-ml/SageAttention" + ) + if debug: + debug.log(error_msg, level="WARNING", category="setup", force=True) + return 'sdpa' + + # SageAttention 2 + if requested_mode == 'sageattn_2': + if SAGE_ATTN_2_AVAILABLE: + return requested_mode + error_msg = ( + "Cannot use 'sageattn_2' attention mode: SageAttention is not installed.\n" + "\n" + "SageAttention provides speedup on NVIDIA GPUs through optimized CUDA kernels.\n" + "Falling back to PyTorch SDPA (scaled dot-product attention).\n" + "\n" + "To fix this issue:\n" + " 1. Install SageAttention: pip install sageattention\n" + " 2. OR change attention_mode to 'flash_attn_2' or 'sdpa'\n" + "\n" + "For more info: https://github.com/thu-ml/SageAttention" ) if debug: debug.log(error_msg, level="WARNING", category="setup", force=True) @@ -166,17 +256,33 @@ def validate_attention_mode(requested_mode: str, debug=None) -> str: @torch._dynamo.disable -def call_flash_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): +def call_flash_attn_2_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): """ - Wrapper for flash_attn_varlen_func that handles tensor-to-scalar conversion. + Wrapper for Flash Attention 2 flash_attn_varlen_func that handles tensor-to-scalar conversion. + + Flash Attention 2 supports dropout_p and window_size parameters. + Works on Ampere+ GPUs (RTX 30xx, 40xx, A100, etc.). This function is excluded from torch.compile because: 1. flash_attn is a C++ extension that can't be compiled anyway 2. It requires Python int scalars for max_seqlen parameters 3. Disabling compilation here keeps the rest of the model compilable + + Args: + q: Query tensor (total_seq, heads, head_dim) + k: Key tensor (total_seq, heads, head_dim) + v: Value tensor (total_seq, heads, head_dim) + cu_seqlens_q: Cumulative sequence lengths for queries + cu_seqlens_k: Cumulative sequence lengths for keys + max_seqlen_q: Maximum query sequence length (can be tensor or int) + max_seqlen_k: Maximum key sequence length (can be tensor or int) + **kwargs: Additional arguments (dropout_p, softmax_scale, causal, window_size, deterministic) + + Returns: + Attention output tensor (total_seq, heads, head_dim) """ - if not FLASH_ATTN_AVAILABLE: - raise ImportError("flash_attn is not available") + if not FLASH_ATTN_2_AVAILABLE: + raise ImportError("Flash Attention 2 is not available") # Convert tensor max_seqlen to Python int if needed if torch.is_tensor(max_seqlen_q): @@ -184,7 +290,7 @@ def call_flash_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, ma if torch.is_tensor(max_seqlen_k): max_seqlen_k = int(max_seqlen_k.item()) - return flash_attn_varlen_func( + return flash_attn_2_varlen_func( q=q, k=k, v=v, @@ -197,17 +303,34 @@ def call_flash_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, ma @torch._dynamo.disable -def call_sage_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): +def call_flash_attn_3_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): """ - Wrapper for SageAttention sageattn_varlen that handles tensor-to-scalar conversion. + Wrapper for Flash Attention 3 flash_attn_varlen_func that handles tensor-to-scalar conversion. + + Flash Attention 3 is faster than FA2 but does NOT support dropout_p and window_size. + Works on Hopper+ GPUs (H100, etc.) - requires flash_attn_interface package. This function is excluded from torch.compile because: - 1. SageAttention is a C++ extension that can't be compiled anyway + 1. flash_attn is a C++ extension that can't be compiled anyway 2. It requires Python int scalars for max_seqlen parameters 3. Disabling compilation here keeps the rest of the model compilable + + Args: + q: Query tensor (total_seq, heads, head_dim) + k: Key tensor (total_seq, heads, head_dim) + v: Value tensor (total_seq, heads, head_dim) + cu_seqlens_q: Cumulative sequence lengths for queries + cu_seqlens_k: Cumulative sequence lengths for keys + max_seqlen_q: Maximum query sequence length (can be tensor or int) + max_seqlen_k: Maximum key sequence length (can be tensor or int) + **kwargs: Additional arguments (softmax_scale, causal, deterministic) + Note: dropout_p and window_size are ignored (not supported by FA3) + + Returns: + Attention output tensor (total_seq, heads, head_dim) """ - if not SAGE_ATTN_AVAILABLE: - raise ImportError("SageAttention is not available") + if not FLASH_ATTN_3_AVAILABLE: + raise ImportError("Flash Attention 3 is not available") # Convert tensor max_seqlen to Python int if needed if torch.is_tensor(max_seqlen_q): @@ -215,16 +338,183 @@ def call_sage_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max if torch.is_tensor(max_seqlen_k): max_seqlen_k = int(max_seqlen_k.item()) - # SageAttention requires contiguous tensors - q = q.contiguous() - k = k.contiguous() - v = v.contiguous() + # FA3 doesn't support dropout_p and window_size - filter them out + fa3_kwargs = {key: val for key, val in kwargs.items() if key not in ('dropout_p', 'window_size')} + + # FA3 returns a tuple (output, softmax_lse), we only need output + return flash_attn_3_varlen_func( + q=q, + k=k, + v=v, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + max_seqlen_q=max_seqlen_q, + max_seqlen_k=max_seqlen_k, + seqused_q=None, + seqused_k=None, + **fa3_kwargs + )[0] + + +@torch._dynamo.disable +def call_sage_attn_2_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): + """ + Wrapper for SageAttention 2 sageattn_varlen that handles tensor-to-scalar conversion. + + SageAttention 2 provides optimized attention for NVIDIA GPUs with native varlen support. + Works on most modern NVIDIA GPUs. + + This function is excluded from torch.compile because: + 1. SageAttention is a C++ extension that can't be compiled anyway + 2. It requires Python int scalars for max_seqlen parameters + 3. Disabling compilation here keeps the rest of the model compilable + + Args: + q: Query tensor (total_seq, heads, head_dim) + k: Key tensor (total_seq, heads, head_dim) + v: Value tensor (total_seq, heads, head_dim) + cu_seqlens_q: Cumulative sequence lengths for queries + cu_seqlens_k: Cumulative sequence lengths for keys + max_seqlen_q: Maximum query sequence length (can be tensor or int) + max_seqlen_k: Maximum key sequence length (can be tensor or int) + **kwargs: Additional arguments (causal supported, others ignored) + + Returns: + Attention output tensor (total_seq, heads, head_dim) + """ + if not SAGE_ATTN_2_AVAILABLE: + raise ImportError("SageAttention 2 is not available") + + # Convert tensor max_seqlen to Python int if needed + if torch.is_tensor(max_seqlen_q): + max_seqlen_q = int(max_seqlen_q.item()) + if torch.is_tensor(max_seqlen_k): + max_seqlen_k = int(max_seqlen_k.item()) + + # SageAttention requires half precision (fp16/bf16) + out_dtype = q.dtype + half_dtypes = (torch.float16, torch.bfloat16) + + if not (q.dtype == k.dtype == v.dtype): + k = k.to(q.dtype) + v = v.to(q.dtype) + + if q.dtype not in half_dtypes: + q = q.to(torch.bfloat16) + k = k.to(torch.bfloat16) + v = v.to(torch.bfloat16) is_causal = kwargs.get('causal', False) sm_scale = 1.0 / (q.shape[-1] ** 0.5) - return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, - max_seqlen_q, max_seqlen_k, is_causal, sm_scale) + out = sageattn_varlen( + q, k, v, + cu_seqlens_q, cu_seqlens_k, + max_seqlen_q, max_seqlen_k, + is_causal, sm_scale + ) + + return out.to(out_dtype) if out.dtype != out_dtype else out + + +@torch._dynamo.disable +def call_sage_attn_3_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): + """ + Wrapper for SageAttention 3 (Blackwell) that converts varlen format to batched format. + + SageAttention 3 / Blackwell provides maximum performance on RTX 50xx series GPUs. + However, it only supports batched attention (uniform sequence lengths), not varlen. + + This wrapper detects uniform-length batches and reshapes accordingly. + For variable-length sequences, it automatically falls back to SageAttention 2. + + This function is excluded from torch.compile because: + 1. SageAttention is a C++ extension that can't be compiled anyway + 2. It requires Python int scalars for max_seqlen parameters + 3. The varlen-to-batched conversion involves dynamic shapes + 4. Disabling compilation here keeps the rest of the model compilable + + Args: + q: Query tensor (total_seq, heads, head_dim) + k: Key tensor (total_seq, heads, head_dim) + v: Value tensor (total_seq, heads, head_dim) + cu_seqlens_q: Cumulative sequence lengths for queries + cu_seqlens_k: Cumulative sequence lengths for keys + max_seqlen_q: Maximum query sequence length (can be tensor or int) + max_seqlen_k: Maximum key sequence length (can be tensor or int) + **kwargs: Additional arguments (passed to SA2 fallback if needed) + + Returns: + Attention output tensor (total_seq, heads, head_dim) + """ + if not SAGE_ATTN_3_AVAILABLE: + raise ImportError("SageAttention 3 (Blackwell) is not available") + + # Convert tensor max_seqlen to Python int if needed + if torch.is_tensor(max_seqlen_q): + max_seqlen_q = int(max_seqlen_q.item()) + if torch.is_tensor(max_seqlen_k): + max_seqlen_k = int(max_seqlen_k.item()) + + # Check if all sequences have uniform length (required for SA3 batched API) + # SA3/Blackwell uses batched attention, not varlen, so we need uniform lengths + seq_lens_q = cu_seqlens_q[1:] - cu_seqlens_q[:-1] + seq_lens_k = cu_seqlens_k[1:] - cu_seqlens_k[:-1] + + uniform_q = (seq_lens_q == seq_lens_q[0]).all() + uniform_k = (seq_lens_k == seq_lens_k[0]).all() + + if not (uniform_q and uniform_k): + # Fall back to SA2 for variable-length sequences + # This is expected behavior - SA3 Blackwell doesn't support varlen natively + if SAGE_ATTN_2_AVAILABLE: + return call_sage_attn_2_varlen( + q, k, v, cu_seqlens_q, cu_seqlens_k, + max_seqlen_q, max_seqlen_k, **kwargs + ) + raise RuntimeError( + "SageAttention 3 (Blackwell) requires uniform sequence lengths, " + "and SageAttention 2 is not available as fallback. " + "Please install sageattention package or use flash_attn/sdpa instead." + ) + + # Extract batch dimensions + batch_size = len(cu_seqlens_q) - 1 + seq_len_q = int(seq_lens_q[0].item()) + seq_len_k = int(seq_lens_k[0].item()) + heads = q.shape[1] + dim = q.shape[2] + + # SageAttention requires half precision (fp16/bf16) + out_dtype = q.dtype + half_dtypes = (torch.float16, torch.bfloat16) + + if not (q.dtype == k.dtype == v.dtype): + k = k.to(q.dtype) + v = v.to(q.dtype) + + if q.dtype not in half_dtypes: + q = q.to(torch.bfloat16) + k = k.to(torch.bfloat16) + v = v.to(torch.bfloat16) + + # Reshape varlen (total_seq, heads, dim) -> batched (batch, seq, heads, dim) + q_batched = q.view(batch_size, seq_len_q, heads, dim) + k_batched = k.view(batch_size, seq_len_k, heads, dim) + v_batched = v.view(batch_size, seq_len_k, heads, dim) + + # SA3/Blackwell expects (batch, heads, seq, dim) layout + q_batched = q_batched.transpose(1, 2) # (batch, heads, seq, dim) + k_batched = k_batched.transpose(1, 2) + v_batched = v_batched.transpose(1, 2) + + # Call SA3 Blackwell + out = sageattn_blackwell(q_batched, k_batched, v_batched, per_block_mean=False) + + # Reshape back to varlen format (total_seq, heads, dim) + out = out.transpose(1, 2).reshape(-1, heads, dim).contiguous() + + return out.to(out_dtype) if out.dtype != out_dtype else out # 2. Triton - Required for torch.compile with inductor backend diff --git a/src/utils/debug.py b/src/utils/debug.py index 3619084..716c3da 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -214,12 +214,30 @@ class Debug: gpu_str = "CPU" cudnn_ver = "N/A" - # Flash Attn & Triton - reuse existing module constants + # Flash Attn, SageAttn & Triton - reuse existing module constants try: - from ..optimization.compatibility import FLASH_ATTN_AVAILABLE, TRITON_AVAILABLE - flash_str, triton_str = ("✓" if FLASH_ATTN_AVAILABLE else "✗"), ("✓" if TRITON_AVAILABLE else "✗") + from ..optimization.compatibility import ( + FLASH_ATTN_2_AVAILABLE, FLASH_ATTN_3_AVAILABLE, + SAGE_ATTN_2_AVAILABLE, SAGE_ATTN_3_AVAILABLE, + TRITON_AVAILABLE + ) + fa_parts = [] + if FLASH_ATTN_3_AVAILABLE: + fa_parts.append("3") + if FLASH_ATTN_2_AVAILABLE: + fa_parts.append("2") + flash_str = f"v{','.join(fa_parts)} ✓" if fa_parts else "✗" + + sa_parts = [] + if SAGE_ATTN_3_AVAILABLE: + sa_parts.append("3") + if SAGE_ATTN_2_AVAILABLE: + sa_parts.append("2") + sage_str = f"v{','.join(sa_parts)} ✓" if sa_parts else "✗" + + triton_str = "✓" if TRITON_AVAILABLE else "✗" except ImportError: - flash_str = triton_str = "?" + flash_str = sage_str = triton_str = "?" # ComfyUI version comfy_str = None @@ -232,7 +250,7 @@ class Debug: # Print self.log(f"OS: {os_str} | GPU: {gpu_str}", category="info") - self.log(f"Python: {py_ver} | PyTorch: {torch_ver} | Flash Attn: {flash_str} | Triton: {triton_str}", category="info") + self.log(f"Python: {py_ver} | PyTorch: {torch_ver} | FlashAttn: {flash_str} | SageAttn: {sage_str} | Triton: {triton_str}", category="info") cuda_line = f"CUDA: {cuda_ver} | cuDNN: {cudnn_ver}" self.log(f"{cuda_line} | ComfyUI: {comfy_str}" if comfy_str else cuda_line, category="info") self.log("", category="none")