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.
91 lines
3.2 KiB
Python
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)}
|