Revert bfloat16 detection - was causing quality regression / keep ensure_triton_compat()

This commit is contained in:
Adrien Toupet
2025-12-05 15:03:40 -05:00
parent 43c4f00e19
commit f061d97fe7
-35
View File
@@ -175,41 +175,6 @@ def _check_conv3d_memory_bug():
NVIDIA_CONV3D_MEMORY_BUG_WORKAROUND = _check_conv3d_memory_bug()
def get_supported_compute_dtype(debug=None) -> torch.dtype:
"""
Detect the best supported compute dtype for the current hardware.
bfloat16 requires CUDA compute capability 8.0+ (Ampere and newer).
Older CUDA GPUs fall back to float16 to avoid CUBLAS_STATUS_NOT_SUPPORTED errors.
Args:
debug: Optional debug instance for logging warnings
Returns:
torch.bfloat16 if supported, torch.float16 for older CUDA GPUs
"""
try:
# CUDA: Check compute capability
if torch.cuda.is_available():
major, minor = torch.cuda.get_device_capability()
# bfloat16 requires compute capability 8.0+ (Ampere: A100, RTX 30xx)
if major < 8:
if debug:
gpu_name = torch.cuda.get_device_name()
debug.log(
f"GPU '{gpu_name}' (compute capability {major}.{minor}) does not support bfloat16. "
f"Using float16 instead - 7B models are not supported and will render black, 3B models may have artifacts.",
level="WARNING", category="precision", force=True
)
return torch.float16
# Default: bfloat16 (MPS, CPU, or CUDA 8.0+)
return torch.bfloat16
except Exception:
return torch.bfloat16
# Log all optimization status once globally (cross-process) using environment variable
if not os.environ.get("SEEDVR2_OPTIMIZATIONS_LOGGED"):
os.environ["SEEDVR2_OPTIMIZATIONS_LOGGED"] = "1"