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