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