Azan commited on
Commit ·
9e9ce7f
1
Parent(s): 91de477
Fix: Force Bfloat16 on CPU to satisfy PyTorch autocast requirements
Browse files
ylff/utils/model_loader.py
CHANGED
|
@@ -236,10 +236,11 @@ def load_da3_model(
|
|
| 236 |
device = "cpu"
|
| 237 |
|
| 238 |
# Monkeypatch torch.cuda.is_bf16_supported to avoid initializing CUDA on CPU machines
|
| 239 |
-
#
|
|
|
|
| 240 |
if hasattr(torch.cuda, "is_bf16_supported"):
|
| 241 |
-
logger.info("Patching torch.cuda.is_bf16_supported to
|
| 242 |
-
torch.cuda.is_bf16_supported = lambda:
|
| 243 |
|
| 244 |
logger.info(f"Loading DA3 model: {model_name} on {device}")
|
| 245 |
model = DepthAnything3.from_pretrained(model_name)
|
|
|
|
| 236 |
device = "cpu"
|
| 237 |
|
| 238 |
# Monkeypatch torch.cuda.is_bf16_supported to avoid initializing CUDA on CPU machines
|
| 239 |
+
# AND return True to force usage of bfloat16, which is required for CPU autocast
|
| 240 |
+
# (PyTorch CPU AMP does not support float16, only bfloat16)
|
| 241 |
if hasattr(torch.cuda, "is_bf16_supported"):
|
| 242 |
+
logger.info("Patching torch.cuda.is_bf16_supported=True to enable CPU bfloat16 autocast")
|
| 243 |
+
torch.cuda.is_bf16_supported = lambda: True
|
| 244 |
|
| 245 |
logger.info(f"Loading DA3 model: {model_name} on {device}")
|
| 246 |
model = DepthAnything3.from_pretrained(model_name)
|