Files
feedmill-auto-label/backend/hardware.py
T

67 lines
2.2 KiB
Python

"""Training defaults derived from the machine this happens to be running on.
The point is REQ-062: moving to a bigger GPU should change the numbers in the
form, not the code. Everything here is a *default* — the user can override all
of it per run.
"""
from typing import Optional
SAM3_RESIDENT_GB = 3.9
SAM3_HEADROOM_GB = 0.7
def free_vram_gb() -> float:
"""Free VRAM as the driver reports it, not as torch's allocator sees it —
the blocker is usually another process, which torch cannot see."""
import torch
if not torch.cuda.is_available():
return 0.0
free, _total = torch.cuda.mem_get_info()
return round(free / (1024 ** 3), 2)
def detect() -> dict:
import torch
if not torch.cuda.is_available():
return {"device": "cpu", "gpu": None, "vram_gb": 0.0}
properties = torch.cuda.get_device_properties(0)
return {
"device": "cuda",
"gpu": properties.name,
"vram_gb": round(properties.total_memory / (1024 ** 3), 1),
}
def defaults(epochs: int = 50) -> dict:
"""Batch size and image size that should fit, given the VRAM we can see."""
info = detect()
vram = info["vram_gb"]
if info["device"] == "cpu":
settings = {"batch": 4, "imgsz": 512, "device": "cpu", "workers": 2}
note = "No GPU visible — training on CPU will be very slow."
elif vram < 8:
settings = {"batch": 8, "imgsz": 640, "device": 0, "workers": 2}
note = f"{vram} GB of VRAM: small batches, 640 px."
elif vram <= 16:
settings = {"batch": 32, "imgsz": 640, "device": 0, "workers": 8}
note = f"{vram} GB of VRAM: optimized batch 32, 640 px."
else:
settings = {"batch": 32, "imgsz": 768, "device": 0, "workers": 8}
note = f"{vram} GB of VRAM: room for larger batches and 768 px."
return {**info, **settings, "epochs": epochs, "note": note}
def resolve(overrides: Optional[dict] = None, epochs: int = 50) -> dict:
"""Defaults with the user's overrides applied on top."""
settings = defaults(epochs)
for key, value in (overrides or {}).items():
if value is not None and key in ("batch", "imgsz", "device", "epochs", "workers"):
settings[key] = value
return settings