diff --git a/studio/backend/core/training/diffusion_train_common.py b/studio/backend/core/training/diffusion_train_common.py index 9d1dfc3475..dbbcda32d7 100644 --- a/studio/backend/core/training/diffusion_train_common.py +++ b/studio/backend/core/training/diffusion_train_common.py @@ -264,7 +264,6 @@ def bf16_unsupported_reason(resolved_family: str) -> Optional[str]: return None try: import torch - if torch.cuda.is_available() and not torch.cuda.is_bf16_supported(): return ( "This trainer requires a bfloat16-capable GPU (Ampere or newer); this CUDA " diff --git a/studio/backend/tests/test_diffusion_base_precision.py b/studio/backend/tests/test_diffusion_base_precision.py index 88f83ce307..c6cf0ec33c 100644 --- a/studio/backend/tests/test_diffusion_base_precision.py +++ b/studio/backend/tests/test_diffusion_base_precision.py @@ -81,7 +81,9 @@ def test_base_precision_denies_fp8_for_corrupted_family(): # The deny is fp8-specific: int8 (per-token, unaffected) and the other dense modes stay # allowed for the same Qwen base. for mode in ("nf4", "bf16", "int8", "auto"): - norm = _cfg(base_model = _QWEN_DENSE, base_precision = mode, mixed_precision = "bf16").normalized() + norm = _cfg( + base_model = _QWEN_DENSE, base_precision = mode, mixed_precision = "bf16" + ).normalized() assert norm.resolved_family == "qwen-image" assert norm.base_precision == mode