feat: add SageAttention (sa2/sa3) support, centralize attention wrappers

- Add sa2/sa3 attention modes for SageAttention v2/v3 kernels
- Centralize call_flash_attn_varlen and call_sage_attn_varlen in compatibility.py
- Remove duplicated attention wrapper code from dit_3b/dit_7b attention.py
- Rename validate_flash_attention_availability to validate_attention_mode
- Remove unnecessary precision control feature (auto/fp16/bf16/bf32)
- Remove unused detect_high_end_system() and log_system_capabilities()
- Update startup logging to show SageAttention availability status
- Update CLI and ComfyUI node to expose sa2/sa3 options
This commit is contained in:
Adrien Toupet
2025-12-10 13:58:46 -05:00
parent ab3284982f
commit bcfbca6ae3
7 changed files with 158 additions and 395 deletions
+2 -6
View File
@@ -793,8 +793,7 @@ def _process_frames_core(
dit_offload_device=dit_offload,
vae_offload_device=vae_offload,
tensor_offload_device=tensor_offload,
debug=debug,
precision=args.precision
debug=debug
)
if runner_cache is not None:
runner_cache['ctx'] = ctx
@@ -1353,10 +1352,7 @@ Examples:
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' (faster), 'sa2' (SageAttention v2), 'sa3' (SageAttention v3)")
perf_group.add_argument("--precision", type=str, default="auto",
choices=["auto", "fp16", "bf16", "bf32"],
help="Compute precision: 'auto' (default), 'fp16', 'bf16', or 'bf32' (TF32)")
help="Attention backend: 'sdpa' (default), 'flash_attn', 'sa2', or 'sa3'")
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",
+9 -30
View File
@@ -318,8 +318,7 @@ def setup_generation_context(
dit_offload_device: Optional[Union[str, torch.device]] = None,
vae_offload_device: Optional[Union[str, torch.device]] = None,
tensor_offload_device: Optional[Union[str, torch.device]] = None,
debug: Optional['Debug'] = None,
precision: str = 'auto'
debug: Optional['Debug'] = None
) -> Dict[str, Any]:
"""
Initialize generation context with device configuration.
@@ -334,7 +333,6 @@ def setup_generation_context(
vae_offload_device: Device to offload VAE to when not in use (optional)
tensor_offload_device: Device to offload intermediate tensors to (optional)
debug: Debug instance for logging
precision: Compute precision ('auto', 'fp16', 'bf16', 'bf32')
Returns:
Dict[str, Any]: Generation context dictionary with torch.device objects
@@ -367,29 +365,6 @@ def setup_generation_context(
interrupt_fn = None
comfyui_available = False
# Determine compute dtype based on precision request
if precision == 'fp16':
compute_dtype = torch.float16
reason = "user requested fp16"
elif precision == 'bf16':
compute_dtype = torch.bfloat16
reason = "user requested bf16"
elif precision == 'bf32':
# BF32 is usually implemented as float32 tensors with specific matmul settings (TF32)
# For torch dtype context, we use float32
compute_dtype = torch.float32
reason = "user requested bf32 (TF32)"
# Note: TF32 enablement should be handled globally or in model config
else:
# 'auto' - existing logic
compute_dtype = COMPUTE_DTYPE
if compute_dtype == torch.float32:
reason = "quality"
elif not BFLOAT16_SUPPORTED:
reason = "compatibility (GPU lacks bfloat16 CUBLAS - 7B models unsupported, 3B may have artifacts)"
else:
reason = "performance"
# Create generation context
ctx = {
'dit_device': dit_device,
@@ -397,7 +372,7 @@ def setup_generation_context(
'dit_offload_device': dit_offload_device,
'vae_offload_device': vae_offload_device,
'tensor_offload_device': tensor_offload_device,
'compute_dtype': compute_dtype,
'compute_dtype': COMPUTE_DTYPE,
'interrupt_fn': interrupt_fn,
'video_transform': None,
'text_embeds': None,
@@ -427,6 +402,12 @@ def setup_generation_context(
f"LOCAL_RANK={os.environ['LOCAL_RANK']}",
category="setup"
)
if ctx['compute_dtype'] == torch.float32:
reason = "quality"
elif not BFLOAT16_SUPPORTED:
reason = "compatibility (GPU lacks bfloat16 CUBLAS - 7B models unsupported, 3B may have artifacts)"
else:
reason = "performance"
debug.log(f"Unified compute dtype: {ctx['compute_dtype']} across entire pipeline for maximum {reason}", category="precision")
return ctx
@@ -451,7 +432,6 @@ def prepare_runner(
decode_tile_overlap: Optional[Tuple[int, int]] = None,
tile_debug: str = "false",
attention_mode: str = 'sdpa',
precision: str = 'auto',
torch_compile_args_dit: Optional[Dict[str, Any]] = None,
torch_compile_args_vae: Optional[Dict[str, Any]] = None
) -> Tuple['VideoDiffusionInfer', Dict[str, Any]]:
@@ -522,7 +502,6 @@ def prepare_runner(
decode_tile_overlap=decode_tile_overlap,
tile_debug=tile_debug,
attention_mode=attention_mode,
precision=precision,
torch_compile_args_dit=torch_compile_args_dit,
torch_compile_args_vae=torch_compile_args_vae
)
@@ -840,4 +819,4 @@ def ensure_precision_initialized(
debug.log(f"Model precision: {', '.join(parts)}", category="precision")
except Exception as e:
debug.log(f"Could not log model dtypes: {e}", level="WARNING", category="precision", force=True)
debug.log(f"Could not log model dtypes: {e}", level="WARNING", category="precision", force=True)
+9 -44
View File
@@ -72,9 +72,7 @@ from ..models.video_vae_v3.modules.causal_inflation_lib import InflatedCausalCon
from ..optimization.compatibility import (
FP8CompatibleDiT,
TRITON_AVAILABLE,
validate_flash_attention_availability,
detect_high_end_system,
log_system_capabilities
validate_attention_mode
)
from ..optimization.blockswap import is_blockswap_enabled, apply_block_swap_to_dit, cleanup_blockswap
from ..optimization.memory_manager import cleanup_dit, cleanup_vae
@@ -749,7 +747,6 @@ def configure_runner(
decode_tile_overlap: Optional[Tuple[int, int]] = None,
tile_debug: str = "false",
attention_mode: str = 'sdpa',
precision: str = 'auto',
torch_compile_args_dit: Optional[Dict[str, Any]] = None,
torch_compile_args_vae: Optional[Dict[str, Any]] = None
) -> Tuple[VideoDiffusionInfer, Dict[str, Any]]:
@@ -797,9 +794,6 @@ def configure_runner(
if debug is None:
raise ValueError("Debug instance must be provided to configure_runner")
# Log installed attention backends and versions
log_system_capabilities(debug)
# Phase 1: Initialize cache and get cached models
cache_context = _initialize_cache_context(
dit_cache, vae_cache, dit_id, vae_id,
@@ -822,9 +816,6 @@ def configure_runner(
block_swap_config, debug
)
# Store precision setting
runner._precision = precision
# Phase 4: Setup models (load from cache or create new)
_setup_models(
runner, cache_context, dit_model, vae_model,
@@ -906,16 +897,6 @@ def _configure_runner_settings(
runner._tensor_offload_device = ctx['tensor_offload_device']
runner._compute_dtype = ctx['compute_dtype']
# Auto-detection for 5070ti/similar hardware
system_opts = detect_high_end_system()
if system_opts.get('high_vram', False):
if debug:
debug.log(f"Detected high-end system optimizations: {system_opts}", category="setup")
# Apply recommended settings if not overridden
# For example, we might favor speed/quality trade-offs differently
# Here we just log it as the user has control via UI, but we could set defaults if they were None
pass
runner.debug = debug
@@ -1186,43 +1167,27 @@ def apply_model_specific_config(model: torch.nn.Module, runner: VideoDiffusionIn
"""
if is_dit:
# DiT-specific
# Determine compute_dtype upfront (respect precision setting)
compute_dtype = getattr(runner, '_compute_dtype', torch.bfloat16)
# Apply precision override if set (redundant if ctx already handled it, but safe)
precision_override = getattr(runner, '_precision', 'auto')
if precision_override == 'fp16':
compute_dtype = torch.float16
elif precision_override == 'bf16':
compute_dtype = torch.bfloat16
elif precision_override == 'bf32':
# TF32 context
compute_dtype = torch.float32
# Apply FP8 compatibility wrapper with correct compute_dtype
# Apply FP8 compatibility wrapper with compute_dtype
if not isinstance(model, FP8CompatibleDiT):
debug.log("Applying FP8/RoPE compatibility wrapper to DiT model", category="setup")
debug.start_timer("FP8CompatibleDiT")
# Get compute_dtype from runner if available, fallback to bfloat16
compute_dtype = getattr(runner, '_compute_dtype', torch.bfloat16)
model = FP8CompatibleDiT(model, debug, compute_dtype=compute_dtype, skip_conversion=False)
debug.end_timer("FP8CompatibleDiT", "FP8/RoPE compatibility wrapper application")
else:
debug.log("Reusing existing FP8/RoPE compatibility wrapper", category="reuse")
# Update compute_dtype if wrapper exists
if model.compute_dtype != compute_dtype:
debug.log(f"Updating FP8 wrapper compute_dtype to {compute_dtype}", category="setup")
model.compute_dtype = compute_dtype
# Apply attention mode and compute_dtype to all FlashAttentionVarlen modules
if hasattr(runner, '_dit_attention_mode'):
requested_attention_mode = runner._dit_attention_mode or 'sdpa'
# Validate and get final attention_mode (with warning if fallback needed)
attention_mode = validate_flash_attention_availability(requested_attention_mode, debug)
attention_mode = validate_attention_mode(requested_attention_mode, debug)
# Log final decision prominently
mode_desc = _describe_attention_mode(attention_mode)
debug.log(f"Using Attention Mode: {mode_desc}", category="info", force=True)
debug.log(f"Using Compute Dtype: {compute_dtype}", category="info", force=True)
# Get compute_dtype from runner
compute_dtype = getattr(runner, '_compute_dtype', torch.bfloat16)
debug.log(f"Applying {attention_mode} attention mode and {compute_dtype} compute dtype to model", category="setup")
# Get the actual model (unwrap if needed)
actual_model = model.dit_model if hasattr(model, 'dit_model') else model
@@ -1503,4 +1468,4 @@ def _propagate_debug_to_modules(module: torch.nn.Module, debug: 'Debug') -> None
for name, submodule in module.named_modules():
if submodule.__class__.__name__ in target_modules:
if not hasattr(submodule, 'debug'): # Only set if not already present
submodule.debug = debug
submodule.debug = debug
+2 -20
View File
@@ -204,18 +204,6 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
"• 'cuda:X': Offload to another GPU (good balance if available, faster than CPU)"
)
),
io.Combo.Input("precision",
options=["auto", "fp16", "bf16", "bf32"],
default="auto",
optional=True,
tooltip=(
"Precision for main generation process (default: auto).\n"
"• auto: Automatically select based on device capabilities\n"
"• fp16: Half precision (fastest, standard)\n"
"• bf16: BFloat16 (better dynamic range, requires Ampere+ GPU)\n"
"• bf32: Float32 with TF32 enabled (Ampere+ GPU)"
)
),
io.Boolean.Input("enable_debug",
default=False,
optional=True,
@@ -239,7 +227,7 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
uniform_batch_size: bool = False, temporal_overlap: int = 0, prepend_frames: int = 0,
color_correction: str = "wavelet", input_noise_scale: float = 0.0,
latent_noise_scale: float = 0.0, offload_device: str = "none",
precision: str = "auto", enable_debug: bool = False) -> io.NodeOutput:
enable_debug: bool = False) -> io.NodeOutput:
"""
Execute SeedVR2 video upscaling with progress reporting
@@ -354,10 +342,6 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
attention_mode = dit.get("attention_mode", "sdpa")
vae_cache = vae.get("cache_model", False)
# Override attention mode if specified in dit config but allow validation later
if "attention_mode" in dit:
attention_mode = dit["attention_mode"]
# BlockSwap configuration - construct from individual values
blocks_to_swap = dit.get("blocks_to_swap", 0)
swap_io_components = dit.get("swap_io_components", False)
@@ -424,8 +408,7 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
dit_offload_device=dit_offload_device,
vae_offload_device=vae_offload_device,
tensor_offload_device=tensor_offload_device,
debug=debug,
precision=precision
debug=debug
)
# Prepare runner with model state management and global cache
@@ -448,7 +431,6 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
decode_tile_overlap=(decode_tile_overlap, decode_tile_overlap),
tile_debug=tile_debug,
attention_mode=attention_mode,
precision=precision,
torch_compile_args_dit=dit_torch_compile_args,
torch_compile_args_vae=vae_torch_compile_args
)
+7 -92
View File
@@ -15,16 +15,11 @@
import torch
import torch.nn.functional as F
# Import flash_attn with automatic fallback from compatibility layer
from ...optimization.compatibility import flash_attn_varlen_func, FLASH_ATTN_AVAILABLE, SAGE_ATTN_AVAILABLE
# Import flash/sage attn with automatic fallback from compatibility layer
from ...optimization.compatibility import call_flash_attn_varlen, call_sage_attn_varlen
from torch import nn
# Safe import for SageAttention
try:
import sageattention
except ImportError:
sageattention = None
def pytorch_varlen_attention(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.0, softmax_scale=None, causal=False, deterministic=False):
"""
@@ -66,81 +61,6 @@ def pytorch_varlen_attention(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q=N
return torch.cat(output_splits, dim=0)
@torch._dynamo.disable
def _call_flash_attn_varlen_func(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.
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
"""
if not FLASH_ATTN_AVAILABLE:
raise ImportError("flash_attn 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())
return flash_attn_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,
**kwargs
)
@torch._dynamo.disable
def _call_sage_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal=False, implementation="sa2"):
"""
Wrapper for SageAttention variable length function.
Args:
implementation: "sa2" (SageAttention v2) or "sa3" (SageAttention v3)
"""
if not SAGE_ATTN_AVAILABLE:
raise ImportError("SageAttention is not available")
# SageAttention expects q, k, v as (total_tokens, heads, head_dim)
# The input q, k, v here are (total_tokens, heads, head_dim)
# 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())
# Ensure tensors are contiguous
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
# SageAttention API usage
# sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal, sm_scale)
try:
from sageattention import sageattn_varlen
except ImportError:
# Fallback or error
raise ImportError("sageattn_varlen not found in sageattention package")
# Check if sm_scale is needed (usually 1/sqrt(head_dim))
sm_scale = 1.0 / (q.shape[-1] ** 0.5)
# Calling sageattn_varlen
# Signature assumptions: q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal, sm_scale
if not hasattr(_call_sage_attn_varlen_func, "_logged"):
print(f"🚀 Executing SageAttention ({implementation}) kernel for the first time")
_call_sage_attn_varlen_func._logged = True
return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal, sm_scale)
class TorchAttention(nn.Module):
def tflops(self, args, kwargs, output) -> float:
assert len(args) == 0 or len(args) > 2, "query, key should both provided by args / kwargs"
@@ -169,7 +89,7 @@ class FlashAttentionVarlen(nn.Module):
Initialize with specified attention backend.
Args:
attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_flash_attention_availability)
attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_attention_mode)
compute_dtype: Compute dtype for attention (set by pipeline, defaults to None for auto-detection)
"""
super().__init__()
@@ -194,19 +114,14 @@ class FlashAttentionVarlen(nn.Module):
v = v.to(self.compute_dtype)
if self.attention_mode == 'flash_attn':
return _call_flash_attn_varlen_func(
return call_flash_attn_varlen(
q, k, v, cu_seqlens_q, cu_seqlens_k,
max_seqlen_q, max_seqlen_k, **kwargs
)
elif self.attention_mode in ['sa2', 'sa3']:
# Use SageAttention
# Extract causal flag if present in kwargs, default to False
is_causal = kwargs.get('causal', False)
return _call_sage_attn_varlen_func(
elif self.attention_mode in ('sa2', 'sa3'):
return call_sage_attn_varlen(
q, k, v, cu_seqlens_q, cu_seqlens_k,
max_seqlen_q, max_seqlen_k,
is_causal=is_causal,
implementation=self.attention_mode
max_seqlen_q, max_seqlen_k, **kwargs
)
else:
# PyTorch SDPA
+7 -92
View File
@@ -15,16 +15,11 @@
import torch
import torch.nn.functional as F
# Import flash_attn with automatic fallback from compatibility layer
from ...optimization.compatibility import flash_attn_varlen_func, FLASH_ATTN_AVAILABLE, SAGE_ATTN_AVAILABLE
# Import flash/sage attn with automatic fallback from compatibility layer
from ...optimization.compatibility import call_flash_attn_varlen, call_sage_attn_varlen
from torch import nn
# Safe import for SageAttention
try:
import sageattention
except ImportError:
sageattention = None
def pytorch_varlen_attention(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.0, softmax_scale=None, causal=False, deterministic=False):
"""
@@ -66,81 +61,6 @@ def pytorch_varlen_attention(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q=N
return torch.cat(output_splits, dim=0)
@torch._dynamo.disable
def _call_flash_attn_varlen_func(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.
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
"""
if not FLASH_ATTN_AVAILABLE:
raise ImportError("flash_attn 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())
return flash_attn_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,
**kwargs
)
@torch._dynamo.disable
def _call_sage_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal=False, implementation="sa2"):
"""
Wrapper for SageAttention variable length function.
Args:
implementation: "sa2" (SageAttention v2) or "sa3" (SageAttention v3)
"""
if not SAGE_ATTN_AVAILABLE:
raise ImportError("SageAttention is not available")
# SageAttention expects q, k, v as (total_tokens, heads, head_dim)
# The input q, k, v here are (total_tokens, heads, head_dim)
# 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())
# Ensure tensors are contiguous
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
# SageAttention API usage
# sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal, sm_scale)
try:
from sageattention import sageattn_varlen
except ImportError:
# Fallback or error
raise ImportError("sageattn_varlen not found in sageattention package")
# Check if sm_scale is needed (usually 1/sqrt(head_dim))
sm_scale = 1.0 / (q.shape[-1] ** 0.5)
# Calling sageattn_varlen
# Signature assumptions: q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal, sm_scale
if not hasattr(_call_sage_attn_varlen_func, "_logged"):
print(f"🚀 Executing SageAttention ({implementation}) kernel for the first time")
_call_sage_attn_varlen_func._logged = True
return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, is_causal, sm_scale)
class TorchAttention(nn.Module):
def tflops(self, args, kwargs, output) -> float:
assert len(args) == 0 or len(args) > 2, "query, key should both provided by args / kwargs"
@@ -169,7 +89,7 @@ class FlashAttentionVarlen(nn.Module):
Initialize with specified attention backend.
Args:
attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_flash_attention_availability)
attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_attention_mode)
compute_dtype: Compute dtype for attention (set by pipeline, defaults to None for auto-detection)
"""
super().__init__()
@@ -194,19 +114,14 @@ class FlashAttentionVarlen(nn.Module):
v = v.to(self.compute_dtype)
if self.attention_mode == 'flash_attn':
return _call_flash_attn_varlen_func(
return call_flash_attn_varlen(
q, k, v, cu_seqlens_q, cu_seqlens_k,
max_seqlen_q, max_seqlen_k, **kwargs
)
elif self.attention_mode in ['sa2', 'sa3']:
# Use SageAttention
# Extract causal flag if present in kwargs, default to False
is_causal = kwargs.get('causal', False)
return _call_sage_attn_varlen_func(
elif self.attention_mode in ('sa2', 'sa3'):
return call_sage_attn_varlen(
q, k, v, cu_seqlens_q, cu_seqlens_k,
max_seqlen_q, max_seqlen_k,
is_causal=is_causal,
implementation=self.attention_mode
max_seqlen_q, max_seqlen_k, **kwargs
)
else:
# PyTorch SDPA
+122 -111
View File
@@ -89,10 +89,9 @@ ensure_xformers_flash_compat()
import torch
import os
from typing import Dict, Any, Optional
# Flash Attention & Triton Compatibility Layer
# Flash/Sage Attention & Triton Compatibility Layer
# 1. Flash Attention - speedup for attention operations
try:
from flash_attn import flash_attn_varlen_func
@@ -103,33 +102,50 @@ except (ImportError, AttributeError, OSError):
flash_attn_varlen_func = None
FLASH_ATTN_AVAILABLE = False
# 1.1 SageAttention - speedup for attention operations
# 2. SageAttention - speedup for attention operations
try:
import sageattention
from sageattention import sageattn_varlen
SAGE_ATTN_AVAILABLE = True
try:
from sageattention import sageattn_varlen
# Basic check to see if it's functional or mock
SAGE_ATTN_VARLEN_AVAILABLE = True
except ImportError:
SAGE_ATTN_VARLEN_AVAILABLE = False
except ImportError:
except (ImportError, AttributeError, OSError):
sageattn_varlen = None
SAGE_ATTN_AVAILABLE = False
SAGE_ATTN_VARLEN_AVAILABLE = False
def validate_flash_attention_availability(requested_mode: str, debug=None) -> str:
def validate_attention_mode(requested_mode: str, debug=None) -> str:
"""
Validate attention mode availability and warn if fallback needed.
Validate attention mode availability with automatic fallback to sdpa.
Args:
requested_mode: 'flash_attn', 'sdpa', 'sd2', or 'sd3'
requested_mode: 'sdpa', 'flash_attn', 'sa2', or 'sa3'
debug: Optional debug instance for logging
Returns:
Validated mode
Validated mode that is available
"""
if requested_mode == 'flash_attn' and not FLASH_ATTN_AVAILABLE:
# SageAttention modes
if requested_mode in ('sa2', 'sa3'):
if SAGE_ATTN_AVAILABLE:
return requested_mode
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"
)
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:
return requested_mode
error_msg = (
f"Cannot use 'flash_attn' attention mode: Flash Attention is not installed.\n"
f"\n"
@@ -144,42 +160,73 @@ def validate_flash_attention_availability(requested_mode: str, debug=None) -> st
)
if debug:
debug.log(error_msg, level="WARNING", category="setup", force=True)
return 'sdpa'
if requested_mode in ['sa2', 'sa3']:
if not SAGE_ATTN_AVAILABLE:
if debug:
debug.log(f"SageAttention not installed. Falling back from '{requested_mode}' to Flash Attention 2...", level="WARNING", category="setup", force=True)
# Fallback to check FA2
return validate_flash_attention_availability('flash_attn', debug)
elif not SAGE_ATTN_VARLEN_AVAILABLE:
if debug:
debug.log(f"SageAttention installed but 'sageattn_varlen' not found. Falling back from '{requested_mode}' to Flash Attention 2...", level="WARNING", category="setup", force=True)
# Fallback to check FA2
return validate_flash_attention_availability('flash_attn', debug)
# If the user explicitly requested sa3, we check for version compatibility.
# If version is unknown or insufficient, we fallback to sa2.
if requested_mode == 'sa3':
try:
version = sageattention.__version__
# Assuming sa3 requires at least a certain version or just presence of version string.
# If we can read version, we assume it's compliant enough or user knows what they are doing.
if debug:
debug.log(f"SageAttention version {version} detected. Using installed kernel for 'sa3' mode.", category="setup", force=True)
except AttributeError:
# Version unknown -> Assume it's an older version (sa2) and fallback
if debug:
debug.log("SageAttention version unknown (likely v2 or older). Falling back from 'sa3' to 'sa2'...", level="WARNING", category="setup", force=True)
return validate_flash_attention_availability('sa2', debug)
pass
return requested_mode
@torch._dynamo.disable
def call_flash_attn_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.
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
"""
if not FLASH_ATTN_AVAILABLE:
raise ImportError("flash_attn 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())
return flash_attn_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,
**kwargs
)
@torch._dynamo.disable
def call_sage_attn_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.
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
"""
if not SAGE_ATTN_AVAILABLE:
raise ImportError("SageAttention 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 contiguous tensors
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
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)
# 2. Triton - Required for torch.compile with inductor backend
try:
import triton
@@ -278,21 +325,34 @@ NVIDIA_CONV3D_MEMORY_BUG_WORKAROUND = _check_conv3d_memory_bug()
if not os.environ.get("SEEDVR2_OPTIMIZATIONS_LOGGED"):
os.environ["SEEDVR2_OPTIMIZATIONS_LOGGED"] = "1"
# Flash Attention & Triton status
has_both = FLASH_ATTN_AVAILABLE and TRITON_AVAILABLE
has_neither = not FLASH_ATTN_AVAILABLE and not TRITON_AVAILABLE
# Build status strings
sage_status = "✅" if SAGE_ATTN_AVAILABLE else "❌"
flash_status = "✅" if FLASH_ATTN_AVAILABLE else "❌"
triton_status = "✅" if TRITON_AVAILABLE else "❌"
if has_both:
print("⚡ SeedVR2 optimizations check: Flash Attention ✅ | Triton ✅")
elif has_neither:
print("⚠️ SeedVR2 optimizations check: Flash Attention ❌ | Triton ❌")
print("💡 For best performance: pip install flash-attn triton")
elif FLASH_ATTN_AVAILABLE:
print("⚡ SeedVR2 optimizations check: Flash Attention ✅ | Triton ❌")
print("💡 Install Triton for torch.compile: pip install triton")
else: # TRITON_AVAILABLE only
print("⚠️ SeedVR2 optimizations check: Flash Attention ❌ | Triton ✅")
print("💡 Install Flash Attention for faster inference: pip install flash-attn")
# Count available optimizations
available = [SAGE_ATTN_AVAILABLE, FLASH_ATTN_AVAILABLE, TRITON_AVAILABLE]
num_available = sum(available)
if num_available == 3:
print(f"⚡ SeedVR2 optimizations check: SageAttention {sage_status} | Flash Attention {flash_status} | Triton {triton_status}")
elif num_available == 0:
print(f"⚠️ SeedVR2 optimizations check: SageAttention {sage_status} | Flash Attention {flash_status} | Triton {triton_status}")
print("💡 For best performance: pip install sageattention flash-attn triton")
else:
icon = "⚡" if num_available >= 2 else "⚠️ "
print(f"{icon} SeedVR2 optimizations check: SageAttention {sage_status} | Flash Attention {flash_status} | Triton {triton_status}")
# Build install suggestions for missing packages
missing = []
if not SAGE_ATTN_AVAILABLE:
missing.append("sageattention")
if not FLASH_ATTN_AVAILABLE:
missing.append("flash-attn")
if not TRITON_AVAILABLE:
missing.append("triton")
if missing:
print(f"💡 Optional: pip install {' '.join(missing)}")
# Conv3d workaround status (if applicable)
if NVIDIA_CONV3D_MEMORY_BUG_WORKAROUND:
@@ -318,55 +378,6 @@ def _probe_bfloat16_support() -> bool:
BFLOAT16_SUPPORTED = _probe_bfloat16_support()
COMPUTE_DTYPE = torch.bfloat16 if BFLOAT16_SUPPORTED else torch.float16
def log_system_capabilities(debug=None):
"""Log installed attention backends and versions at startup."""
if not debug:
return
# SageAttention
sa_status = "Available" if SAGE_ATTN_AVAILABLE else "Not Installed"
if SAGE_ATTN_AVAILABLE:
try:
sa_version = sageattention.__version__
sa_status += f" (v{sa_version})"
except AttributeError:
sa_status += " (Version Unknown)"
# FlashAttention
fa_status = "Available" if FLASH_ATTN_AVAILABLE else "Not Installed"
# Triton
triton_status = "Available" if TRITON_AVAILABLE else "Not Installed"
debug.log(f"Attention Backends: SageAttention={sa_status} | FlashAttention={fa_status} | Triton={triton_status}", category="info", force=True)
def detect_high_end_system() -> Dict[str, Any]:
"""
Detect high-end systems (16GB+ VRAM, etc.) and return optimized defaults.
Returns:
Dict with recommended settings or empty if no specific optimizations found.
"""
optimizations = {}
try:
# Basic VRAM check
if torch.cuda.is_available():
device = torch.device("cuda:0")
props = torch.cuda.get_device_properties(device)
total_vram_gb = props.total_memory / (1024**3)
# High-end GPU check (e.g., 5070ti/4080/4090/etc with >15GB VRAM)
if total_vram_gb >= 15.5:
optimizations['high_vram'] = True
optimizations['recommended_dtype'] = 'bf16' if BFLOAT16_SUPPORTED else 'fp16'
# For 16GB cards, BlockSwap might still be useful for 7B models but maybe less aggressive
optimizations['block_swap_recommendation'] = 'moderate'
except Exception:
pass
return optimizations
def call_rope_with_stability(method, *args, **kwargs):
"""