train_utility/src/yolo_tui/config.py

219 lines
8.3 KiB
Python

from __future__ import annotations
from dataclasses import dataclass, field
from typing import Literal, Any
YoloTask = Literal["detect", "segment", "classify", "pose", "obb"]
AutoAugmentPolicy = Literal["randaugment", "autoaugment", "augmix"]
CopyPasteMode = Literal["flip", "mixup"]
SUPPORTED_TASKS: tuple[YoloTask, ...] = (
"detect",
"segment",
"classify",
"pose",
"obb",
)
SUPPORTED_AUTO_AUGMENT_POLICIES: tuple[AutoAugmentPolicy, ...] = (
"randaugment",
"autoaugment",
"augmix",
)
SUPPORTED_COPY_PASTE_MODES: tuple[CopyPasteMode, ...] = ("flip", "mixup")
@dataclass(frozen=True, slots=True)
class AugmentationConfig:
"""Training augmentation overrides accepted by Ultralytics."""
enabled: bool = True
hsv_h: float = 0.015
hsv_s: float = 0.7
hsv_v: float = 0.4
degrees: float = 0.0
translate: float = 0.1
scale: float = 0.5
shear: float = 0.0
perspective: float = 0.0
flipud: float = 0.0
fliplr: float = 0.5
bgr: float = 0.0
mosaic: float = 1.0
mixup: float = 0.0
cutmix: float = 0.0
copy_paste: float = 0.0
copy_paste_mode: CopyPasteMode = "flip"
auto_augment: AutoAugmentPolicy = "randaugment"
erasing: float = 0.4
close_mosaic: int = 10
def validate(self) -> None:
if not self.enabled:
return
fractions = {
"HSV hue": self.hsv_h,
"HSV saturation": self.hsv_s,
"HSV brightness": self.hsv_v,
"Translate": self.translate,
"Scale": self.scale,
"Perspective": self.perspective,
"Flip up/down": self.flipud,
"Flip left/right": self.fliplr,
"BGR": self.bgr,
"Mosaic": self.mosaic,
"MixUp": self.mixup,
"CutMix": self.cutmix,
"Copy-paste": self.copy_paste,
"Erasing": self.erasing,
}
for label, value in fractions.items():
if not 0.0 <= value <= 1.0:
raise ValueError(f"Параметр «{label}» должен быть от 0 до 1.")
if self.degrees < 0:
raise ValueError("Угол поворота не может быть отрицательным.")
if self.shear < 0:
raise ValueError("Угол сдвига не может быть отрицательным.")
if self.close_mosaic < 0:
raise ValueError("Close mosaic не может быть отрицательным.")
if self.copy_paste_mode not in SUPPORTED_COPY_PASTE_MODES:
raise ValueError(f"Неизвестный режим copy-paste: {self.copy_paste_mode}.")
if self.auto_augment not in SUPPORTED_AUTO_AUGMENT_POLICIES:
raise ValueError(f"Неизвестная политика AutoAugment: {self.auto_augment}.")
def train_kwargs(self) -> dict[str, str | int | float]:
if not self.enabled:
return {}
return {
"hsv_h": self.hsv_h,
"hsv_s": self.hsv_s,
"hsv_v": self.hsv_v,
"degrees": self.degrees,
"translate": self.translate,
"scale": self.scale,
"shear": self.shear,
"perspective": self.perspective,
"flipud": self.flipud,
"fliplr": self.fliplr,
"bgr": self.bgr,
"mosaic": self.mosaic,
"mixup": self.mixup,
"cutmix": self.cutmix,
"copy_paste": self.copy_paste,
"copy_paste_mode": self.copy_paste_mode,
"auto_augment": self.auto_augment,
"erasing": self.erasing,
"close_mosaic": self.close_mosaic,
}
@dataclass(frozen=True, slots=True)
class MlflowConfig:
enabled: bool = True
tracking_uri: str = "sqlite:///mlflow.db"
experiment_name: str = "yolo-tui"
run_name: str = ""
def validate(self) -> None:
if self.enabled and not self.tracking_uri.strip():
raise ValueError("Укажите URI хранилища MLflow.")
if self.enabled and not self.experiment_name.strip():
raise ValueError("Укажите название эксперимента MLflow.")
@dataclass(frozen=True, slots=True)
class DatasetSplitConfig:
enabled: bool = False
train_ratio: float = 0.8
classes_path: str = ""
def validate(self) -> None:
if self.enabled:
if not 0.1 <= self.train_ratio <= 0.95:
raise ValueError("Доля обучающей выборки (Train) должна быть от 0.1 до 0.95.")
@dataclass(frozen=True, slots=True)
class TrainingConfig:
dataset: str
model: str
task: YoloTask = "detect"
epochs: int = 100
image_size: int = 640
batch_size: int = 16
device: str = ""
workers: int = 8
patience: int = 100
project: str = "runs/train"
run_name: str = ""
augmentation: AugmentationConfig = field(default_factory=AugmentationConfig)
mlflow: MlflowConfig = field(default_factory=MlflowConfig)
split: DatasetSplitConfig = field(default_factory=DatasetSplitConfig)
def validate(self) -> None:
if not self.dataset.strip():
raise ValueError("Укажите путь или имя датасета.")
if not self.model.strip():
raise ValueError("Укажите путь или имя модели.")
if self.task not in SUPPORTED_TASKS:
raise ValueError(f"Неизвестный тип задачи: {self.task}.")
if self.epochs < 1:
raise ValueError("Количество эпох должно быть не меньше 1.")
if self.image_size < 32:
raise ValueError("Размер изображения должен быть не меньше 32.")
if self.batch_size == 0 or self.batch_size < -1:
raise ValueError("Batch должен быть положительным числом или -1 для автоподбора.")
if self.workers < 0:
raise ValueError("Количество workers не может быть отрицательным.")
if self.patience < 0:
raise ValueError("Patience не может быть отрицательным.")
self.augmentation.validate()
self.mlflow.validate()
self.split.validate()
def train_kwargs(self) -> dict[str, str | int | float | bool]:
"""Convert the form values to arguments accepted by YOLO.train()."""
values: dict[str, str | int | float | bool] = {
"data": self.dataset.strip(),
"epochs": self.epochs,
"imgsz": self.image_size,
"batch": self.batch_size,
"workers": self.workers,
"patience": self.patience,
"project": self.project.strip() or "runs/train",
# Ultralytics' tqdm output would otherwise paint over Textual's screen.
# Epoch metrics are sent to the in-app log by callbacks instead.
"verbose": False,
}
if self.device.strip():
values["device"] = self.device.strip()
if self.run_name.strip():
values["name"] = self.run_name.strip()
values.update(self.augmentation.train_kwargs())
return values
def to_dict(self) -> dict[str, Any]:
import dataclasses
return dataclasses.asdict(self)
@classmethod
def from_dict(cls, data: dict[str, Any]) -> TrainingConfig:
aug_data = data.get("augmentation", {})
mlflow_data = data.get("mlflow", {})
split_data = data.get("split", {})
return cls(
dataset=data["dataset"],
model=data["model"],
task=data.get("task", "detect"),
epochs=data.get("epochs", 100),
image_size=data.get("image_size", 640),
batch_size=data.get("batch_size", 16),
device=data.get("device", ""),
workers=data.get("workers", 8),
patience=data.get("patience", 100),
project=data.get("project", "runs/train"),
run_name=data.get("run_name", ""),
augmentation=AugmentationConfig(**aug_data) if aug_data else AugmentationConfig(),
mlflow=MlflowConfig(**mlflow_data) if mlflow_data else MlflowConfig(),
split=DatasetSplitConfig(**split_data) if split_data else DatasetSplitConfig(),
)