Files
reTraining/backend/augment.py
T
asus 5c7c122105 feat: add counting bench, triage, and dataset modules
This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
2026-08-14 16:28:52 +07:00

91 lines
3.2 KiB
Python

"""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)}