diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index d154dd0..c63b007 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -282,11 +282,6 @@ class FP8CompatibleDiT(torch.nn.Module): self.debug.start_timer("_stabilize_rope_computations") self._stabilize_rope_computations() self.debug.end_timer("_stabilize_rope_computations", "RoPE stabilization") - - # 🚀 FLASH ATTENTION OPTIMIZATION (Phase 2) - self.debug.start_timer("_apply_flash_attention_optimization") - self._apply_flash_attention_optimization() - self.debug.end_timer("_apply_flash_attention_optimization", "Flash Attention application") def _detect_model_dtype(self) -> torch.dtype: """Detect main model dtype""" @@ -409,210 +404,6 @@ class FP8CompatibleDiT(torch.nn.Module): if rope_count > 0: self.debug.log(f"Stabilized {rope_count} RoPE modules", category="success") - - def _apply_flash_attention_optimization(self) -> None: - """🚀 FLASH ATTENTION OPTIMIZATION - 30-50% speedup of attention layers""" - attention_layers_optimized = 0 - flash_attention_available = self._check_flash_attention_support() - - for name, module in self.dit_model.named_modules(): - # Identify all attention layers - if self._is_attention_layer(name, module): - # Apply optimization based on availability - if self._optimize_attention_layer(name, module, flash_attention_available): - attention_layers_optimized += 1 - - if not flash_attention_available: - self.debug.log("Flash Attention not available, using PyTorch SDPA as fallback", category="info", force=True) - - def _check_flash_attention_support(self) -> bool: - """Check if Flash Attention is available""" - # Check PyTorch SDPA (includes Flash Attention on H100/A100) - if hasattr(torch.nn.functional, 'scaled_dot_product_attention'): - return True - - # Check flash-attn package (uses module-level check from top of file) - return FLASH_ATTN_AVAILABLE - - def _is_attention_layer(self, name: str, module: torch.nn.Module) -> bool: - """Identify if a module is an attention layer""" - attention_keywords = [ - 'attention', 'attn', 'self_attn', 'cross_attn', 'mhattn', 'multihead', - 'transformer_block', 'dit_block' - ] - - # Check by name - if any(keyword in name.lower() for keyword in attention_keywords): - return True - - # Check by module type - module_type = type(module).__name__.lower() - if any(keyword in module_type for keyword in attention_keywords): - return True - - # Check by attributes (modules with q, k, v projections) - if hasattr(module, 'q_proj') or hasattr(module, 'qkv') or hasattr(module, 'to_q'): - return True - - return False - - def _optimize_attention_layer(self, name: str, module: torch.nn.Module, flash_attention_available: bool) -> bool: - """Optimize a specific attention layer""" - try: - # Save original forward method - if not hasattr(module, '_original_forward'): - module._original_forward = module.forward - - # Create new optimized forward method - if flash_attention_available: - optimized_forward = self._create_flash_attention_forward(module, name) - else: - optimized_forward = self._create_sdpa_forward(module, name) - - # Replace forward method - module.forward = optimized_forward - return True - - except Exception as e: - self.debug.log(f"Failed to optimize attention layer '{name}': {e}", level="WARNING", category="dit", force=True) - return False - - def _create_flash_attention_forward(self, module: torch.nn.Module, layer_name: str): - """Create optimized forward with Flash Attention""" - original_forward = module._original_forward - - def flash_attention_forward(*args, **kwargs): - try: - # Try to use Flash Attention via SDPA - return self._sdpa_attention_forward(original_forward, module, *args, **kwargs) - except Exception as e: - # Fallback to original implementation - self.debug.log(f"Flash Attention failed for {layer_name}, using original: {e}", level="WARNING", category="dit", force=True) - return original_forward(*args, **kwargs) - - return flash_attention_forward - - def _create_sdpa_forward(self, module: torch.nn.Module, layer_name: str): - """Create optimized forward with PyTorch SDPA""" - original_forward = module._original_forward - - def sdpa_forward(*args, **kwargs): - try: - return self._sdpa_attention_forward(original_forward, module, *args, **kwargs) - except Exception as e: - # Fallback to original implementation - return original_forward(*args, **kwargs) - - return sdpa_forward - - def _sdpa_attention_forward(self, original_forward, module: torch.nn.Module, *args, **kwargs): - """Optimized forward pass using SDPA (Scaled Dot Product Attention)""" - # Detect if we can intercept and optimize this layer - if len(args) >= 1 and isinstance(args[0], torch.Tensor): - input_tensor = args[0] - - # Check dimensions to ensure it's standard attention - if len(input_tensor.shape) >= 3: # [batch, seq_len, hidden_dim] or similar - try: - return self._optimized_attention_computation(module, input_tensor, *args[1:], **kwargs) - except: - pass - - # Fallback to original implementation - return original_forward(*args, **kwargs) - - def _optimized_attention_computation(self, module: torch.nn.Module, input_tensor: torch.Tensor, *args, **kwargs): - """Optimized attention computation with SDPA""" - # Try to detect standard attention format - batch_size, seq_len = input_tensor.shape[:2] - - # Check if module has standard Q, K, V projections - if hasattr(module, 'qkv') or (hasattr(module, 'q_proj') and hasattr(module, 'k_proj') and hasattr(module, 'v_proj')): - return self._compute_sdpa_attention(module, input_tensor, *args, **kwargs) - - # If no standard format detected, use original - return module._original_forward(input_tensor, *args, **kwargs) - - def _compute_sdpa_attention(self, module: torch.nn.Module, x: torch.Tensor, *args, **kwargs): - """Optimized SDPA computation for standard attention modules""" - try: - # Case 1: Module with combined QKV projection - if hasattr(module, 'qkv'): - qkv = module.qkv(x) - # Reshape to separate Q, K, V - batch_size, seq_len, _ = qkv.shape - qkv = qkv.reshape(batch_size, seq_len, 3, -1) - q, k, v = qkv.unbind(dim=2) - - # Case 2: Separate Q, K, V projections - elif hasattr(module, 'q_proj') and hasattr(module, 'k_proj') and hasattr(module, 'v_proj'): - q = module.q_proj(x) - k = module.k_proj(x) - v = module.v_proj(x) - else: - # Unsupported format, use original - return module._original_forward(x, *args, **kwargs) - - # Detect number of heads - head_dim = getattr(module, 'head_dim', None) - num_heads = getattr(module, 'num_heads', None) - - if head_dim is None or num_heads is None: - # Try to guess from dimensions - hidden_dim = q.shape[-1] - if hasattr(module, 'num_heads'): - num_heads = module.num_heads - head_dim = hidden_dim // num_heads - else: - # Reasonable defaults - head_dim = 64 - num_heads = hidden_dim // head_dim - - # Reshape for multi-head attention - batch_size, seq_len = q.shape[:2] - q = q.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2) - k = k.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2) - v = v.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2) - - if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): - attn_output = torch.nn.functional.scaled_dot_product_attention( - q, k, v, - dropout_p=0.0, - is_causal=False - ) - else: - # Use optimized SDPA - PyTorch 2.3+ API with CUDNN support, fallback for older versions - if hasattr(torch.nn.attention, 'sdpa_kernel'): - ctx = torch.nn.attention.sdpa_kernel([ - torch.nn.attention.SDPBackend.FLASH_ATTENTION, - torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION, - torch.nn.attention.SDPBackend.CUDNN_ATTENTION, - torch.nn.attention.SDPBackend.MATH]) - else: - ctx = torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=True, enable_mem_efficient=True) - - with ctx: - attn_output = torch.nn.functional.scaled_dot_product_attention( - q, k, v, - dropout_p=0.0, - is_causal=False - ) - - # Reshape back - attn_output = attn_output.transpose(1, 2).contiguous().view( - batch_size, seq_len, num_heads * head_dim - ) - - # Output projection if it exists - if hasattr(module, 'out_proj') or hasattr(module, 'o_proj'): - proj = getattr(module, 'out_proj', None) or getattr(module, 'o_proj', None) - attn_output = proj(attn_output) - - return attn_output - - except Exception as e: - # In case of error, use original implementation - return module._original_forward(x, *args, **kwargs) def forward(self, *args, **kwargs): """