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(), )