train_utility/tests/test_splitter.py

100 lines
3.3 KiB
Python

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