diff --git a/src/core/generation_utils.py b/src/core/generation_utils.py index e84bd98..0aba160 100644 --- a/src/core/generation_utils.py +++ b/src/core/generation_utils.py @@ -36,6 +36,7 @@ from .model_configuration import configure_runner from .infer import VideoDiffusionInfer from ..data.image.transforms.divisible_crop import DivisiblePad from ..data.image.transforms.na_resize import NaResize +from ..optimization.compatibility import COMPUTE_DTYPE, BFLOAT16_SUPPORTED from ..optimization.memory_manager import manage_tensor from ..utils.constants import get_script_directory @@ -371,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': torch.bfloat16, # Hardcoded - gives the best compromise between memory & quality without artifacts + 'compute_dtype': COMPUTE_DTYPE, 'interrupt_fn': interrupt_fn, 'video_transform': None, 'text_embeds': None, @@ -401,7 +402,12 @@ def setup_generation_context( f"LOCAL_RANK={os.environ['LOCAL_RANK']}", category="setup" ) - reason = "quality" if ctx['compute_dtype'] == torch.float32 else "compatibility" + 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 diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index 4326308..d154dd0 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -41,39 +41,6 @@ import types import os -# Automatic bfloat16 SDPA fallback for GPUs that don't support it (e.g., GTX 970) -_BFLOAT16_SDPA_WORKS = None # None=untested, True=works, False=needs float16 fallback -_ORIGINAL_SDPA = torch.nn.functional.scaled_dot_product_attention - -def _safe_scaled_dot_product_attention(query, key, value, *args, **kwargs): - """SDPA wrapper with automatic bfloat16 -> float16 fallback for old GPUs.""" - global _BFLOAT16_SDPA_WORKS - - original_dtype = query.dtype - - # Fast path: already know bfloat16 fails on this GPU - if original_dtype == torch.bfloat16 and _BFLOAT16_SDPA_WORKS is False: - out = _ORIGINAL_SDPA(query.half(), key.half(), value.half(), *args, **kwargs) - return out.to(original_dtype) - - try: - out = _ORIGINAL_SDPA(query, key, value, *args, **kwargs) - if _BFLOAT16_SDPA_WORKS is None and original_dtype == torch.bfloat16: - _BFLOAT16_SDPA_WORKS = True - return out - except RuntimeError as e: - if "CUBLAS_STATUS_NOT_SUPPORTED" in str(e) and original_dtype == torch.bfloat16: - _BFLOAT16_SDPA_WORKS = False - print("⚠️ [SeedVR2] GPU does not support bfloat16 SDPA, using float16 fallback. " - "Tiling artifacts or black frames may occur.") - out = _ORIGINAL_SDPA(query.half(), key.half(), value.half(), *args, **kwargs) - return out.to(original_dtype) - raise - -# Apply SDPA patch at module load -torch.nn.functional.scaled_dot_product_attention = _safe_scaled_dot_product_attention - - # Flash Attention & Triton Compatibility Layer # 1. Flash Attention - speedup for attention operations try: @@ -236,6 +203,24 @@ if not os.environ.get("SEEDVR2_OPTIMIZATIONS_LOGGED"): print(f"🔧 Conv3d workaround active: PyTorch {torch_ver}, cuDNN {cudnn_ver} (fixing VAE 3x memory bug)") +# Bfloat16 CUBLAS support +def _probe_bfloat16_support() -> bool: + if not torch.cuda.is_available(): + return True + try: + a = torch.randn(8, 8, dtype=torch.bfloat16, device='cuda:0') + _ = torch.matmul(a, a) + del a + return True + except RuntimeError as e: + if "CUBLAS_STATUS_NOT_SUPPORTED" in str(e): + return False + raise + +BFLOAT16_SUPPORTED = _probe_bfloat16_support() +COMPUTE_DTYPE = torch.bfloat16 if BFLOAT16_SUPPORTED else torch.float16 + + def call_rope_with_stability(method, *args, **kwargs): """ Call RoPE method with stability fixes: