bone_2026/tests/test_labels.py

173 lines
7.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
Тесты разбора датасета DXA и построения меток.
Проверяют правила, от которых зависит обучение: метка из имени файла,
склейка побайтных дублей и разбиение по исследованиям без утечки.
"""
import sys
from pathlib import Path
import pytest
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from src.dxa.labels import ( # noqa: E402
QUALITY_BAD,
QUALITY_GOOD,
ImageRecord,
format_summary,
label_summary,
marker_from_filename,
region_from_filename,
scan_dataset,
stratified_group_split,
)
DATASET_ROOT = Path("dataset_hack")
HAVE_DATASET = (DATASET_ROOT / "НД_для_обучения" / "Исследования").is_dir()
needs_dataset = pytest.mark.skipif(not HAVE_DATASET, reason="dataset_hack is not available")
class TestRegionFromFilename:
"""Имена файлов в датасете неоднородны; все встречающиеся варианты должны разбираться."""
@pytest.mark.parametrize("name,expected", [
("spine_01.dcm", "spine"),
("spine_1_bad.dcm", "spine"),
("spine-1.dcm", "spine"),
("Spine.dcm", "spine"),
("r_spine_03.dcm", "spine"),
("l_hip_01.dcm", "hip_left"),
("l_hip-2.dcm", "hip_left"),
("l_hip_1_good.dcm", "hip_left"),
("r_hip_12.dcm", "hip_right"),
("r_hop_1.dcm", "hip_right"),
])
def test_recognised_variants(self, name, expected):
assert region_from_filename(name) == expected
def test_unknown_name_returns_none(self):
assert region_from_filename("scan_0001.dcm") is None
class TestMarkerFromFilename:
@pytest.mark.parametrize("name,expected", [
("spine_01_bad.dcm", "bad"),
("l_hip_02_good.dcm", "good"),
("spine-1_bad.dcm", "bad"),
("r_hip_1_good.dcm", "good"),
("bad.dcm", "bad"),
("good.dcm", "good"),
])
def test_explicit_markers(self, name, expected):
assert marker_from_filename(name) == expected
@pytest.mark.parametrize("name", ["spine_01.dcm", "l_hip_2.dcm", "r_hip_03.dcm"])
def test_missing_marker_is_none(self, name):
assert marker_from_filename(name) is None
def test_marker_not_matched_inside_word(self):
"""«bad» в середине имени не является меткой."""
assert marker_from_filename("spine_bad_extra.dcm") is None
def _record(study, region, label, name="img.dcm"):
return ImageRecord(path=Path(f"/tmp/{study}/{name}"), study=study, region=region, label=label)
class TestStratifiedGroupSplit:
"""Разбиение не должно допускать утечки между train и val."""
def test_no_study_appears_in_both_parts(self):
records = []
for i in range(40):
study = f"study_{i:02d}"
# Каждое исследование даёт 1–3 снимка, часть из них с нарушением.
records.append(_record(study, "spine", QUALITY_BAD if i % 4 == 0 else QUALITY_GOOD, "a.dcm"))
records.append(_record(study, "spine", QUALITY_GOOD, "b.dcm"))
if i % 3 == 0:
records.append(_record(study, "hip_left", QUALITY_GOOD, "c.dcm"))
train, val = stratified_group_split(records, val_fraction=0.25, seed=0)
train_studies = {r.study for r in train}
val_studies = {r.study for r in val}
assert not (train_studies & val_studies), "study leaked between train and val"
assert len(train) + len(val) == len(records)
assert train and val
def test_both_splits_contain_both_classes(self):
records = []
for i in range(30):
study = f"s{i:02d}"
records.append(_record(study, "spine", QUALITY_BAD if i % 3 == 0 else QUALITY_GOOD))
train, val = stratified_group_split(records, val_fraction=0.3, seed=7)
assert {r.label for r in train} == {QUALITY_GOOD, QUALITY_BAD}
assert {r.label for r in val} == {QUALITY_GOOD, QUALITY_BAD}
def test_deterministic_for_same_seed(self):
records = [_record(f"s{i:02d}", "spine", i % 2) for i in range(20)]
assert stratified_group_split(records, seed=5) == stratified_group_split(records, seed=5)
def test_rejects_invalid_fraction(self):
with pytest.raises(ValueError):
stratified_group_split([_record("s0", "spine", 0)], val_fraction=0.0)
def test_empty_input(self):
assert stratified_group_split([]) == ([], [])
class TestLabelSummary:
def test_counts_by_region_and_class(self):
records = [
_record("s1", "spine", QUALITY_BAD),
_record("s1", "spine", QUALITY_GOOD),
_record("s2", "hip_left", QUALITY_GOOD),
_record("s3", None, QUALITY_GOOD),
]
summary = label_summary(records)
assert summary["spine"] == {"good": 1, "bad": 1, "total": 2}
assert summary["hip_left"] == {"good": 1, "bad": 0, "total": 1}
assert summary["unknown"]["total"] == 1
def test_format_summary_mentions_totals(self):
text = format_summary([_record("s1", "spine", QUALITY_BAD), _record("s1", "spine", QUALITY_GOOD)])
assert "bad=" in text and "good=" in text
assert "Unique images: 2" in text
@needs_dataset
class TestRealDataset:
"""Проверки на реальном датасете: инварианты, влияющие на обучение."""
@pytest.fixture(scope="class")
def records(self):
return scan_dataset(DATASET_ROOT, with_pixel_dedup=True)
def test_no_label_conflicts_after_dedup(self):
"""Склейка дублей не должна порождать изображения с противоречивой меткой."""
records = scan_dataset(DATASET_ROOT, with_pixel_dedup=True)
assert all(r.label in (QUALITY_GOOD, QUALITY_BAD) for r in records)
def test_dedup_reduces_more_than_file_count(self):
"""Файлов на диске больше, чем уникальных изображений."""
raw_files = list((DATASET_ROOT / "НД_для_обучения" / "Исследования").rglob("*.dcm"))
with_dedup = scan_dataset(DATASET_ROOT, with_pixel_dedup=True)
without_dedup = scan_dataset(DATASET_ROOT, with_pixel_dedup=False)
assert len(with_dedup) < len(without_dedup)
assert len(without_dedup) <= len(raw_files)
def test_most_records_have_a_region(self, records):
known = sum(1 for r in records if r.region in ("spine", "hip_left", "hip_right"))
assert known / len(records) > 0.95
def test_bad_share_is_minority(self, records):
bad = sum(1 for r in records if r.label == QUALITY_BAD)
assert 0.05 < bad / len(records) < 0.30
def test_split_on_real_data_leaks_nothing(self, records):
train, val = stratified_group_split(records, val_fraction=0.2, seed=42)
assert not ({r.study for r in train} & {r.study for r in val})
assert train and val