454 lines
20 KiB
Python
454 lines
20 KiB
Python
"""
|
||
Тесты разметки снимков по экспертной таблице `разметка.xlsx`.
|
||
|
||
Проверяют правила, от которых зависит итоговая метка: чтение критериев,
|
||
трактовку «1 = нарушение», голосование по анатомической области, перенос
|
||
оценки на единственный снимок бедра и объединение источников
|
||
(экспертная таблица ИЛИ имя файла).
|
||
"""
|
||
import sys
|
||
from pathlib import Path
|
||
|
||
import pandas as pd
|
||
import pytest
|
||
|
||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||
|
||
from src.dxa.excel_labels import ( # noqa: E402
|
||
ARTIFACT,
|
||
AXIS_DEVIATION,
|
||
POSITIONING,
|
||
ROTATION,
|
||
UNSPECIFIED,
|
||
RegionCriteria,
|
||
StudyCriteria,
|
||
_pick_criteria,
|
||
apply_excel_labels,
|
||
format_summary,
|
||
label_dataset,
|
||
load_study_criteria,
|
||
resolve_region,
|
||
)
|
||
from src.dxa.labels import ImageRecord # noqa: E402
|
||
from src.dxa.dataset import make_datasets # noqa: E402
|
||
from src.dxa.train import resolve_labels_csv # noqa: E402
|
||
|
||
DATASET_ROOT = Path("dataset_hack")
|
||
EXCEL_PATH = DATASET_ROOT / "НД_для_обучения" / "разметка.xlsx"
|
||
LABELS_CSV = Path("labels/labels_images.csv")
|
||
HAVE_DATASET = (DATASET_ROOT / "НД_для_обучения" / "Исследования").is_dir()
|
||
|
||
needs_dataset = pytest.mark.skipif(not HAVE_DATASET, reason="dataset_hack is not available")
|
||
|
||
_WIDTH = 19
|
||
|
||
|
||
def _write_excel(path: Path, studies: list) -> Path:
|
||
"""Собрать файл разметки того же вида, что и настоящий: две строки заголовков."""
|
||
table = [[None] * _WIDTH for _ in range(2 + len(studies))]
|
||
for offset, spec in enumerate(studies):
|
||
row = table[2 + offset]
|
||
row[1] = spec.get("study")
|
||
for i, value in enumerate(spec.get("spine", [])):
|
||
row[2 + i] = value
|
||
for i, value in enumerate(spec.get("hip_right", [])):
|
||
row[5 + i] = value
|
||
for i, value in enumerate(spec.get("hip_left", [])):
|
||
row[7 + i] = value
|
||
for column, key in ((9, "spine_total"), (10, "hip_right_total"), (11, "hip_left_total")):
|
||
if key in spec:
|
||
row[column] = spec[key]
|
||
row[12] = spec.get("comment")
|
||
pd.DataFrame(table).to_excel(path, header=False, index=False)
|
||
return path
|
||
|
||
|
||
def _criteria(**regions) -> StudyCriteria:
|
||
"""StudyCriteria из коротких описаний: criteria=(0, 1), total=1."""
|
||
cells = {}
|
||
for name, spec in regions.items():
|
||
criteria, total = spec
|
||
keys = ("positioning", "axis_deviation", "artifact") if name == "spine" else ("rotation", "roi_incorrect")
|
||
cells[name] = RegionCriteria(violations=dict(zip(keys, criteria)), total=total)
|
||
return StudyCriteria(study="s", regions=cells)
|
||
|
||
|
||
class TestLoadStudyCriteria:
|
||
"""Файл разметки содержит заголовки и разреженные ячейки; разбор не должен падать."""
|
||
|
||
def test_reads_criteria_totals_and_comment(self, tmp_path):
|
||
path = _write_excel(tmp_path / "a.xlsx", [
|
||
{"study": "S1", "spine": [0, 1, 0], "spine_total": 1,
|
||
"hip_right": [1, 0], "hip_right_total": 1, "comment": "сколиоз"},
|
||
{"study": "S2", "spine": [0, 0, 0], "spine_total": 0},
|
||
])
|
||
criteria = load_study_criteria(path)
|
||
|
||
assert set(criteria) == {"S1", "S2"}
|
||
spine = criteria["S1"].for_region("spine")
|
||
assert spine.total == 1 and spine.bad
|
||
assert spine.violated() == [AXIS_DEVIATION]
|
||
assert criteria["S1"].comment == "сколиоз"
|
||
assert criteria["S1"].for_region("hip_right").violated() == [ROTATION]
|
||
assert not criteria["S2"].for_region("spine").bad
|
||
|
||
def test_missing_region_is_unknown(self, tmp_path):
|
||
path = _write_excel(tmp_path / "a.xlsx", [{"study": "S1", "spine": [0, 0, 0], "spine_total": 0}])
|
||
criteria = load_study_criteria(path)
|
||
cell = criteria["S1"].for_region("hip_left")
|
||
assert cell is not None and not cell.known
|
||
|
||
def test_row_without_study_is_skipped(self, tmp_path):
|
||
path = _write_excel(tmp_path / "a.xlsx", [
|
||
{"study": None, "spine": [1, 0, 0], "spine_total": 1},
|
||
{"study": "S1", "spine": [0, 0, 0], "spine_total": 0},
|
||
{"study": "S2", "spine": [1, 0, 0], "spine_total": 1},
|
||
])
|
||
assert set(load_study_criteria(path)) == {"S1", "S2"}
|
||
|
||
|
||
class TestRegionCriteria:
|
||
"""Трактовка значений: 1 в критерии и в итоге означает нарушение."""
|
||
|
||
def test_total_without_criteria_is_unspecified(self):
|
||
cell = RegionCriteria(violations={"positioning": 0, "axis_deviation": 0, "artifact": 0}, total=1)
|
||
assert cell.bad and cell.violated() == [UNSPECIFIED]
|
||
|
||
def test_criterion_without_total_still_flags(self):
|
||
# В датасете есть случай: критерий отмечен, итог оставлен нулевым.
|
||
cell = RegionCriteria(violations={"positioning": 0, "axis_deviation": 1, "artifact": 0}, total=0)
|
||
assert cell.bad and cell.violated() == [AXIS_DEVIATION]
|
||
|
||
def test_all_zero_is_good(self):
|
||
cell = RegionCriteria(violations={"rotation": 0, "roi_incorrect": 0}, total=0)
|
||
assert not cell.bad and cell.violated() == []
|
||
assert cell.known
|
||
|
||
def test_all_missing_is_unknown(self):
|
||
cell = RegionCriteria(violations={"rotation": None, "roi_incorrect": None}, total=None)
|
||
assert not cell.known and not cell.bad
|
||
|
||
def test_multiple_criteria_are_all_reported(self):
|
||
cell = RegionCriteria(violations={"positioning": 1, "axis_deviation": 1, "artifact": 1}, total=1)
|
||
assert cell.violated() == [POSITIONING, AXIS_DEVIATION, ARTIFACT]
|
||
|
||
|
||
class TestResolveRegion:
|
||
"""Побайтные дубли иногда названы по-разному; область решается голосованием."""
|
||
|
||
def test_single_name(self):
|
||
assert resolve_region([Path("spine_01.dcm")]) == ("spine", False)
|
||
|
||
def test_majority_wins(self):
|
||
sources = [Path(f"r_hip_{i}.dcm") for i in range(4)] + [Path("r_spine_03.dcm")]
|
||
assert resolve_region(sources) == ("hip_right", False)
|
||
|
||
def test_tie_is_ambiguous(self):
|
||
region, ambiguous = resolve_region([Path("spine_03_bad.dcm"), Path("l_hip_01.dcm")])
|
||
assert region is None and ambiguous
|
||
|
||
def test_unknown_names(self):
|
||
assert resolve_region([Path("bad.dcm")]) == (None, True)
|
||
|
||
|
||
class TestPickCriteria:
|
||
"""Оценка берётся со своей стороны; зеркальный перенос — только для единственного бедра."""
|
||
|
||
def test_same_side_is_used(self):
|
||
study = _criteria(hip_right=((1, 0), 1), hip_left=((0, 0), 0))
|
||
cell, mirrored = _pick_criteria(study, "hip_right", hip_image_count=2)
|
||
assert cell.violated() == [ROTATION] and not mirrored
|
||
|
||
def test_single_hip_borrows_opposite_side(self):
|
||
study = _criteria(hip_right=((None, None), None), hip_left=((1, 0), 1))
|
||
cell, mirrored = _pick_criteria(study, "hip_right", hip_image_count=1)
|
||
assert cell.violated() == [ROTATION] and mirrored
|
||
|
||
def test_single_hip_without_opposite_stays_unknown(self):
|
||
study = _criteria(hip_right=((None, None), None), hip_left=((None, None), None))
|
||
cell, mirrored = _pick_criteria(study, "hip_right", hip_image_count=1)
|
||
assert not cell.known and not mirrored
|
||
|
||
def test_two_hips_never_borrow(self):
|
||
study = _criteria(hip_right=((None, None), None), hip_left=((1, 0), 1))
|
||
cell, mirrored = _pick_criteria(study, "hip_right", hip_image_count=2)
|
||
assert cell is not None and not cell.known and not mirrored
|
||
|
||
def test_spine_never_borrows(self):
|
||
study = _criteria(spine=((None, None, None), None), hip_left=((1, 0), 1))
|
||
cell, mirrored = _pick_criteria(study, "spine", hip_image_count=0)
|
||
assert not cell.known and not mirrored
|
||
|
||
|
||
class TestApplyExcelLabels:
|
||
"""Подстановка построенной разметки в записи датасета."""
|
||
|
||
def _csv(self, tmp_path, rows):
|
||
path = tmp_path / "labels.csv"
|
||
path.write_text(
|
||
"path_to_image,quality_class,anatomical_region\n"
|
||
+ "".join(f"{p},{q},{r}\n" for p, q, r in rows),
|
||
encoding="utf-8",
|
||
)
|
||
return path
|
||
|
||
def _record(self, path, region=None, label=0, marker=None):
|
||
return ImageRecord(path=Path(path), study="S1", region=region, label=label, marker=marker)
|
||
|
||
def test_replaces_labels_and_regions(self, tmp_path):
|
||
csv = self._csv(tmp_path, [("/s/a.dcm", 1, "hip_right")])
|
||
records = [self._record("/s/a.dcm", region=None, label=0)]
|
||
out = apply_excel_labels(records, csv)
|
||
assert out[0].label == 1
|
||
assert out[0].region == "hip_right"
|
||
|
||
def test_keeps_filename_label_when_image_absent(self, tmp_path):
|
||
csv = self._csv(tmp_path, [("/s/a.dcm", 1, "spine")])
|
||
records = [
|
||
self._record("/s/a.dcm", region="spine"),
|
||
self._record("/s/b.dcm", region="spine", label=1, marker="bad"),
|
||
]
|
||
out = apply_excel_labels(records, csv)
|
||
assert [r.label for r in out] == [1, 1]
|
||
# запись, которой нет в CSV, сохраняет метку из имени файла
|
||
assert out[1].path == Path("/s/b.dcm")
|
||
|
||
def test_other_dataset_raises(self, tmp_path):
|
||
csv = self._csv(tmp_path, [("/other/a.dcm", 1, "spine")])
|
||
with pytest.raises(ValueError, match="None of the"):
|
||
apply_excel_labels([self._record("/s/a.dcm")], csv)
|
||
|
||
def test_missing_columns_raise(self, tmp_path):
|
||
path = tmp_path / "bad.csv"
|
||
path.write_text("file,label\n/s/a.dcm,1\n", encoding="utf-8")
|
||
with pytest.raises(ValueError, match="lacks columns"):
|
||
apply_excel_labels([self._record("/s/a.dcm")], path)
|
||
|
||
def test_empty_region_becomes_none(self, tmp_path):
|
||
csv = self._csv(tmp_path, [("/s/a.dcm", 1, "")])
|
||
out = apply_excel_labels([self._record("/s/a.dcm", region="spine")], csv)
|
||
assert out[0].region is None
|
||
|
||
|
||
class TestResolveLabelsCsv:
|
||
"""Отсутствующий файл разметки не должен молча менять источник меток."""
|
||
|
||
def test_empty_value_disables(self):
|
||
assert resolve_labels_csv("") is None
|
||
assert resolve_labels_csv(None) is None
|
||
|
||
def test_missing_file_falls_back_with_warning(self, tmp_path, caplog):
|
||
with caplog.at_level("WARNING"):
|
||
assert resolve_labels_csv(str(tmp_path / "nope.csv")) is None
|
||
assert "not found" in caplog.text
|
||
|
||
def test_existing_file_is_returned(self, tmp_path):
|
||
path = tmp_path / "labels.csv"
|
||
path.write_text("path_to_image,quality_class\n/s/a.dcm,1\n", encoding="utf-8")
|
||
assert resolve_labels_csv(str(path)) == path
|
||
|
||
|
||
@needs_dataset
|
||
class TestLabelRules:
|
||
"""Правило метки: объединение источников, только таблица или чистый эталон."""
|
||
|
||
@pytest.fixture(scope="class")
|
||
def union(self):
|
||
return label_dataset(DATASET_ROOT, EXCEL_PATH, label_rule="union")
|
||
|
||
@pytest.fixture(scope="class")
|
||
def table(self):
|
||
return label_dataset(DATASET_ROOT, EXCEL_PATH, label_rule="table")
|
||
|
||
@pytest.fixture(scope="class")
|
||
def expert(self):
|
||
return label_dataset(DATASET_ROOT, EXCEL_PATH, label_rule="expert")
|
||
|
||
def test_union_uses_filename_evidence(self, union, table):
|
||
assert sum(1 for l in union if l.quality == 1) == 92
|
||
assert sum(1 for l in table if l.quality == 1) == 77
|
||
|
||
def test_table_rule_ignores_filename_marker(self, table):
|
||
# Снимок, помеченный `_bad`, но признанный экспертом качественным,
|
||
# в режиме «только таблица» остаётся качественным.
|
||
mismatched = [l for l in table if l.quality_excel == 0 and l.quality_filename == 1]
|
||
assert mismatched, "в наборе должны быть расхождения источников"
|
||
assert all(l.quality == 0 for l in mismatched)
|
||
|
||
def test_table_rule_keeps_filename_only_as_fallback(self, table):
|
||
fallback = [l for l in table if l.used_filename_fallback]
|
||
assert len(fallback) == 3
|
||
assert all(l.quality_excel is None for l in fallback)
|
||
|
||
def test_expert_rule_drops_unscored_images(self, expert):
|
||
assert len(expert) == 249
|
||
assert all(l.quality_excel is not None for l in expert)
|
||
assert not any(l.used_filename_fallback for l in expert)
|
||
|
||
def test_rules_agree_where_labels_agree(self, union, table):
|
||
"""Там, где правила дают одну метку, совпадают и типы нарушений."""
|
||
by_uid = {l.dicom_image_uid: l for l in union}
|
||
agreed = 0
|
||
for label in table:
|
||
reference = by_uid[label.dicom_image_uid]
|
||
if label.quality != reference.quality:
|
||
continue
|
||
assert label.violations == reference.violations
|
||
agreed += 1
|
||
# Расходятся ровно 15 снимков (см. следующий тест), остальные совпадают.
|
||
assert agreed == 252 - 15
|
||
|
||
def test_filename_only_violations_become_unspecified(self, union, table):
|
||
"""
|
||
Снимок, который эксперт считает качественным, а имя файла — нарушением:
|
||
в режиме объединения он получает `unspecified`, в режиме таблицы — ничего.
|
||
"""
|
||
by_uid = {l.dicom_image_uid: l for l in union}
|
||
disputed = [l for l in table if l.quality_excel == 0 and l.quality_filename == 1]
|
||
assert len(disputed) == 15
|
||
for label in disputed:
|
||
assert label.quality == 0 and label.violations == ()
|
||
assert by_uid[label.dicom_image_uid].quality == 1
|
||
assert by_uid[label.dicom_image_uid].violations == (UNSPECIFIED,)
|
||
|
||
def test_unknown_rule_is_rejected(self):
|
||
with pytest.raises(ValueError, match="label_rule"):
|
||
label_dataset(DATASET_ROOT, EXCEL_PATH, label_rule="nonsense")
|
||
|
||
def test_rule_is_recorded_in_output(self, union, table, expert):
|
||
assert {l.label_rule for l in union} == {"union"}
|
||
assert {l.label_rule for l in table} == {"table"}
|
||
assert {l.label_rule for l in expert} == {"expert"}
|
||
|
||
|
||
class TestParseVariants:
|
||
"""Разбор описания вариантов для сравнения правил разметки."""
|
||
|
||
def test_default(self):
|
||
from src.dxa.compare_labels import DEFAULT_VARIANTS, parse_variants
|
||
|
||
assert parse_variants(None) == DEFAULT_VARIANTS
|
||
|
||
def test_custom(self):
|
||
from src.dxa.compare_labels import parse_variants
|
||
|
||
assert parse_variants("a=/tmp/a.csv, b=") == {"a": "/tmp/a.csv", "b": ""}
|
||
|
||
def test_single_variant_is_rejected(self):
|
||
from src.dxa.compare_labels import parse_variants
|
||
|
||
with pytest.raises(ValueError, match="минимум два"):
|
||
parse_variants("only=/tmp/a.csv")
|
||
|
||
|
||
@needs_dataset
|
||
class TestLabelDataset:
|
||
"""Сквозная проверка на реальном датасете: инварианты разметки."""
|
||
|
||
@pytest.fixture(scope="class")
|
||
def labels(self):
|
||
return label_dataset(DATASET_ROOT, EXCEL_PATH)
|
||
|
||
def test_covers_every_unique_image(self, labels):
|
||
assert len(labels) == 252
|
||
|
||
def test_region_appears_once_per_study(self, labels):
|
||
seen = {}
|
||
for label in labels:
|
||
if label.region is None:
|
||
continue
|
||
seen.setdefault((label.record.study, label.region), 0)
|
||
seen[(label.record.study, label.region)] += 1
|
||
assert all(count == 1 for count in seen.values())
|
||
assert len(seen) == 251
|
||
|
||
def test_union_rule_holds(self):
|
||
"""Арифметика объединения источников (правило `union`, не по умолчанию)."""
|
||
labels = label_dataset(DATASET_ROOT, EXCEL_PATH, label_rule="union")
|
||
excel = sum(1 for l in labels if l.quality_excel == 1)
|
||
filename = sum(1 for l in labels if l.quality_filename == 1)
|
||
both = sum(1 for l in labels if l.quality_excel == 1 and l.quality_filename == 1)
|
||
assert sum(1 for l in labels if l.quality == 1) == excel + filename - both
|
||
assert (excel, filename, both) == (74, 37, 19)
|
||
assert sum(1 for l in labels if l.quality == 1) == 92
|
||
|
||
def test_default_rule_is_table(self, labels):
|
||
"""По умолчанию в метке участвует только экспертная таблица."""
|
||
assert {l.label_rule for l in labels} == {"table"}
|
||
assert sum(1 for l in labels if l.quality == 1) == 77
|
||
|
||
def test_violations_match_excel_criteria(self, labels):
|
||
criteria = load_study_criteria(EXCEL_PATH)
|
||
checked = 0
|
||
for label in labels:
|
||
if label.region is None or label.quality_excel != 1:
|
||
continue
|
||
cell = criteria[label.record.study].for_region(label.region)
|
||
if cell is None or not cell.known:
|
||
continue
|
||
assert label.violations == tuple(cell.violated())
|
||
checked += 1
|
||
assert checked > 60
|
||
|
||
def test_good_images_have_no_violations(self, labels):
|
||
assert all(not l.violations for l in labels if l.quality == 0)
|
||
|
||
def test_flagged_laterality_mirroring(self, labels):
|
||
flagged = [l for l in labels if l.laterality_mirrored]
|
||
assert len(flagged) == 6
|
||
assert all(l.region in ("hip_left", "hip_right") for l in flagged)
|
||
assert all(l.quality_excel is not None for l in flagged)
|
||
|
||
def test_single_ambiguous_region_is_labelled_from_filename(self, labels):
|
||
ambiguous = [l for l in labels if l.region_ambiguous]
|
||
assert len(ambiguous) == 1
|
||
label = ambiguous[0]
|
||
assert label.region is None
|
||
assert label.quality == 1 and label.violations == (UNSPECIFIED,)
|
||
|
||
def test_uids_are_populated(self, labels):
|
||
assert all(l.dicom_image_uid and l.dicom_study_uid for l in labels)
|
||
assert len({l.dicom_image_uid for l in labels}) == 252
|
||
assert len({l.dicom_study_uid for l in labels}) == 100
|
||
|
||
def test_summary_mentions_regions(self, labels):
|
||
text = format_summary(labels)
|
||
for region in ("spine", "hip_right", "hip_left"):
|
||
assert region in text
|
||
assert "ALL" in text
|
||
|
||
|
||
@needs_dataset
|
||
class TestTrainingWiring:
|
||
"""Обучение должно брать метки из построенного CSV, сохраняя разбиение по исследованиям."""
|
||
|
||
def test_datasets_use_csv_labels(self):
|
||
train_ds, val_ds, _ = make_datasets(
|
||
DATASET_ROOT, EXCEL_PATH, labels_csv=LABELS_CSV, seed=42
|
||
)
|
||
records = list(train_ds.records) + list(val_ds.records)
|
||
assert len(records) == 252
|
||
# Официальная разметка — правило `table`: 74 нарушения эксперта плюс 3 снимка,
|
||
# для которых таблица область не оценивала.
|
||
assert sum(r.label for r in records) == 77
|
||
|
||
def test_split_still_has_no_study_leak(self):
|
||
train_ds, val_ds, _ = make_datasets(
|
||
DATASET_ROOT, EXCEL_PATH, labels_csv=LABELS_CSV, seed=42
|
||
)
|
||
overlap = {r.study for r in train_ds.records} & {r.study for r in val_ds.records}
|
||
assert overlap == set()
|
||
|
||
def test_fallback_to_filename_labels(self):
|
||
train_ds, val_ds, _ = make_datasets(DATASET_ROOT, EXCEL_PATH, labels_csv=None, seed=42)
|
||
records = list(train_ds.records) + list(val_ds.records)
|
||
assert sum(r.label for r in records) == 37
|
||
|
||
def test_region_head_sees_resolved_region(self):
|
||
# Область одного конфликтного дубля решается голосованием имён и приходит из CSV.
|
||
train_ds, val_ds, _ = make_datasets(
|
||
DATASET_ROOT, EXCEL_PATH, labels_csv=LABELS_CSV, seed=42
|
||
)
|
||
regions = {r.region for r in list(train_ds.records) + list(val_ds.records)}
|
||
assert regions == {"spine", "hip_right", "hip_left", None}
|