223 lines
9.2 KiB
Python
223 lines
9.2 KiB
Python
"""
|
||
Тесты разбора датасета DXA и построения меток.
|
||
|
||
Проверяют правила, от которых зависит обучение: метка из имени файла,
|
||
склейка побайтных дублей и разбиение по исследованиям без утечки.
|
||
"""
|
||
import json
|
||
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,
|
||
export_split,
|
||
format_summary,
|
||
label_summary,
|
||
load_split,
|
||
marker_from_filename,
|
||
region_from_filename,
|
||
scan_dataset,
|
||
split_by_studies,
|
||
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
|
||
|
||
|
||
class TestFixedSplit:
|
||
"""Фиксированное разбиение: один held-out набор для сравнения вариантов меток."""
|
||
|
||
def _records(self, studies):
|
||
return [
|
||
ImageRecord(path=Path(f"{study}/img.dcm"), study=study, region="spine", label=label)
|
||
for study, label in studies
|
||
]
|
||
|
||
def test_export_and_load_roundtrip(self, tmp_path):
|
||
records = self._records([(f"s{i}", i % 2) for i in range(10)])
|
||
path = export_split(records, tmp_path / "split.json", val_fraction=0.3, seed=1)
|
||
|
||
payload = json.loads(path.read_text())
|
||
assert payload["val_fraction"] == 0.3 and payload["seed"] == 1
|
||
assert payload["val_images"] == len(payload["val_studies"])
|
||
|
||
val_studies = load_split(path)
|
||
assert val_studies == set(payload["val_studies"])
|
||
assert val_studies <= {r.study for r in records}
|
||
|
||
def test_split_by_studies_matches_contract(self):
|
||
records = self._records([(f"s{i}", i % 2) for i in range(10)])
|
||
train, val = split_by_studies(records, ["s0", "s1", "s2"])
|
||
assert {r.study for r in val} == {"s0", "s1", "s2"}
|
||
assert {r.study for r in train} == {f"s{i}" for i in range(3, 10)}
|
||
assert len(train) + len(val) == len(records)
|
||
|
||
def test_split_by_studies_rejects_empty_side(self):
|
||
records = self._records([("s0", 0), ("s1", 1)])
|
||
with pytest.raises(ValueError, match="degenerated"):
|
||
split_by_studies(records, ["s0", "s1"])
|
||
|
||
def test_load_split_rejects_empty_file(self, tmp_path):
|
||
path = tmp_path / "bad.json"
|
||
path.write_text(json.dumps({"val_studies": []}))
|
||
with pytest.raises(ValueError, match="val_studies"):
|
||
load_split(path)
|
||
|
||
def test_exported_split_is_stable_for_same_inputs(self, tmp_path):
|
||
records = self._records([(f"s{i}", i % 2) for i in range(10)])
|
||
first = load_split(export_split(records, tmp_path / "a.json", seed=7))
|
||
second = load_split(export_split(records, tmp_path / "b.json", seed=7))
|
||
assert first == second
|
||
|
||
|
||
@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
|