194 lines
6.7 KiB
Python
194 lines
6.7 KiB
Python
from __future__ import annotations
|
|
|
|
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"}
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("names", "expected"),
|
|
[
|
|
(["cat", "dog"], {0: "cat", 1: "dog"}),
|
|
({0: "cat", 1: "dog"}, {0: "cat", 1: "dog"}),
|
|
],
|
|
)
|
|
def test_read_classes_custom_yaml(
|
|
tmp_path: Path, names: object, expected: dict[int, str]
|
|
) -> None:
|
|
custom_file = tmp_path / "custom_classes.yaml"
|
|
custom_file.write_text(
|
|
yaml.safe_dump({"names": names}, allow_unicode=True),
|
|
encoding="utf-8",
|
|
)
|
|
|
|
assert read_classes(tmp_path, str(custom_file)) == expected
|
|
|
|
|
|
def test_missing_custom_classes_path_does_not_fall_back(tmp_path: Path) -> None:
|
|
(tmp_path / "classes.txt").write_text("fallback\n", encoding="utf-8")
|
|
|
|
with pytest.raises(ValueError, match="не существует"):
|
|
read_classes(tmp_path, str(tmp_path / "typo.yaml"))
|
|
|
|
|
|
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)
|
|
relative_output_dir = Path(yaml_path).parent.relative_to(tmp_path)
|
|
assert data["train"] == (relative_output_dir / "train.txt").as_posix()
|
|
assert data["val"] == (relative_output_dir / "val.txt").as_posix()
|
|
assert data["names"] == {0: "dummy_class"}
|
|
|
|
# Verify lists content
|
|
output_dir = Path(yaml_path).parent
|
|
train_list = (output_dir / "train.txt").read_text(encoding="utf-8").strip().split("\n")
|
|
val_list = (output_dir / "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()
|
|
|
|
|
|
def test_split_dataset_rejects_single_image_without_writing_output(
|
|
tmp_path: Path,
|
|
) -> None:
|
|
images_dir = tmp_path / "images"
|
|
labels_dir = tmp_path / "labels"
|
|
images_dir.mkdir()
|
|
labels_dir.mkdir()
|
|
(images_dir / "only.jpg").write_bytes(b"")
|
|
(labels_dir / "only.txt").write_text("0 0.5 0.5 1 1\n", encoding="utf-8")
|
|
(tmp_path / "classes.txt").write_text("item\n", encoding="utf-8")
|
|
|
|
with pytest.raises(ValueError, match="минимум 2"):
|
|
split_dataset(str(tmp_path), 0.8, "")
|
|
|
|
assert not (tmp_path / ".yolo-tui").exists()
|
|
|
|
|
|
def test_split_dataset_finds_nested_images_and_labels(tmp_path: Path) -> None:
|
|
images_dir = tmp_path / "images" / "day"
|
|
labels_dir = tmp_path / "labels" / "day"
|
|
images_dir.mkdir(parents=True)
|
|
labels_dir.mkdir(parents=True)
|
|
for index in range(2):
|
|
(images_dir / f"nested-{index}.jpg").write_bytes(b"")
|
|
(labels_dir / f"nested-{index}.txt").write_text(
|
|
"2 0.5 0.5 0.2 0.2\n",
|
|
encoding="utf-8",
|
|
)
|
|
|
|
train_count, val_count, yaml_path = split_dataset(str(tmp_path), 0.8, "")
|
|
|
|
assert (train_count, val_count) == (1, 1)
|
|
data = yaml.safe_load(Path(yaml_path).read_text(encoding="utf-8"))
|
|
assert data["names"] == {0: "class_0", 1: "class_1", 2: "class_2"}
|
|
listed_images = "".join(
|
|
(Path(yaml_path).parent / filename).read_text(encoding="utf-8")
|
|
for filename in ("train.txt", "val.txt")
|
|
)
|
|
assert "images/day/nested-0.jpg" in listed_images
|
|
assert "images/day/nested-1.jpg" in listed_images
|
|
|
|
|
|
def test_split_dataset_preserves_existing_split_files(tmp_path: Path) -> None:
|
|
images_dir = tmp_path / "images"
|
|
labels_dir = tmp_path / "labels"
|
|
images_dir.mkdir()
|
|
labels_dir.mkdir()
|
|
for index in range(2):
|
|
(images_dir / f"image-{index}.jpg").write_bytes(b"")
|
|
(labels_dir / f"image-{index}.txt").write_text(
|
|
"0 0.5 0.5 0.2 0.2\n",
|
|
encoding="utf-8",
|
|
)
|
|
(tmp_path / "classes.txt").write_text("item\n", encoding="utf-8")
|
|
user_split = tmp_path / "split"
|
|
user_split.mkdir()
|
|
(user_split / "train.txt").write_text("user data\n", encoding="utf-8")
|
|
|
|
first_yaml = Path(split_dataset(str(tmp_path), 0.5, "")[2])
|
|
second_yaml = Path(split_dataset(str(tmp_path), 0.5, "")[2])
|
|
|
|
assert (user_split / "train.txt").read_text(encoding="utf-8") == "user data\n"
|
|
assert first_yaml.parent != second_yaml.parent
|
|
assert user_split not in first_yaml.parents
|