67 lines
2.2 KiB
Python
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
|