"""Augmentation settings, stored per project and passed to Ultralytics (REQ-110…113). Ultralytics augments during training whether or not we ask it to. Before this module, `training.py` passed no augmentation arguments at all, so every run used library defaults invisibly. MEDIUM below *is* that default set — a project that has never been touched trains exactly as it did, only now it is written down. Validation is never augmented; that is Ultralytics' own behaviour and REQ-112 only requires that we do not defeat it. """ import json from typing import Optional from backend import db # name -> (minimum, maximum). Bounds are Ultralytics' own accepted ranges. FIELDS = { "fliplr": (0.0, 1.0), "flipud": (0.0, 1.0), "degrees": (0.0, 180.0), "translate": (0.0, 1.0), "scale": (0.0, 1.0), "hsv_h": (0.0, 1.0), "hsv_s": (0.0, 1.0), "hsv_v": (0.0, 1.0), "mosaic": (0.0, 1.0), } OFF = {name: 0.0 for name in FIELDS} LIGHT = {"fliplr": 0.5, "flipud": 0.0, "degrees": 0.0, "translate": 0.05, "scale": 0.2, "hsv_h": 0.010, "hsv_s": 0.4, "hsv_v": 0.3, "mosaic": 0.0} # Ultralytics' defaults, spelled out. MEDIUM = {"fliplr": 0.5, "flipud": 0.0, "degrees": 0.0, "translate": 0.1, "scale": 0.5, "hsv_h": 0.015, "hsv_s": 0.7, "hsv_v": 0.4, "mosaic": 1.0} AGGRESSIVE = {"fliplr": 0.5, "flipud": 0.1, "degrees": 10.0, "translate": 0.2, "scale": 0.9, "hsv_h": 0.020, "hsv_s": 0.9, "hsv_v": 0.5, "mosaic": 1.0} PRESETS = {"off": OFF, "light": LIGHT, "medium": MEDIUM, "aggressive": AGGRESSIVE} class AugmentError(Exception): pass def normalise(incoming: Optional[dict]) -> dict: """Fill in missing keys from MEDIUM and reject out-of-range values.""" settings = dict(MEDIUM) for name, value in (incoming or {}).items(): if name not in FIELDS: raise AugmentError(f"'{name}' is not an augmentation setting") if not isinstance(value, (int, float)) or isinstance(value, bool): raise AugmentError(f"'{name}' must be a number") low, high = FIELDS[name] if not low <= value <= high: raise AugmentError(f"'{name}' must be between {low} and {high}") settings[name] = float(value) return settings def preset_name(settings: dict) -> str: """Which preset these settings match, or 'custom'.""" for name, preset in PRESETS.items(): if all(abs(settings[field] - preset[field]) < 1e-9 for field in FIELDS): return name return "custom" def get(project_id: int) -> dict: with db.cursor() as cur: cur.execute("SELECT augment FROM projects WHERE id = ?", (project_id,)) row = cur.fetchone() if row is None: raise AugmentError("No such project") stored = json.loads(row[0]) if row[0] else None settings = normalise(stored) return {"settings": settings, "preset": preset_name(settings)} def save(project_id: int, incoming: dict) -> dict: settings = normalise(incoming) with db.cursor() as cur: cur.execute("UPDATE projects SET augment = ? WHERE id = ?", (json.dumps(settings), project_id)) if cur.rowcount == 0: raise AugmentError("No such project") return {"settings": settings, "preset": preset_name(settings)}