219 lines
8.3 KiB
Python
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(),
|
|
)
|