feat: Separate Flash Attention 2/3 and SageAttention 2/3 backends

- 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
This commit is contained in:
Adrien Toupet
2025-12-10 15:26:16 -05:00
parent bcfbca6ae3
commit 2911b78288
9 changed files with 450 additions and 103 deletions
+6 -3
View File
@@ -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)
+2 -2
View File
@@ -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",
+1 -1
View File
@@ -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
+8 -7
View File
@@ -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
+6 -5
View File
@@ -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:
+28 -11
View File
@@ -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
)
+28 -11
View File
@@ -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
)
+348 -58
View File
@@ -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
+23 -5
View File
@@ -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")