Remove dead flash attention wrapper from FP8CompatibleDiT

The wrapper methods (_apply_flash_attention_optimization and related)
matched NaDiT attention modules by name but required qkv or q_proj+k_proj+v_proj
attributes to optimize. NaDiT uses proj_qkv instead, so the optimization
path was never taken - always falling back to original forward.

FlashAttentionVarlen already handles flash_attn vs sdpa switching via
its attention_mode attribute, making this wrapper redundant.

Removes ~200 lines of dead code.
This commit is contained in:
Adrien Toupet
2025-12-09 12:29:15 -05:00
parent a06afb5956
commit e65e7fa418
-209
View File
@@ -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):
"""