100 lines
3.3 KiB
Python
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()
|