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