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:
@@ -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
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user