Добавлены тесты, исправлены ошибки запуска обучения
This commit is contained in:
parent
1dcc6a5249
commit
06deeb25d1
9 changed files with 580 additions and 36 deletions
2
.gitignore
vendored
2
.gitignore
vendored
|
|
@ -10,3 +10,5 @@ mlruns/
|
|||
mlflow.db
|
||||
mlflow.db-shm
|
||||
mlflow.db-wal
|
||||
passport_obb_up/
|
||||
yolo11n.pt
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from textual.widgets import (
|
|||
|
||||
from .config import (
|
||||
AugmentationConfig,
|
||||
DatasetSplitConfig,
|
||||
MlflowConfig,
|
||||
SUPPORTED_AUTO_AUGMENT_POLICIES,
|
||||
SUPPORTED_COPY_PASTE_MODES,
|
||||
|
|
@ -194,7 +195,7 @@ class YoloTrainApp(App[None]):
|
|||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.runner = TrainingRunner()
|
||||
self._running = False
|
||||
self._training_running = False
|
||||
|
||||
def compose(self) -> ComposeResult:
|
||||
yield Header(show_clock=True)
|
||||
|
|
@ -215,10 +216,18 @@ class YoloTrainApp(App[None]):
|
|||
Input(value="yolo11n.pt", placeholder="/models/best.pt", id="model"),
|
||||
)
|
||||
yield Field(
|
||||
"Датасет — путь к YAML/каталогу или имя",
|
||||
Input(value="coco8.yaml", placeholder="/data/dataset.yaml", id="dataset"),
|
||||
"Датасет — путь к папке датасета",
|
||||
Input(value="coco8.yaml", placeholder="/path/to/dataset", id="dataset"),
|
||||
)
|
||||
|
||||
yield Static("Разделение датасета (Train/Val)", classes="section-title")
|
||||
with Horizontal(classes="toggle-row"):
|
||||
yield Label("Разделить автоматически на train/val")
|
||||
yield Switch(value=False, id="split-enabled")
|
||||
with Horizontal(classes="row split-field"):
|
||||
yield Field("Доля train (0.1…0.95)", Input(value="0.8", id="split-ratio"))
|
||||
yield Field("Путь к classes.txt / YAML (необязательно)", Input(placeholder="Автопоиск", id="split-classes"))
|
||||
|
||||
yield Static("Параметры обучения", classes="section-title")
|
||||
with Horizontal(classes="row"):
|
||||
yield Field("Эпохи", Input(value="100", type="integer", id="epochs"))
|
||||
|
|
@ -309,8 +318,19 @@ class YoloTrainApp(App[None]):
|
|||
yield Footer()
|
||||
|
||||
def on_mount(self) -> None:
|
||||
import os
|
||||
os.environ["MPLBACKEND"] = "Agg"
|
||||
import ultralytics
|
||||
|
||||
self.query_one("#progress", ProgressBar).update(progress=0)
|
||||
self._write_log("[dim]Интерфейс готов. Обучение еще не запускалось.[/dim]")
|
||||
for control in self.query(".split-field Input"):
|
||||
control.disabled = True
|
||||
|
||||
@on(Switch.Changed, "#split-enabled")
|
||||
def toggle_split(self, event: Switch.Changed) -> None:
|
||||
for control in self.query(".split-field Input"):
|
||||
control.disabled = not event.value
|
||||
|
||||
@on(Switch.Changed, "#mlflow-enabled")
|
||||
def toggle_mlflow(self, event: Switch.Changed) -> None:
|
||||
|
|
@ -333,7 +353,7 @@ class YoloTrainApp(App[None]):
|
|||
self.action_stop_training()
|
||||
|
||||
def action_start_training(self) -> None:
|
||||
if self._running:
|
||||
if self._training_running:
|
||||
self.notify("Обучение уже выполняется.", severity="warning")
|
||||
return
|
||||
try:
|
||||
|
|
@ -363,7 +383,7 @@ class YoloTrainApp(App[None]):
|
|||
self._train_in_background(config)
|
||||
|
||||
def action_stop_training(self) -> None:
|
||||
if not self._running:
|
||||
if not self._training_running:
|
||||
return
|
||||
self.runner.request_stop()
|
||||
self.query_one("#status-title", Static).update("ОСТАНОВКА")
|
||||
|
|
@ -373,15 +393,86 @@ class YoloTrainApp(App[None]):
|
|||
|
||||
@work(thread=True, exclusive=True, group="yolo-training")
|
||||
def _train_in_background(self, config: TrainingConfig) -> None:
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from rich.markup import escape
|
||||
|
||||
temp_config_path = None
|
||||
process = None
|
||||
try:
|
||||
output_dir = self.runner.train(
|
||||
config,
|
||||
lambda event: self.app.call_from_thread(self._handle_training_event, event),
|
||||
with tempfile.NamedTemporaryFile("w", suffix=".json", delete=False, encoding="utf-8") as f:
|
||||
json.dump(config.to_dict(), f)
|
||||
temp_config_path = f.name
|
||||
|
||||
cmd = [sys.executable, "-m", "yolo_tui.subprocess_runner", temp_config_path]
|
||||
process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
bufsize=1,
|
||||
)
|
||||
except Exception as exc: # errors must be surfaced in the TUI, not hidden in a worker
|
||||
self.runner.set_subprocess(process)
|
||||
|
||||
output_dir = None
|
||||
|
||||
while True:
|
||||
line = process.stdout.readline()
|
||||
if not line:
|
||||
break
|
||||
line_str = line.strip()
|
||||
if not line_str:
|
||||
continue
|
||||
|
||||
if line_str.startswith("__YOLO_TUI_EVENT__:"):
|
||||
try:
|
||||
event_data = json.loads(line_str[len("__YOLO_TUI_EVENT__:"):])
|
||||
event = TrainingEvent(
|
||||
kind=event_data["kind"],
|
||||
message=event_data["message"],
|
||||
epoch=event_data["epoch"],
|
||||
total_epochs=event_data["total_epochs"],
|
||||
)
|
||||
self.app.call_from_thread(self._handle_training_event, event)
|
||||
except Exception:
|
||||
pass
|
||||
elif line_str.startswith("__YOLO_TUI_RESULT__:"):
|
||||
output_dir = line_str[len("__YOLO_TUI_RESULT__:"):]
|
||||
else:
|
||||
self.app.call_from_thread(self._write_log, escape(line_str))
|
||||
|
||||
process.wait()
|
||||
rc = process.returncode
|
||||
|
||||
if rc == 0:
|
||||
self.app.call_from_thread(self._training_finished, output_dir)
|
||||
else:
|
||||
if self.runner.stop_requested:
|
||||
self.app.call_from_thread(self._training_finished, None)
|
||||
else:
|
||||
self.app.call_from_thread(
|
||||
self._training_failed,
|
||||
Exception("Процесс обучения завершился с ошибкой. Проверьте логи выше."),
|
||||
)
|
||||
|
||||
except BaseException as exc:
|
||||
import traceback
|
||||
try:
|
||||
with open("/Users/vadim/.gemini/antigravity/brain/ab216120-71db-4a55-a2bf-2eff6a9caf26/error.log", "w", encoding="utf-8") as f:
|
||||
traceback.print_exc(file=f)
|
||||
except Exception:
|
||||
pass
|
||||
self.app.call_from_thread(self._training_failed, exc)
|
||||
else:
|
||||
self.app.call_from_thread(self._training_finished, output_dir)
|
||||
finally:
|
||||
if temp_config_path and os.path.exists(temp_config_path):
|
||||
try:
|
||||
os.unlink(temp_config_path)
|
||||
except Exception:
|
||||
pass
|
||||
self.runner.clear_subprocess()
|
||||
|
||||
def _handle_training_event(self, event: TrainingEvent) -> None:
|
||||
styles = {
|
||||
|
|
@ -429,7 +520,7 @@ class YoloTrainApp(App[None]):
|
|||
self.notify(str(error), title="Обучение не запущено", severity="error", timeout=10)
|
||||
|
||||
def _set_running(self, running: bool) -> None:
|
||||
self._running = running
|
||||
self._training_running = running
|
||||
self.query_one("#start-button", Button).disabled = running
|
||||
self.query_one("#stop-button", Button).disabled = not running
|
||||
|
||||
|
|
@ -437,20 +528,11 @@ class YoloTrainApp(App[None]):
|
|||
task = self.query_one("#task", Select).value
|
||||
if task not in SUPPORTED_TASKS:
|
||||
raise ValueError("Выберите тип задачи YOLO.")
|
||||
return TrainingConfig(
|
||||
dataset=self._input("dataset"),
|
||||
model=self._input("model"),
|
||||
task=task,
|
||||
epochs=self._integer("epochs", "Эпохи"),
|
||||
image_size=self._integer("image-size", "Размер изображения"),
|
||||
batch_size=self._integer("batch-size", "Batch"),
|
||||
device=self._input("device"),
|
||||
workers=self._integer("workers", "Workers"),
|
||||
patience=self._integer("patience", "Patience"),
|
||||
project=self._input("project"),
|
||||
run_name=self._input("run-name"),
|
||||
augmentation=AugmentationConfig(
|
||||
enabled=self.query_one("#augmentation-enabled", Switch).value,
|
||||
|
||||
augmentation_enabled = self.query_one("#augmentation-enabled", Switch).value
|
||||
if augmentation_enabled:
|
||||
augmentation = AugmentationConfig(
|
||||
enabled=True,
|
||||
hsv_h=self._float("hsv-h", "HSV hue"),
|
||||
hsv_s=self._float("hsv-s", "HSV saturation"),
|
||||
hsv_v=self._float("hsv-v", "HSV brightness"),
|
||||
|
|
@ -478,13 +560,46 @@ class YoloTrainApp(App[None]):
|
|||
),
|
||||
erasing=self._float("erasing", "Erasing"),
|
||||
close_mosaic=self._integer("close-mosaic", "Close mosaic"),
|
||||
),
|
||||
mlflow=MlflowConfig(
|
||||
enabled=self.query_one("#mlflow-enabled", Switch).value,
|
||||
)
|
||||
else:
|
||||
augmentation = AugmentationConfig(enabled=False)
|
||||
|
||||
mlflow_enabled = self.query_one("#mlflow-enabled", Switch).value
|
||||
if mlflow_enabled:
|
||||
mlflow = MlflowConfig(
|
||||
enabled=True,
|
||||
tracking_uri=self._input("tracking-uri"),
|
||||
experiment_name=self._input("experiment-name"),
|
||||
run_name=self._input("mlflow-run-name"),
|
||||
),
|
||||
)
|
||||
else:
|
||||
mlflow = MlflowConfig(enabled=False)
|
||||
|
||||
split_enabled = self.query_one("#split-enabled", Switch).value
|
||||
if split_enabled:
|
||||
split_config = DatasetSplitConfig(
|
||||
enabled=True,
|
||||
train_ratio=self._float("split-ratio", "Доля train"),
|
||||
classes_path=self._input("split-classes"),
|
||||
)
|
||||
else:
|
||||
split_config = DatasetSplitConfig(enabled=False)
|
||||
|
||||
return TrainingConfig(
|
||||
dataset=self._input("dataset"),
|
||||
model=self._input("model"),
|
||||
task=task,
|
||||
epochs=self._integer("epochs", "Эпохи"),
|
||||
image_size=self._integer("image-size", "Размер изображения"),
|
||||
batch_size=self._integer("batch-size", "Batch"),
|
||||
device=self._input("device"),
|
||||
workers=self._integer("workers", "Workers"),
|
||||
patience=self._integer("patience", "Patience"),
|
||||
project=self._input("project"),
|
||||
run_name=self._input("run-name"),
|
||||
augmentation=augmentation,
|
||||
mlflow=mlflow,
|
||||
split=split_config,
|
||||
)
|
||||
|
||||
def _input(self, widget_id: str) -> str:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Literal
|
||||
from typing import Literal, Any
|
||||
|
||||
|
||||
YoloTask = Literal["detect", "segment", "classify", "pose", "obb"]
|
||||
|
|
@ -121,6 +121,18 @@ class MlflowConfig:
|
|||
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
|
||||
|
|
@ -136,6 +148,7 @@ class TrainingConfig:
|
|||
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():
|
||||
|
|
@ -156,6 +169,7 @@ class TrainingConfig:
|
|||
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()."""
|
||||
|
|
@ -177,3 +191,29 @@ class TrainingConfig:
|
|||
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(),
|
||||
)
|
||||
|
|
|
|||
139
src/yolo_tui/dataset_splitter.py
Normal file
139
src/yolo_tui/dataset_splitter.py
Normal file
|
|
@ -0,0 +1,139 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
from pathlib import Path
|
||||
import yaml
|
||||
|
||||
|
||||
def _parse_classes_file(path: Path) -> dict[int, str]:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
lines = [line.strip() for line in f if line.strip()]
|
||||
return {i: name for i, name in enumerate(lines)}
|
||||
|
||||
|
||||
def read_classes(dataset_dir: Path, custom_classes_path: str) -> dict[int, str]:
|
||||
# 1. Custom path specified by user
|
||||
if custom_classes_path.strip():
|
||||
path = Path(custom_classes_path.strip())
|
||||
if path.exists():
|
||||
return _parse_classes_file(path)
|
||||
|
||||
# 2. classes.txt in dataset_dir root
|
||||
path = dataset_dir / "classes.txt"
|
||||
if path.exists():
|
||||
return _parse_classes_file(path)
|
||||
|
||||
# 3. classes.txt inside dataset_dir/labels
|
||||
path = dataset_dir / "labels" / "classes.txt"
|
||||
if path.exists():
|
||||
return _parse_classes_file(path)
|
||||
|
||||
# 4. Any .yaml file in dataset_dir (excluding split/dataset.yaml)
|
||||
yaml_files = list(dataset_dir.glob("*.yaml"))
|
||||
yaml_files = [f for f in yaml_files if f.name != "dataset.yaml"]
|
||||
if yaml_files:
|
||||
try:
|
||||
with open(yaml_files[0], "r", encoding="utf-8") as f:
|
||||
data = yaml.safe_load(f)
|
||||
if isinstance(data, dict) and "names" in data:
|
||||
names = data["names"]
|
||||
if isinstance(names, dict):
|
||||
return {int(k): str(v) for k, v in names.items()}
|
||||
elif isinstance(names, list):
|
||||
return {i: str(v) for i, v in enumerate(names)}
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 5. Fallback: Scan label files to determine number of classes and use class_i names
|
||||
class_ids = set()
|
||||
labels_dir = dataset_dir / "labels"
|
||||
if labels_dir.exists():
|
||||
for txt_file in labels_dir.glob("*.txt"):
|
||||
if txt_file.name == "classes.txt":
|
||||
continue
|
||||
try:
|
||||
with open(txt_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
parts = line.strip().split()
|
||||
if parts:
|
||||
class_ids.add(int(parts[0]))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if class_ids:
|
||||
max_id = max(class_ids)
|
||||
return {i: f"class_{i}" for i in range(max_id + 1)}
|
||||
|
||||
raise ValueError(
|
||||
"Не удалось найти список классов. Пожалуйста, создайте файл classes.txt "
|
||||
"в корневой папке датасета или укажите путь к нему."
|
||||
)
|
||||
|
||||
|
||||
def split_dataset(
|
||||
dataset_dir: str, train_ratio: float, classes_path: str
|
||||
) -> tuple[int, int, str]:
|
||||
base_dir = Path(dataset_dir.strip()).absolute()
|
||||
images_dir = base_dir / "images"
|
||||
labels_dir = base_dir / "labels"
|
||||
|
||||
if not base_dir.exists():
|
||||
raise ValueError(f"Каталог датасета '{base_dir}' не существует.")
|
||||
if not images_dir.exists():
|
||||
raise ValueError(f"Папка с изображениями '{images_dir}' не найдена.")
|
||||
if not labels_dir.exists():
|
||||
raise ValueError(f"Папка с разметкой '{labels_dir}' не найдена.")
|
||||
|
||||
# Find images
|
||||
valid_extensions = {".jpg", ".jpeg", ".png", ".bmp", ".webp", ".tif", ".tiff"}
|
||||
image_files = [
|
||||
f for f in images_dir.iterdir()
|
||||
if f.is_file() and f.suffix.lower() in valid_extensions
|
||||
]
|
||||
|
||||
if not image_files:
|
||||
raise ValueError(f"В папке '{images_dir}' не найдено изображений.")
|
||||
|
||||
# Shuffle deterministically using a fixed seed
|
||||
rng = random.Random(42)
|
||||
rng.shuffle(image_files)
|
||||
|
||||
split_idx = int(len(image_files) * train_ratio)
|
||||
if split_idx == 0:
|
||||
split_idx = 1
|
||||
if split_idx >= len(image_files):
|
||||
split_idx = len(image_files) - 1
|
||||
|
||||
train_images = image_files[:split_idx]
|
||||
val_images = image_files[split_idx:]
|
||||
|
||||
split_dir = base_dir / "split"
|
||||
split_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
train_txt_path = split_dir / "train.txt"
|
||||
val_txt_path = split_dir / "val.txt"
|
||||
|
||||
with open(train_txt_path, "w", encoding="utf-8") as f:
|
||||
for img in train_images:
|
||||
f.write(f"{img}\n")
|
||||
|
||||
with open(val_txt_path, "w", encoding="utf-8") as f:
|
||||
for img in val_images:
|
||||
f.write(f"{img}\n")
|
||||
|
||||
# Load class mapping
|
||||
classes = read_classes(base_dir, classes_path)
|
||||
|
||||
# Write data.yaml equivalent
|
||||
dataset_yaml_path = split_dir / "dataset.yaml"
|
||||
dataset_data = {
|
||||
"path": str(base_dir),
|
||||
"train": f"split/{train_txt_path.name}",
|
||||
"val": f"split/{val_txt_path.name}",
|
||||
"names": classes,
|
||||
}
|
||||
|
||||
with open(dataset_yaml_path, "w", encoding="utf-8") as f:
|
||||
yaml.safe_dump(dataset_data, f, allow_unicode=True, sort_keys=False)
|
||||
|
||||
return len(train_images), len(val_images), str(dataset_yaml_path)
|
||||
49
src/yolo_tui/subprocess_runner.py
Normal file
49
src/yolo_tui/subprocess_runner.py
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# Force headless Matplotlib to avoid any thread/process GUI issues
|
||||
os.environ["MPLBACKEND"] = "Agg"
|
||||
|
||||
from yolo_tui.config import TrainingConfig
|
||||
from yolo_tui.trainer import TrainingEvent, TrainingRunner
|
||||
|
||||
def main():
|
||||
if len(sys.argv) < 2:
|
||||
print("Usage: python -m yolo_tui.subprocess_runner <config_json_path>", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
config_path = sys.argv[1]
|
||||
try:
|
||||
with open(config_path, "r", encoding="utf-8") as f:
|
||||
config_dict = json.load(f)
|
||||
config = TrainingConfig.from_dict(config_dict)
|
||||
except Exception as e:
|
||||
print(f"Error loading config: {e}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
runner = TrainingRunner()
|
||||
|
||||
def handle_event(event: TrainingEvent) -> None:
|
||||
event_dict = {
|
||||
"kind": event.kind,
|
||||
"message": event.message,
|
||||
"epoch": event.epoch,
|
||||
"total_epochs": event.total_epochs,
|
||||
}
|
||||
# Print structured JSON event so the parent process can parse it
|
||||
print(f"__YOLO_TUI_EVENT__:{json.dumps(event_dict)}", flush=True)
|
||||
|
||||
try:
|
||||
output_dir = runner.train(config, handle_event)
|
||||
if output_dir:
|
||||
print(f"__YOLO_TUI_RESULT__:{output_dir}", flush=True)
|
||||
sys.exit(0)
|
||||
except BaseException as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -55,10 +55,24 @@ class TrainingRunner:
|
|||
self._model: Any | None = None
|
||||
self._state_lock = Lock()
|
||||
self._stop_requested = Event()
|
||||
self._subprocess: Any | None = None
|
||||
|
||||
def set_subprocess(self, process: Any) -> None:
|
||||
with self._state_lock:
|
||||
self._subprocess = process
|
||||
|
||||
def clear_subprocess(self) -> None:
|
||||
with self._state_lock:
|
||||
self._subprocess = None
|
||||
|
||||
def request_stop(self) -> None:
|
||||
self._stop_requested.set()
|
||||
with self._state_lock:
|
||||
if self._subprocess is not None:
|
||||
try:
|
||||
self._subprocess.terminate()
|
||||
except Exception:
|
||||
pass
|
||||
trainer = getattr(self._model, "trainer", None)
|
||||
if trainer is not None:
|
||||
trainer.stop = True
|
||||
|
|
@ -71,6 +85,28 @@ class TrainingRunner:
|
|||
config.validate()
|
||||
self._stop_requested.clear()
|
||||
|
||||
train_args = config.train_kwargs()
|
||||
|
||||
if config.split.enabled:
|
||||
on_event(TrainingEvent("info", "Разделение датасета на train/val…"))
|
||||
try:
|
||||
from .dataset_splitter import split_dataset
|
||||
train_count, val_count, yaml_path = split_dataset(
|
||||
dataset_dir=config.dataset,
|
||||
train_ratio=config.split.train_ratio,
|
||||
classes_path=config.split.classes_path,
|
||||
)
|
||||
on_event(
|
||||
TrainingEvent(
|
||||
"info",
|
||||
f"Разделение завершено: train={train_count}, val={val_count}",
|
||||
)
|
||||
)
|
||||
train_args["data"] = yaml_path
|
||||
except Exception as exc:
|
||||
on_event(TrainingEvent("warning", f"Ошибка разделения датасета: {exc}"))
|
||||
raise exc
|
||||
|
||||
on_event(TrainingEvent("info", "Загрузка Ultralytics и подготовка модели…"))
|
||||
from ultralytics import YOLO, settings
|
||||
|
||||
|
|
@ -86,7 +122,7 @@ class TrainingRunner:
|
|||
model.add_callback("on_train_end", self._on_train_end(on_event))
|
||||
|
||||
try:
|
||||
model.train(**config.train_kwargs())
|
||||
model.train(**train_args)
|
||||
trainer = getattr(model, "trainer", None)
|
||||
save_dir = getattr(trainer, "save_dir", None)
|
||||
return Path(save_dir) if save_dir else None
|
||||
|
|
@ -130,9 +166,11 @@ class TrainingRunner:
|
|||
@staticmethod
|
||||
def _metrics_summary(metrics: dict[str, Any]) -> str:
|
||||
result: list[str] = []
|
||||
for key, value in list(metrics.items())[:3]:
|
||||
for key, value in metrics.items():
|
||||
try:
|
||||
result.append(f"{key.split('/')[-1]}={float(value):.4g}")
|
||||
if len(result) == 3:
|
||||
break
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
return " · ".join(result)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import asyncio
|
|||
from textual.widgets import Button, Input, Select, Switch
|
||||
|
||||
from yolo_tui.app import YoloTrainApp
|
||||
from yolo_tui.config import AugmentationConfig, DatasetSplitConfig
|
||||
|
||||
|
||||
def test_app_mounts_with_expected_defaults() -> None:
|
||||
|
|
@ -50,3 +51,50 @@ def test_mlflow_fields_follow_switch() -> None:
|
|||
assert app.query_one("#experiment-name", Input).disabled is True
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
||||
|
||||
def test_read_config_ignores_disabled_augmentation() -> None:
|
||||
async def exercise() -> None:
|
||||
app = YoloTrainApp()
|
||||
async with app.run_test(size=(140, 45)) as pilot:
|
||||
app.query_one("#augmentation-enabled", Switch).value = False
|
||||
app.query_one("#hsv-h", Input).value = "not a number"
|
||||
await pilot.pause()
|
||||
|
||||
config = app._read_config()
|
||||
assert config.augmentation.enabled is False
|
||||
# When disabled, config.augmentation uses defaults, doesn't parse from UI input
|
||||
assert config.augmentation.hsv_h == AugmentationConfig(enabled=False).hsv_h
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
||||
|
||||
def test_split_fields_follow_switch() -> None:
|
||||
async def exercise() -> None:
|
||||
app = YoloTrainApp()
|
||||
async with app.run_test(size=(140, 45)) as pilot:
|
||||
assert app.query_one("#split-ratio", Input).disabled is True
|
||||
assert app.query_one("#split-classes", Input).disabled is True
|
||||
|
||||
app.query_one("#split-enabled", Switch).value = True
|
||||
await pilot.pause()
|
||||
|
||||
assert app.query_one("#split-ratio", Input).disabled is False
|
||||
assert app.query_one("#split-classes", Input).disabled is False
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
||||
|
||||
def test_read_config_ignores_disabled_split() -> None:
|
||||
async def exercise() -> None:
|
||||
app = YoloTrainApp()
|
||||
async with app.run_test(size=(140, 45)) as pilot:
|
||||
app.query_one("#split-enabled", Switch).value = False
|
||||
app.query_one("#split-ratio", Input).value = "not a float"
|
||||
await pilot.pause()
|
||||
|
||||
config = app._read_config()
|
||||
assert config.split.enabled is False
|
||||
assert config.split.train_ratio == DatasetSplitConfig(enabled=False).train_ratio
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import os
|
|||
|
||||
import pytest
|
||||
|
||||
from yolo_tui.config import AugmentationConfig, MlflowConfig, TrainingConfig
|
||||
from yolo_tui.config import AugmentationConfig, DatasetSplitConfig, MlflowConfig, TrainingConfig
|
||||
from yolo_tui.trainer import TrainingRunner, mlflow_environment
|
||||
|
||||
|
||||
|
|
@ -79,7 +79,20 @@ def test_mlflow_environment_is_restored(monkeypatch: pytest.MonkeyPatch) -> None
|
|||
|
||||
def test_metrics_summary_skips_non_numeric_values() -> None:
|
||||
summary = TrainingRunner._metrics_summary(
|
||||
{"metrics/mAP50": 0.81234, "label": "invalid", "val/loss": 0.12345}
|
||||
{
|
||||
"metrics/mAP50": 0.81234,
|
||||
"label": "invalid",
|
||||
"val/loss": 0.12345,
|
||||
"metrics/precision": 0.95678,
|
||||
"another": 12.34,
|
||||
}
|
||||
)
|
||||
|
||||
assert summary == "mAP50=0.8123 · loss=0.1235"
|
||||
assert summary == "mAP50=0.8123 · loss=0.1235 · precision=0.9568"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ratio", [0.05, 0.98])
|
||||
def test_dataset_split_ratio_is_validated(ratio: float) -> None:
|
||||
config = DatasetSplitConfig(enabled=True, train_ratio=ratio)
|
||||
with pytest.raises(ValueError, match="Доля обучающей выборки"):
|
||||
config.validate()
|
||||
|
|
|
|||
100
tests/test_splitter.py
Normal file
100
tests/test_splitter.py
Normal file
|
|
@ -0,0 +1,100 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
from yolo_tui.dataset_splitter import read_classes, split_dataset
|
||||
|
||||
|
||||
def test_read_classes_custom_path(tmp_path: Path) -> None:
|
||||
custom_file = tmp_path / "custom_classes.txt"
|
||||
custom_file.write_text("classA\nclassB\n", encoding="utf-8")
|
||||
|
||||
classes = read_classes(tmp_path, str(custom_file))
|
||||
assert classes == {0: "classA", 1: "classB"}
|
||||
|
||||
|
||||
def test_read_classes_root_classes_txt(tmp_path: Path) -> None:
|
||||
classes_file = tmp_path / "classes.txt"
|
||||
classes_file.write_text("class0\nclass1\nclass2\n", encoding="utf-8")
|
||||
|
||||
classes = read_classes(tmp_path, "")
|
||||
assert classes == {0: "class0", 1: "class1", 2: "class2"}
|
||||
|
||||
|
||||
def test_read_classes_labels_classes_txt(tmp_path: Path) -> None:
|
||||
labels_dir = tmp_path / "labels"
|
||||
labels_dir.mkdir()
|
||||
classes_file = labels_dir / "classes.txt"
|
||||
classes_file.write_text("lbl0\nlbl1\n", encoding="utf-8")
|
||||
|
||||
classes = read_classes(tmp_path, "")
|
||||
assert classes == {0: "lbl0", 1: "lbl1"}
|
||||
|
||||
|
||||
def test_read_classes_from_yaml(tmp_path: Path) -> None:
|
||||
yaml_file = tmp_path / "dataset_config.yaml"
|
||||
yaml_data = {"names": ["yaml_cls0", "yaml_cls1"]}
|
||||
yaml_file.write_text(yaml.dump(yaml_data), encoding="utf-8")
|
||||
|
||||
classes = read_classes(tmp_path, "")
|
||||
assert classes == {0: "yaml_cls0", 1: "yaml_cls1"}
|
||||
|
||||
|
||||
def test_read_classes_fallback_scanning_labels(tmp_path: Path) -> None:
|
||||
labels_dir = tmp_path / "labels"
|
||||
labels_dir.mkdir()
|
||||
# Write some mock label files containing class IDs: 0, 2
|
||||
(labels_dir / "img1.txt").write_text("0 0.5 0.5 0.2 0.2\n", encoding="utf-8")
|
||||
(labels_dir / "img2.txt").write_text("2 0.4 0.4 0.1 0.1\n", encoding="utf-8")
|
||||
|
||||
classes = read_classes(tmp_path, "")
|
||||
# Should generate class_0, class_1, class_2 since max_id is 2
|
||||
assert classes == {0: "class_0", 1: "class_1", 2: "class_2"}
|
||||
|
||||
|
||||
def test_split_dataset_flow(tmp_path: Path) -> None:
|
||||
images_dir = tmp_path / "images"
|
||||
labels_dir = tmp_path / "labels"
|
||||
images_dir.mkdir()
|
||||
labels_dir.mkdir()
|
||||
|
||||
# Create 4 image files
|
||||
for i in range(1, 5):
|
||||
(images_dir / f"img{i}.png").write_text("", encoding="utf-8")
|
||||
(labels_dir / f"img{i}.txt").write_text(f"0 0.5 0.5 0.1 0.1\n", encoding="utf-8")
|
||||
|
||||
# Create classes.txt
|
||||
(tmp_path / "classes.txt").write_text("dummy_class\n", encoding="utf-8")
|
||||
|
||||
# Split with 75% train ratio -> 3 train, 1 val
|
||||
train_count, val_count, yaml_path = split_dataset(
|
||||
dataset_dir=str(tmp_path),
|
||||
train_ratio=0.75,
|
||||
classes_path=""
|
||||
)
|
||||
|
||||
assert train_count == 3
|
||||
assert val_count == 1
|
||||
assert Path(yaml_path).exists()
|
||||
|
||||
# Verify dataset.yaml content
|
||||
with open(yaml_path, "r", encoding="utf-8") as f:
|
||||
data = yaml.safe_load(f)
|
||||
assert data["path"] == str(tmp_path)
|
||||
assert data["train"] == "split/train.txt"
|
||||
assert data["val"] == "split/val.txt"
|
||||
assert data["names"] == {0: "dummy_class"}
|
||||
|
||||
# Verify lists content
|
||||
train_list = (tmp_path / "split" / "train.txt").read_text(encoding="utf-8").strip().split("\n")
|
||||
val_list = (tmp_path / "split" / "val.txt").read_text(encoding="utf-8").strip().split("\n")
|
||||
|
||||
assert len(train_list) == 3
|
||||
assert len(val_list) == 1
|
||||
|
||||
# Verify paths are absolute
|
||||
assert Path(train_list[0]).is_absolute()
|
||||
assert Path(val_list[0]).is_absolute()
|
||||
Loading…
Add table
Reference in a new issue