From bcfbca6ae382471df01a607d1cd3988d56d9ba1f Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 10 Dec 2025 13:58:46 -0500 Subject: [PATCH] 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 --- inference_cli.py | 8 +- src/core/generation_utils.py | 39 ++--- src/core/model_configuration.py | 53 ++----- src/interfaces/video_upscaler.py | 22 +-- src/models/dit_3b/attention.py | 99 +------------ src/models/dit_7b/attention.py | 99 +------------ src/optimization/compatibility.py | 233 ++++++++++++++++-------------- 7 files changed, 158 insertions(+), 395 deletions(-) diff --git a/inference_cli.py b/inference_cli.py index c995887..f19b52d 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -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", diff --git a/src/core/generation_utils.py b/src/core/generation_utils.py index 3a90ed5..6b21d96 100644 --- a/src/core/generation_utils.py +++ b/src/core/generation_utils.py @@ -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) \ No newline at end of file diff --git a/src/core/model_configuration.py b/src/core/model_configuration.py index 5bde2b7..4190861 100644 --- a/src/core/model_configuration.py +++ b/src/core/model_configuration.py @@ -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 \ No newline at end of file diff --git a/src/interfaces/video_upscaler.py b/src/interfaces/video_upscaler.py index e6e212f..64d8815 100644 --- a/src/interfaces/video_upscaler.py +++ b/src/interfaces/video_upscaler.py @@ -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 ) diff --git a/src/models/dit_3b/attention.py b/src/models/dit_3b/attention.py index cd6741b..d6ace0d 100644 --- a/src/models/dit_3b/attention.py +++ b/src/models/dit_3b/attention.py @@ -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 diff --git a/src/models/dit_7b/attention.py b/src/models/dit_7b/attention.py index cd6741b..d6ace0d 100644 --- a/src/models/dit_7b/attention.py +++ b/src/models/dit_7b/attention.py @@ -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 diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index 7c6a254..5b79853 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -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): """