develop - hack_2026
This commit is contained in:
parent
2f279d2816
commit
ddfeb01174
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
Binary file not shown.
|
|
@ -0,0 +1,176 @@
|
||||||
|
"""
|
||||||
|
Сравнение классификатора с «эвристикой области».
|
||||||
|
|
||||||
|
Отвечает на вопрос проверки: не выучила ли модель просто «где какая анатомия»
|
||||||
|
вместо «есть ли нарушение». Сравниваются два предсказателя на одних данных:
|
||||||
|
|
||||||
|
1. Модель (чекпоинт) — логит класса «нарушение».
|
||||||
|
2. Эвристика области — использует только знание анатомической области
|
||||||
|
(позвоночник = нарушение), без анализа изображения.
|
||||||
|
|
||||||
|
Если модель не лучше эвристики области, она не даёт клинической ценности.
|
||||||
|
|
||||||
|
Запуск:
|
||||||
|
python -m src.dxa.discriminator --model-path models/dxa_model.pth
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from sklearn.metrics import roc_auc_score
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||||
|
|
||||||
|
from src.dxa.dataset import build_records
|
||||||
|
from src.dxa.labels import REGIONS, QUALITY_BAD
|
||||||
|
from src.dxa.model import create_model
|
||||||
|
from src.dxa.preprocess import PreprocessConfig, preprocess_dicom
|
||||||
|
from src.dxa.train import get_device
|
||||||
|
|
||||||
|
logger = logging.getLogger("dxa.discriminator")
|
||||||
|
|
||||||
|
|
||||||
|
def model_logits(records, model_path: str, device: str, input_size: int = 224) -> np.ndarray:
|
||||||
|
"""Логиты класса «нарушение» для каждого снимка."""
|
||||||
|
model = create_model(device=device)
|
||||||
|
metadata = model.load(model_path)
|
||||||
|
model.model.eval()
|
||||||
|
|
||||||
|
cfg = PreprocessConfig.from_dict(
|
||||||
|
{**(metadata.get("preprocess") or {}), "input_size": input_size}
|
||||||
|
)
|
||||||
|
logits = []
|
||||||
|
with np.errstate(all="ignore"):
|
||||||
|
for rec in records:
|
||||||
|
arr = preprocess_dicom(rec.path, cfg)
|
||||||
|
import torch
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
out = model.model(torch.from_numpy(arr).float().unsqueeze(0).to(device))
|
||||||
|
logits.append(float(out["quality_logits"][0, 1]))
|
||||||
|
return np.array(logits)
|
||||||
|
|
||||||
|
|
||||||
|
def region_prior_scores(records) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Оценка «по области»: монотонная функция доли нарушений в области.
|
||||||
|
|
||||||
|
Позвоночник получает больший балл, потому что в обучающем наборе в нём
|
||||||
|
нарушений ~29 % против ~4–5 % у бёдер. Изображение при этом не
|
||||||
|
анализируется вообще.
|
||||||
|
"""
|
||||||
|
per_region = {r: [] for r in REGIONS}
|
||||||
|
for rec in records:
|
||||||
|
if rec.region in per_region:
|
||||||
|
per_region[rec.region].append(rec.label)
|
||||||
|
rate = {r: (np.mean(v) if v else 0.0) for r, v in per_region.items()}
|
||||||
|
return np.array([rate.get(rec.region, 0.0) for rec in records])
|
||||||
|
|
||||||
|
|
||||||
|
def evaluate(records, scores: np.ndarray) -> Dict:
|
||||||
|
"""Общие и внутриобластные метрики для одного набора оценок."""
|
||||||
|
labels = np.array([r.label for r in records])
|
||||||
|
regions = np.array([r.region for r in records])
|
||||||
|
out: Dict = {"n": len(records), "n_pos": int(labels.sum())}
|
||||||
|
|
||||||
|
if len(set(labels.tolist())) > 1:
|
||||||
|
out["roc_auc"] = float(roc_auc_score(labels, scores))
|
||||||
|
else:
|
||||||
|
out["roc_auc"] = None
|
||||||
|
|
||||||
|
per_region = {}
|
||||||
|
for region in REGIONS:
|
||||||
|
mask = regions == region
|
||||||
|
if mask.sum() == 0 or len(set(labels[mask].tolist())) < 2:
|
||||||
|
continue
|
||||||
|
per_region[region] = {
|
||||||
|
"n": int(mask.sum()),
|
||||||
|
"n_pos": int(labels[mask].sum()),
|
||||||
|
"roc_auc": float(roc_auc_score(labels[mask], scores[mask])),
|
||||||
|
}
|
||||||
|
out["per_region"] = per_region
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def run_model_comparison(model_path: str, device: Optional[str] = None) -> Dict:
|
||||||
|
"""Сравнение модели с эвристикой области на всём наборе."""
|
||||||
|
device = device or get_device()
|
||||||
|
records = build_records(
|
||||||
|
"dataset_hack", "dataset_hack/НД_для_обучения/разметка.xlsx"
|
||||||
|
)
|
||||||
|
logger.info("Device: %s, images: %d", device, len(records))
|
||||||
|
|
||||||
|
report = {
|
||||||
|
"model_path": str(model_path),
|
||||||
|
"device": device,
|
||||||
|
"n_images": len(records),
|
||||||
|
"n_violations": int(sum(r.label == QUALITY_BAD for r in records)),
|
||||||
|
"model": evaluate(records, model_logits(records, model_path, device)),
|
||||||
|
"region_prior": evaluate(records, region_prior_scores(records)),
|
||||||
|
}
|
||||||
|
|
||||||
|
model_auc = report["model"].get("roc_auc")
|
||||||
|
prior_auc = report["region_prior"].get("roc_auc")
|
||||||
|
if model_auc is not None and prior_auc is not None:
|
||||||
|
report["model_minus_region_prior_auc"] = round(model_auc - prior_auc, 4)
|
||||||
|
return report
|
||||||
|
|
||||||
|
|
||||||
|
def _format_report(report: Dict) -> str:
|
||||||
|
"""Текстовое представление отчёта сравнения."""
|
||||||
|
lines = [
|
||||||
|
"Сравнение модели с эвристикой области",
|
||||||
|
"=" * 44,
|
||||||
|
f"Снимков: {report['n_images']}, нарушений: {report['n_violations']}",
|
||||||
|
"",
|
||||||
|
f"Модель ROC-AUC: {report['model'].get('roc_auc')}",
|
||||||
|
f"Эвристика ROC-AUC: {report['region_prior'].get('roc_auc')}",
|
||||||
|
f"Преимущество модели: {report.get('model_minus_region_prior_auc')}",
|
||||||
|
"",
|
||||||
|
"Внутри областей (общий AUC здесь не помогает — область уже известна):",
|
||||||
|
]
|
||||||
|
for name, key in (("Модель", "model"), ("Эвристика области", "region_prior")):
|
||||||
|
lines.append(f" {name}:")
|
||||||
|
for region, m in (report[key].get("per_region") or {}).items():
|
||||||
|
lines.append(
|
||||||
|
f" {region:10} n={m['n']:4} нарушений={m['n_pos']:3} ROC-AUC={m['roc_auc']:.3f}"
|
||||||
|
)
|
||||||
|
lines += [
|
||||||
|
"",
|
||||||
|
"Как читать: превышение моделью эвристики области означает, что модель",
|
||||||
|
"использует содержимое снимка, а не только его анатомию. Значения внутри",
|
||||||
|
"областей показывают качество на «чистой» задаче, где подсказка по области",
|
||||||
|
"недоступна. При малом числе нарушений в области доверительный интервал",
|
||||||
|
"очень широкий.",
|
||||||
|
]
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Optional[List[str]] = None) -> int:
|
||||||
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||||
|
parser = argparse.ArgumentParser(description="Сравнение модели с эвристикой области")
|
||||||
|
parser.add_argument("--model-path", default="models/dxa_model.pth")
|
||||||
|
parser.add_argument("--device", default=None)
|
||||||
|
parser.add_argument("--json-out", default=None, help="Куда сохранить отчёт в JSON")
|
||||||
|
args = parser.parse_args(argv)
|
||||||
|
|
||||||
|
if not Path(args.model_path).exists():
|
||||||
|
logger.error("Checkpoint not found: %s", args.model_path)
|
||||||
|
return 1
|
||||||
|
|
||||||
|
report = run_model_comparison(args.model_path, args.device)
|
||||||
|
print(_format_report(report))
|
||||||
|
if args.json_out:
|
||||||
|
Path(args.json_out).write_text(json.dumps(report, ensure_ascii=False, indent=2))
|
||||||
|
logger.info("Report: %s", args.json_out)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
|
|
@ -0,0 +1,329 @@
|
||||||
|
"""
|
||||||
|
Разбор датасета DXA и построение меток качества.
|
||||||
|
|
||||||
|
Метки берутся из имён DICOM-файлов, где суффикс кодирует экспертную оценку
|
||||||
|
качества снимка. Отсутствие суффикса означает «изображение хорошее».
|
||||||
|
|
||||||
|
Приоритет меток (политика согласована с заказчиком):
|
||||||
|
_bad > _good > нет суффикса (= good)
|
||||||
|
|
||||||
|
Изображения с побайтово совпадающим содержимым считаются одним примером.
|
||||||
|
Побеждает наиболее сильное свидетельство: явная метка перекрывает неявное
|
||||||
|
«good». Без этого склейка дублей сделала бы метку одного и того же снимка
|
||||||
|
противоречивой (например, `spine_01.dcm` и `spine_4_bad.dcm` — один и тот же
|
||||||
|
снимок).
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, Iterable, List, Optional, Sequence, Tuple
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from src.dxa.preprocess import load_dicom_array
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
QUALITY_GOOD = 0
|
||||||
|
QUALITY_BAD = 1
|
||||||
|
|
||||||
|
REGIONS = ("spine", "hip_right", "hip_left")
|
||||||
|
|
||||||
|
_MARKER_RE = re.compile(r"[_\-\s](good|bad)$", re.IGNORECASE)
|
||||||
|
_BARE_MARKER_RE = re.compile(r"^(good|bad)$", re.IGNORECASE)
|
||||||
|
|
||||||
|
# Сила свидетельства: чем больше, тем приоритетнее.
|
||||||
|
_MARKER_PRIORITY = {None: 0, "good": 1, "bad": 2}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ImageRecord:
|
||||||
|
"""Один уникальный DICOM снимок с анатомической областью и меткой качества."""
|
||||||
|
|
||||||
|
path: Path
|
||||||
|
study: str
|
||||||
|
region: Optional[str]
|
||||||
|
label: int
|
||||||
|
marker: Optional[str] = None
|
||||||
|
sources: Tuple[Path, ...] = ()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def stem(self) -> str:
|
||||||
|
return self.path.stem
|
||||||
|
|
||||||
|
def as_dict(self) -> Dict:
|
||||||
|
return {
|
||||||
|
"path_to_image": str(self.path),
|
||||||
|
"dcm_path": str(self.path),
|
||||||
|
"study_uid": self.study,
|
||||||
|
"anatomical_region": self.region,
|
||||||
|
"quality": self.label,
|
||||||
|
"marker": self.marker,
|
||||||
|
"sources": [str(p) for p in self.sources],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_extension(stem: str) -> str:
|
||||||
|
"""Убрать расширение .dcm, если оно есть: функции принимают имя со stem."""
|
||||||
|
return stem[:-4] if stem.lower().endswith(".dcm") else stem
|
||||||
|
|
||||||
|
|
||||||
|
def region_from_filename(stem: str) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
Определить анатомическую область по имени файла.
|
||||||
|
|
||||||
|
Принимает как stem (`spine_01_bad`), так и имя с расширением
|
||||||
|
(`spine_01_bad.dcm`). Поддерживает встречающиеся в датасете варианты:
|
||||||
|
`spine`, `r_spine`, `Spine`, `spine-1`, `l_hip`, `l_hip-2`, `r_hip`,
|
||||||
|
`r_hop`, `r_hip03`.
|
||||||
|
"""
|
||||||
|
name = _MARKER_RE.sub("", _strip_extension(stem)).lower()
|
||||||
|
if name.startswith("spine") or name.startswith("r_spine"):
|
||||||
|
return "spine"
|
||||||
|
if name.startswith("l_hip") or name.startswith("left_hip"):
|
||||||
|
return "hip_left"
|
||||||
|
if name.startswith("r_hip") or name.startswith("r_hop") or name.startswith("right_hip"):
|
||||||
|
return "hip_right"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def marker_from_filename(stem: str) -> Optional[str]:
|
||||||
|
"""Извлечь явную метку качества (`good`/`bad`) из имени файла или stem."""
|
||||||
|
name = _strip_extension(stem)
|
||||||
|
match = _MARKER_RE.search(name)
|
||||||
|
if match:
|
||||||
|
return match.group(1).lower()
|
||||||
|
if _BARE_MARKER_RE.match(name):
|
||||||
|
return name.lower()
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _pixel_key(path: Path) -> str:
|
||||||
|
"""Хэш пиксельного содержимого: одинаковые снимки дают одинаковый ключ."""
|
||||||
|
arr = load_dicom_array(path)
|
||||||
|
return hashlib.md5(np.ascontiguousarray(arr).tobytes()).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def _discover_study_roots(data_root: Path) -> List[Path]:
|
||||||
|
"""Найти директории исследований для двух поддерживаемых раскладок."""
|
||||||
|
nested = data_root / "НД_для_обучения" / "Исследования"
|
||||||
|
if nested.is_dir():
|
||||||
|
return sorted(p for p in nested.iterdir() if p.is_dir())
|
||||||
|
|
||||||
|
direct = data_root / "Исследования"
|
||||||
|
if direct.is_dir():
|
||||||
|
return sorted(p for p in direct.iterdir() if p.is_dir())
|
||||||
|
|
||||||
|
# Раскладка без общего каталога исследований: корень сам содержит DICOM.
|
||||||
|
return [data_root]
|
||||||
|
|
||||||
|
|
||||||
|
def scan_dataset(
|
||||||
|
data_root: str | Path,
|
||||||
|
with_pixel_dedup: bool = True,
|
||||||
|
annotation_path: Optional[str | Path] = None,
|
||||||
|
) -> List[ImageRecord]:
|
||||||
|
"""
|
||||||
|
Собрать уникальные примеры из каталога датасета.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
data_root: корень датасета (`dataset_hack`, `dataset_hack/НД_для_обучения`
|
||||||
|
или каталог с DICOM).
|
||||||
|
with_pixel_dedup: склеивать снимки с одинаковым пиксельным содержимым.
|
||||||
|
annotation_path: необязательный Excel-файл разметки; используется только
|
||||||
|
для предупреждения о расхождении с метками из имён файлов.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Список ImageRecord, по одному на уникальный снимок.
|
||||||
|
"""
|
||||||
|
root = Path(data_root)
|
||||||
|
if not root.exists():
|
||||||
|
raise FileNotFoundError(f"Dataset root not found: {root}")
|
||||||
|
|
||||||
|
studies = _discover_study_roots(root)
|
||||||
|
grouped: Dict[str, List[Dict]] = {}
|
||||||
|
|
||||||
|
for study_dir in studies:
|
||||||
|
dcm_files = sorted(study_dir.rglob("*.dcm"))
|
||||||
|
for dcm in dcm_files:
|
||||||
|
marker = marker_from_filename(dcm.stem)
|
||||||
|
region = region_from_filename(dcm.stem)
|
||||||
|
if region is None and marker is None:
|
||||||
|
logger.warning("Cannot determine region for %s; keeping as unlabeled", dcm)
|
||||||
|
try:
|
||||||
|
key = _pixel_key(dcm) if with_pixel_dedup else str(dcm)
|
||||||
|
except Exception as exc: # битый DICOM не должен ронять разбор датасета
|
||||||
|
logger.error("Failed to read %s: %s", dcm, exc)
|
||||||
|
continue
|
||||||
|
grouped.setdefault(key, []).append(
|
||||||
|
{"path": dcm, "study": study_dir.name, "region": region, "marker": marker}
|
||||||
|
)
|
||||||
|
|
||||||
|
records: List[ImageRecord] = []
|
||||||
|
for members in grouped.values():
|
||||||
|
marker = max((m["marker"] for m in members), key=lambda m: _MARKER_PRIORITY[m])
|
||||||
|
label = QUALITY_BAD if marker == "bad" else QUALITY_GOOD
|
||||||
|
|
||||||
|
regions = {m["region"] for m in members if m["region"]}
|
||||||
|
if len(regions) > 1:
|
||||||
|
logger.warning("Conflicting regions %s for %s", regions, [m["path"].name for m in members])
|
||||||
|
region = regions.pop() if len(regions) == 1 else None
|
||||||
|
|
||||||
|
# Представитель: сначала по силе метки, затем по имени файла.
|
||||||
|
members.sort(key=lambda m: (-_MARKER_PRIORITY[m["marker"]], str(m["path"])))
|
||||||
|
best = members[0]
|
||||||
|
|
||||||
|
records.append(
|
||||||
|
ImageRecord(
|
||||||
|
path=best["path"],
|
||||||
|
study=best["study"],
|
||||||
|
region=region,
|
||||||
|
label=label,
|
||||||
|
marker=marker,
|
||||||
|
sources=tuple(m["path"] for m in members),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
records.sort(key=lambda r: (r.study, str(r.path)))
|
||||||
|
if annotation_path:
|
||||||
|
warn_annotation_mismatch(records, annotation_path)
|
||||||
|
return records
|
||||||
|
|
||||||
|
|
||||||
|
def warn_annotation_mismatch(records: Sequence[ImageRecord], annotation_path: str | Path) -> None:
|
||||||
|
"""
|
||||||
|
Сравнить метки из имён файлов с разметкой Excel и предупредить о расхождениях.
|
||||||
|
|
||||||
|
Метки Excel относятся к исследованию целиком, поэтому расхождения —
|
||||||
|
ожидаемое явление: Excel фиксирует нарушения, не отражённые суффиксом
|
||||||
|
имени файла.
|
||||||
|
"""
|
||||||
|
import pandas as pd
|
||||||
|
|
||||||
|
try:
|
||||||
|
raw = pd.read_excel(annotation_path, header=None).iloc[2:].copy()
|
||||||
|
raw.columns = range(len(raw.columns))
|
||||||
|
raw = raw.dropna(subset=[1])
|
||||||
|
totals = {
|
||||||
|
str(row[1]).strip(): {"spine": row[9], "hip_right": row[10], "hip_left": row[11]}
|
||||||
|
for _, row in raw.iterrows()
|
||||||
|
}
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Could not read annotation %s: %s", annotation_path, exc)
|
||||||
|
return
|
||||||
|
|
||||||
|
mismatched = 0
|
||||||
|
for rec in records:
|
||||||
|
values = totals.get(rec.study)
|
||||||
|
if not values or rec.region is None:
|
||||||
|
continue
|
||||||
|
excel_value = values.get(rec.region)
|
||||||
|
if pd.isna(excel_value):
|
||||||
|
continue
|
||||||
|
if int(excel_value) != rec.label:
|
||||||
|
mismatched += 1
|
||||||
|
|
||||||
|
if mismatched:
|
||||||
|
logger.warning(
|
||||||
|
"%d/%d images have Excel quality different from filename marker; "
|
||||||
|
"filename markers take precedence",
|
||||||
|
mismatched,
|
||||||
|
len(records),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def label_summary(records: Sequence[ImageRecord]) -> Dict[str, Dict[str, int]]:
|
||||||
|
"""Сводка распределения примеров по областям и классам."""
|
||||||
|
summary: Dict[str, Dict[str, int]] = {
|
||||||
|
region: {"good": 0, "bad": 0, "total": 0} for region in REGIONS
|
||||||
|
}
|
||||||
|
summary["unknown"] = {"good": 0, "bad": 0, "total": 0}
|
||||||
|
for rec in records:
|
||||||
|
bucket = summary[rec.region if rec.region in REGIONS else "unknown"]
|
||||||
|
bucket["bad" if rec.label == QUALITY_BAD else "good"] += 1
|
||||||
|
bucket["total"] += 1
|
||||||
|
return summary
|
||||||
|
|
||||||
|
|
||||||
|
def format_summary(records: Sequence[ImageRecord]) -> str:
|
||||||
|
"""Человекочитаемая сводка датасета для логов обучения."""
|
||||||
|
summary = label_summary(records)
|
||||||
|
lines = [f"Unique images: {len(records)}"]
|
||||||
|
total_bad = total_good = 0
|
||||||
|
for region, counts in summary.items():
|
||||||
|
if counts["total"] == 0:
|
||||||
|
continue
|
||||||
|
total_bad += counts["bad"]
|
||||||
|
total_good += counts["good"]
|
||||||
|
lines.append(
|
||||||
|
f" {region:10} good={counts['good']:4} bad={counts['bad']:4} total={counts['total']:4}"
|
||||||
|
)
|
||||||
|
share = total_bad / max(total_bad + total_good, 1)
|
||||||
|
lines.append(f" {'ALL':10} good={total_good:4} bad={total_bad:4} (bad share={share:.1%})")
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def stratified_group_split(
|
||||||
|
records: Sequence[ImageRecord],
|
||||||
|
val_fraction: float = 0.2,
|
||||||
|
seed: int = 42,
|
||||||
|
) -> Tuple[List[ImageRecord], List[ImageRecord]]:
|
||||||
|
"""
|
||||||
|
Разделить примеры на train/val без утечки между исследованиями.
|
||||||
|
|
||||||
|
Разбиение выполняется по исследованиям, а не по отдельным снимкам: снимки
|
||||||
|
одного исследования всегда попадают в одну часть. Внутри каждой страты
|
||||||
|
(есть нарушение / только хорошие) исследования тасуются детерминированно,
|
||||||
|
затем отбирается нужная доля снимков.
|
||||||
|
"""
|
||||||
|
if not records:
|
||||||
|
return [], []
|
||||||
|
if not 0.0 < val_fraction < 1.0:
|
||||||
|
raise ValueError(f"val_fraction must be in (0, 1), got {val_fraction}")
|
||||||
|
|
||||||
|
by_study: Dict[str, List[ImageRecord]] = {}
|
||||||
|
for rec in records:
|
||||||
|
by_study.setdefault(rec.study, []).append(rec)
|
||||||
|
|
||||||
|
strata: Dict[bool, List[str]] = {True: [], False: []}
|
||||||
|
for study, items in by_study.items():
|
||||||
|
strata[any(r.label == QUALITY_BAD for r in items)].append(study)
|
||||||
|
|
||||||
|
rng = np.random.default_rng(seed)
|
||||||
|
val_studies: set[str] = set()
|
||||||
|
for has_violation, studies in strata.items():
|
||||||
|
studies.sort()
|
||||||
|
order = rng.permutation(len(studies))
|
||||||
|
n_images = sum(len(by_study[studies[i]]) for i in range(len(studies)))
|
||||||
|
target = int(round(n_images * val_fraction))
|
||||||
|
taken = 0
|
||||||
|
for idx in order:
|
||||||
|
if taken >= target and len(val_studies) > 0:
|
||||||
|
break
|
||||||
|
study = studies[idx]
|
||||||
|
val_studies.add(study)
|
||||||
|
taken += len(by_study[study])
|
||||||
|
|
||||||
|
train = [r for r in records if r.study not in val_studies]
|
||||||
|
val = [r for r in records if r.study in val_studies]
|
||||||
|
|
||||||
|
if not train or not val:
|
||||||
|
logger.warning("Grouped split degenerated; falling back to single-study val split")
|
||||||
|
ordered = sorted(records, key=lambda r: (r.study, str(r.path)))
|
||||||
|
cutoff = max(1, int(len(ordered) * (1 - val_fraction)))
|
||||||
|
return ordered[:cutoff], ordered[cutoff:]
|
||||||
|
|
||||||
|
return train, val
|
||||||
|
|
||||||
|
|
||||||
|
def iter_dicom_files(data_root: str | Path) -> Iterable[Path]:
|
||||||
|
"""Перечислить все DICOM файлы датасета (для пакетной обработки)."""
|
||||||
|
root = Path(data_root)
|
||||||
|
if root.is_file():
|
||||||
|
yield root
|
||||||
|
return
|
||||||
|
yield from sorted(root.rglob("*.dcm"))
|
||||||
|
|
@ -0,0 +1,152 @@
|
||||||
|
"""
|
||||||
|
Общая предобработка DICOM DXA изображений.
|
||||||
|
|
||||||
|
Модуль используется и при обучении, и при инференсе, чтобы вход модели был
|
||||||
|
идентичным в обоих режимах. Параметры предобработки хранятся в чекпоинте
|
||||||
|
(config['preprocess']) и восстанавливаются при загрузке модели.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import asdict, dataclass, replace
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pydicom
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
NORM_CHOICES = ("percentile", "minmax")
|
||||||
|
|
||||||
|
# Нормализация ImageNet: предобученный backbone ожидает именно такой вход.
|
||||||
|
_IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
|
||||||
|
_IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PreprocessConfig:
|
||||||
|
"""Параметры предобработки одного DICOM изображения."""
|
||||||
|
|
||||||
|
norm: str = "percentile"
|
||||||
|
p_low: float = 0.5
|
||||||
|
p_high: float = 99.5
|
||||||
|
imagenet_norm: bool = True
|
||||||
|
input_size: int = 224
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if self.norm not in NORM_CHOICES:
|
||||||
|
raise ValueError(f"norm must be one of {NORM_CHOICES}, got {self.norm!r}")
|
||||||
|
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
return asdict(self)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, cfg: Dict[str, Any]) -> "PreprocessConfig":
|
||||||
|
known = {f: cfg[f] for f in cls.__dataclass_fields__ if f in cfg}
|
||||||
|
return cls(**known)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def legacy(cls) -> "PreprocessConfig":
|
||||||
|
"""Предобработка старых чекпоинтов: min-max без нормировки ImageNet."""
|
||||||
|
return cls(norm="minmax", imagenet_norm=False)
|
||||||
|
|
||||||
|
|
||||||
|
def load_dicom_array(path: Union[str, Path]) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Прочитать DICOM и вернуть 2D float32 массив без нормализации.
|
||||||
|
|
||||||
|
Учитывает RescaleSlope/RescaleIntercept и MONOCHROME1; многокадровые
|
||||||
|
изображения усредняются по кадрам.
|
||||||
|
"""
|
||||||
|
ds = pydicom.dcmread(str(path))
|
||||||
|
arr = ds.pixel_array.astype(np.float32)
|
||||||
|
|
||||||
|
if arr.ndim == 3:
|
||||||
|
arr = arr.mean(axis=0) if arr.shape[0] > 1 else arr[0]
|
||||||
|
if arr.ndim != 2:
|
||||||
|
raise ValueError(f"Unsupported pixel_array shape {arr.shape} for {path}")
|
||||||
|
|
||||||
|
slope = float(getattr(ds, "RescaleSlope", 1) or 1)
|
||||||
|
intercept = float(getattr(ds, "RescaleIntercept", 0) or 0)
|
||||||
|
if slope != 1.0 or intercept != 0.0:
|
||||||
|
arr = arr * slope + intercept
|
||||||
|
|
||||||
|
if str(getattr(ds, "PhotometricInterpretation", "")).upper() == "MONOCHROME1":
|
||||||
|
arr = arr.max() - arr
|
||||||
|
|
||||||
|
return np.ascontiguousarray(arr, dtype=np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_array(arr: np.ndarray, cfg: PreprocessConfig) -> np.ndarray:
|
||||||
|
"""Привести интенсивности к [0, 1] по стратегии cfg.norm."""
|
||||||
|
if cfg.norm == "percentile":
|
||||||
|
lo, hi = np.percentile(arr, [cfg.p_low, cfg.p_high])
|
||||||
|
else:
|
||||||
|
lo, hi = float(arr.min()), float(arr.max())
|
||||||
|
|
||||||
|
if hi <= lo:
|
||||||
|
hi = lo + 1.0
|
||||||
|
return np.clip((arr - lo) / (hi - lo), 0.0, 1.0).astype(np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def to_model_tensor(arr01: np.ndarray, cfg: PreprocessConfig) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Преобразовать нормализованное изображение [0,1] в CHW float32 тензор.
|
||||||
|
|
||||||
|
Порядок операций (resize в uint8) совпадает с историческим кодом проекта,
|
||||||
|
чтобы не менять распределение входа относительно прошлых моделей.
|
||||||
|
"""
|
||||||
|
size = int(cfg.input_size)
|
||||||
|
img = np.clip(arr01 * 255.0, 0, 255).astype(np.uint8)
|
||||||
|
pil = Image.fromarray(np.stack([img] * 3, axis=-1)).resize((size, size), Image.BILINEAR)
|
||||||
|
chw = np.asarray(pil, dtype=np.float32).transpose(2, 0, 1) / 255.0
|
||||||
|
|
||||||
|
if cfg.imagenet_norm:
|
||||||
|
chw = (chw - _IMAGENET_MEAN[:, None, None]) / _IMAGENET_STD[:, None, None]
|
||||||
|
|
||||||
|
return np.ascontiguousarray(chw, dtype=np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess_dicom(path: Union[str, Path], cfg: PreprocessConfig) -> np.ndarray:
|
||||||
|
"""Полный конвейер: DICOM -> нормализованный CHW массив."""
|
||||||
|
return to_model_tensor(normalize_array(load_dicom_array(path), cfg), cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def load_array_from_bytes(dcm_bytes: bytes) -> "np.ndarray":
|
||||||
|
"""
|
||||||
|
Прочитать DICOM из памяти и вернуть 2D float32 массив.
|
||||||
|
|
||||||
|
Нужен HTTP API: он получает файл потоком и не имеет пути на диске.
|
||||||
|
Повторяет коррекции `load_dicom_array` (RescaleSlope/Intercept, MONOCHROME1).
|
||||||
|
"""
|
||||||
|
import io
|
||||||
|
|
||||||
|
ds = pydicom.dcmread(io.BytesIO(dcm_bytes))
|
||||||
|
arr = ds.pixel_array.astype(np.float32)
|
||||||
|
if arr.ndim == 3:
|
||||||
|
arr = arr.mean(axis=0) if arr.shape[0] > 1 else arr[0]
|
||||||
|
if arr.ndim != 2:
|
||||||
|
raise ValueError(f"Unsupported pixel_array shape {arr.shape}")
|
||||||
|
|
||||||
|
slope = float(getattr(ds, "RescaleSlope", 1) or 1)
|
||||||
|
intercept = float(getattr(ds, "RescaleIntercept", 0) or 0)
|
||||||
|
if slope != 1.0 or intercept != 0.0:
|
||||||
|
arr = arr * slope + intercept
|
||||||
|
|
||||||
|
if str(getattr(ds, "PhotometricInterpretation", "")).upper() == "MONOCHROME1":
|
||||||
|
arr = arr.max() - arr
|
||||||
|
|
||||||
|
return np.ascontiguousarray(arr, dtype=np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess_from_array(arr: np.ndarray, cfg: PreprocessConfig) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Конвейер для уже прочитанного массива пикселей.
|
||||||
|
|
||||||
|
Используется HTTP API, который принимает DICOM потоком и не имеет пути
|
||||||
|
к файлу. Результат идентичен `preprocess_dicom`.
|
||||||
|
"""
|
||||||
|
return to_model_tensor(normalize_array(arr.astype(np.float32), cfg), cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def with_input_size(cfg: PreprocessConfig, input_size: int) -> PreprocessConfig:
|
||||||
|
return replace(cfg, input_size=int(input_size))
|
||||||
|
|
@ -0,0 +1,163 @@
|
||||||
|
/**
|
||||||
|
* Проверка панели деталей веб-интерфейса в реальном браузере.
|
||||||
|
*
|
||||||
|
* Загружает несколько DICOM-файлов, кликает по строкам таблицы и снимает
|
||||||
|
* значения из панели деталей, чтобы убедиться, что они меняются при
|
||||||
|
* переключении строк (это был основной симптом бага: панель показывала одни
|
||||||
|
* и те же значения).
|
||||||
|
*
|
||||||
|
* Требуется запущенный сервер:
|
||||||
|
* python -m uvicorn src.main:app --port 8123
|
||||||
|
* node tests/browser/ui_check.js
|
||||||
|
*
|
||||||
|
* Переменные окружения:
|
||||||
|
* BASE_URL адрес сервера (по умолчанию http://127.0.0.1:8123)
|
||||||
|
* PLAYWRIGHT путь к playwright-core (по умолчанию берётся из bundled
|
||||||
|
* Browser Use в установке Qwen Code)
|
||||||
|
*
|
||||||
|
* Используется временный профиль Chrome через launchPersistentContext,
|
||||||
|
* профиль пользователя не затрагивается.
|
||||||
|
*/
|
||||||
|
const fs = require('fs');
|
||||||
|
const path = require('path');
|
||||||
|
const os = require('os');
|
||||||
|
|
||||||
|
function resolvePlaywright() {
|
||||||
|
if (process.env.PLAYWRIGHT) return process.env.PLAYWRIGHT;
|
||||||
|
// Ищем playwright-core в bundled Browser Use установки Qwen Code.
|
||||||
|
const base = path.join(os.homedir(), '.qwen', 'updates', 'npm');
|
||||||
|
if (!fs.existsSync(base)) return 'playwright-core';
|
||||||
|
const versions = fs.readdirSync(path.join(base, fs.readdirSync(base)[0], 'versions')).sort();
|
||||||
|
for (const v of versions.reverse()) {
|
||||||
|
const candidate = path.join(
|
||||||
|
base, fs.readdirSync(base)[0], 'versions', v,
|
||||||
|
'node_modules/@qwen-code/qwen-code/bundled/browser-use/runtime/node_modules/playwright-core'
|
||||||
|
);
|
||||||
|
if (fs.existsSync(candidate)) return candidate;
|
||||||
|
}
|
||||||
|
return 'playwright-core';
|
||||||
|
}
|
||||||
|
|
||||||
|
const { chromium } = require(resolvePlaywright());
|
||||||
|
|
||||||
|
const BASE = process.env.BASE_URL || 'http://127.0.0.1:8123';
|
||||||
|
const REPO = path.resolve(__dirname, '..', '..');
|
||||||
|
const FILES = [
|
||||||
|
path.join(REPO, 'dataset_hack/Для теста/spine.dcm'),
|
||||||
|
path.join(REPO, 'dataset_hack/Для теста/l_hip.dcm'),
|
||||||
|
path.join(REPO, 'dataset_hack/Для теста/r_hip.dcm'),
|
||||||
|
];
|
||||||
|
|
||||||
|
async function panelState(page) {
|
||||||
|
return page.evaluate(() => {
|
||||||
|
const txt = (id) => {
|
||||||
|
const el = document.getElementById(id);
|
||||||
|
return el ? el.innerText.replace(/\s+/g, ' ').trim() : null;
|
||||||
|
};
|
||||||
|
const cards = (id) => {
|
||||||
|
const el = document.getElementById(id);
|
||||||
|
if (!el) return null;
|
||||||
|
return Array.from(el.children).map((c) => c.innerText.replace(/\s+/g, ' ').trim());
|
||||||
|
};
|
||||||
|
const imgs = Array.from(document.querySelectorAll('#detailImage img, #detailMask img'))
|
||||||
|
.map((i) => ({ src: (i.getAttribute('src') || '').slice(0, 22), w: i.naturalWidth, h: i.naturalHeight }));
|
||||||
|
return {
|
||||||
|
detailVisible: !document.getElementById('detailSection').classList.contains('hidden'),
|
||||||
|
badge: txt('detailBadge'),
|
||||||
|
basic: txt('detailBasic'),
|
||||||
|
quality: txt('detailQuality'),
|
||||||
|
measurements: cards('detailMotion'),
|
||||||
|
roiCards: cards('detailArtifacts'),
|
||||||
|
spineHidden: document.getElementById('spineSection').classList.contains('hidden'),
|
||||||
|
spine: txt('detailSpine'),
|
||||||
|
hipHidden: document.getElementById('hipSection').classList.contains('hidden'),
|
||||||
|
hip: txt('detailHip'),
|
||||||
|
reason: txt('detailReason'),
|
||||||
|
note: txt('detailMetricsNote'),
|
||||||
|
images: imgs,
|
||||||
|
};
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
(async () => {
|
||||||
|
const userDataDir = fs.mkdtempSync(path.join(require('os').tmpdir(), 'dxa-ui-'));
|
||||||
|
// launchPersistentContext с временным профилем: профиль пользователя не затрагивается.
|
||||||
|
const context = await chromium.launchPersistentContext(userDataDir, { channel: 'chrome', headless: true });
|
||||||
|
const browser = context;
|
||||||
|
const page = context.pages()[0] || await context.newPage();
|
||||||
|
await page.setViewportSize({ width: 1400, height: 1100 });
|
||||||
|
|
||||||
|
const consoleErrors = [];
|
||||||
|
page.on('console', (m) => { if (m.type() === 'error') consoleErrors.push(m.text()); });
|
||||||
|
page.on('pageerror', (e) => consoleErrors.push('PAGEERROR ' + e.message));
|
||||||
|
|
||||||
|
await page.goto(BASE, { waitUntil: 'domcontentloaded' });
|
||||||
|
console.log('title:', await page.title());
|
||||||
|
console.log('status banner:', (await page.locator('#statusText').innerText()).trim());
|
||||||
|
|
||||||
|
// Загружаем файлы через input (drag-and-drop эмулировать не нужно).
|
||||||
|
await page.setInputFiles('#fileInput', FILES);
|
||||||
|
await page.waitForFunction(
|
||||||
|
() => document.querySelectorAll('#resultsTable tr').length === 3,
|
||||||
|
{ timeout: 60000 }
|
||||||
|
);
|
||||||
|
console.log('rows rendered:', await page.locator('#resultsTable tr').count());
|
||||||
|
console.log('stats:', await page.evaluate(() => ({
|
||||||
|
total: document.getElementById('totalCount').innerText,
|
||||||
|
ok: document.getElementById('okCount').innerText,
|
||||||
|
violation: document.getElementById('violationCount').innerText,
|
||||||
|
avg: document.getElementById('accuracy').innerText,
|
||||||
|
})));
|
||||||
|
|
||||||
|
const shots = [];
|
||||||
|
const states = [];
|
||||||
|
for (const idx of [0, 1, 2]) {
|
||||||
|
await page.locator('#resultsTable tr').nth(idx).click();
|
||||||
|
await page.waitForTimeout(1200);
|
||||||
|
const state = await panelState(page);
|
||||||
|
states.push(state);
|
||||||
|
const file = path.join(os.tmpdir(), `dxa_panel_row${idx}.png`);
|
||||||
|
await page.screenshot({ path: file, fullPage: false });
|
||||||
|
shots.push({ idx, file, measurements: state.measurements, region: state.basic });
|
||||||
|
// Прокручиваем к панели, чтобы она попала в кадр
|
||||||
|
await page.locator('#detailSection').scrollIntoViewIfNeeded();
|
||||||
|
await page.screenshot({ path: path.join(os.tmpdir(), `dxa_detail_row${idx}.png`) });
|
||||||
|
}
|
||||||
|
|
||||||
|
console.log('\n=== ПАНЕЛЬ ДЕТАЛЕЙ ПО СТРОКАМ ===');
|
||||||
|
for (const s of shots) {
|
||||||
|
console.log(`\nстрока ${s.idx}:`);
|
||||||
|
console.log(' basic:', s.region);
|
||||||
|
console.log(' измерения:', s.measurements);
|
||||||
|
}
|
||||||
|
|
||||||
|
console.log('\n=== СЕКЦИИ / ВИЗУАЛИЗАЦИЯ / ЗАКЛЮЧЕНИЕ ===');
|
||||||
|
states.forEach((s, i) => {
|
||||||
|
console.log(`\nстрока ${i}: visible=${s.detailVisible}`);
|
||||||
|
console.log(' badge:', s.badge);
|
||||||
|
console.log(' quality:', s.quality);
|
||||||
|
console.log(' roi:', s.roiCards);
|
||||||
|
console.log(' spine hidden:', s.spineHidden, '| hip hidden:', s.hipHidden);
|
||||||
|
if (!s.spineHidden) console.log(' spine:', s.spine);
|
||||||
|
if (!s.hipHidden) console.log(' hip:', s.hip);
|
||||||
|
console.log(' note:', (s.note || '').slice(0, 80));
|
||||||
|
console.log(' images:', s.images);
|
||||||
|
});
|
||||||
|
|
||||||
|
// Проверка, что значения действительно различаются между строками
|
||||||
|
const uniq = (arr) => new Set(arr.filter(Boolean)).size;
|
||||||
|
console.log('\n=== РАЗЛИЧИМОСТЬ ЗНАЧЕНИЙ МЕЖДУ СТРОКАМИ ===');
|
||||||
|
console.log('разных basic:', uniq(states.map((s) => s.basic)));
|
||||||
|
console.log('разных measurements:', uniq(states.map((s) => JSON.stringify(s.measurements))));
|
||||||
|
console.log('разных quality:', uniq(states.map((s) => s.quality)));
|
||||||
|
console.log('разных roi:', uniq(states.map((s) => JSON.stringify(s.roiCards))));
|
||||||
|
console.log('разных reason:', uniq(states.map((s) => s.reason)));
|
||||||
|
|
||||||
|
console.log('\nconsole errors:', consoleErrors.length ? consoleErrors : 'none');
|
||||||
|
fs.writeFileSync(path.join(os.tmpdir(), 'dxa_ui_states.json'), JSON.stringify(states, null, 2));
|
||||||
|
console.log('states saved to', path.join(os.tmpdir(), 'dxa_ui_states.json'));
|
||||||
|
console.log('screenshots:', shots.map((s) => s.file).join(', '));
|
||||||
|
|
||||||
|
await browser.close();
|
||||||
|
fs.rmSync(userDataDir, { recursive: true, force: true });
|
||||||
|
})().catch((e) => { console.error('FAILED:', e.message); process.exit(1); });
|
||||||
|
|
@ -0,0 +1,100 @@
|
||||||
|
/**
|
||||||
|
* Проверка работы интерфейса в офлайне.
|
||||||
|
*
|
||||||
|
* Блокирует все запросы к внешним хостам: если оформление или иконки
|
||||||
|
* подгружаются с CDN, они отвалятся. Успешное выполнение означает, что
|
||||||
|
* страница самодостаточна и работает в контейнере без сети.
|
||||||
|
*
|
||||||
|
* Требуется запущенный сервер:
|
||||||
|
* python -m uvicorn src.main:app --port 8123
|
||||||
|
* node tests/browser/ui_offline.js
|
||||||
|
*/
|
||||||
|
const fs = require('fs');
|
||||||
|
const path = require('path');
|
||||||
|
const os = require('os');
|
||||||
|
|
||||||
|
function resolvePlaywright() {
|
||||||
|
if (process.env.PLAYWRIGHT) return process.env.PLAYWRIGHT;
|
||||||
|
const base = path.join(os.homedir(), '.qwen', 'updates', 'npm');
|
||||||
|
if (!fs.existsSync(base)) return 'playwright-core';
|
||||||
|
const versions = fs.readdirSync(path.join(base, fs.readdirSync(base)[0], 'versions')).sort();
|
||||||
|
for (const v of versions.reverse()) {
|
||||||
|
const candidate = path.join(
|
||||||
|
base, fs.readdirSync(base)[0], 'versions', v,
|
||||||
|
'node_modules/@qwen-code/qwen-code/bundled/browser-use/runtime/node_modules/playwright-core'
|
||||||
|
);
|
||||||
|
if (fs.existsSync(candidate)) return candidate;
|
||||||
|
}
|
||||||
|
return 'playwright-core';
|
||||||
|
}
|
||||||
|
|
||||||
|
const { chromium } = require(resolvePlaywright());
|
||||||
|
|
||||||
|
const BASE = process.env.BASE_URL || 'http://127.0.0.1:8123';
|
||||||
|
const REPO = path.resolve(__dirname, '..', '..');
|
||||||
|
const FILE = path.join(REPO, 'dataset_hack/Для теста/spine.dcm');
|
||||||
|
|
||||||
|
(async () => {
|
||||||
|
const userDataDir = fs.mkdtempSync(path.join(os.tmpdir(), 'dxa-offline-'));
|
||||||
|
const context = await chromium.launchPersistentContext(userDataDir, { channel: 'chrome', headless: true });
|
||||||
|
const page = context.pages()[0] || await context.newPage();
|
||||||
|
await page.setViewportSize({ width: 1400, height: 1100 });
|
||||||
|
|
||||||
|
const external = [];
|
||||||
|
const failed = [];
|
||||||
|
await context.route('**/*', (route) => {
|
||||||
|
const url = route.request().url();
|
||||||
|
const isLocal = url.startsWith(BASE) || url.startsWith('data:') || url.startsWith('blob:');
|
||||||
|
if (!isLocal) {
|
||||||
|
external.push(url);
|
||||||
|
return route.abort();
|
||||||
|
}
|
||||||
|
return route.continue();
|
||||||
|
});
|
||||||
|
page.on('requestfailed', (r) => failed.push(r.url()));
|
||||||
|
|
||||||
|
await page.goto(BASE, { waitUntil: 'domcontentloaded' });
|
||||||
|
await page.waitForTimeout(1500);
|
||||||
|
|
||||||
|
// Оформление применилось? Проверяем вычисленные стили, а не наличие файла.
|
||||||
|
const styling = await page.evaluate(() => {
|
||||||
|
const drop = document.getElementById('dropZone');
|
||||||
|
const cs = getComputedStyle(drop);
|
||||||
|
const icon = document.querySelector('i.fas');
|
||||||
|
const iconFont = icon ? getComputedStyle(icon).fontFamily : null;
|
||||||
|
const iconWidth = icon ? icon.getBoundingClientRect().width : 0;
|
||||||
|
return {
|
||||||
|
dropZoneBorderRadius: cs.borderRadius,
|
||||||
|
dropZoneDisplay: cs.display,
|
||||||
|
iconFontFamily: iconFont,
|
||||||
|
iconWidth,
|
||||||
|
};
|
||||||
|
});
|
||||||
|
console.log('styling:', styling);
|
||||||
|
|
||||||
|
// Загрузка файла и открытие панели деталей в офлайне
|
||||||
|
await page.setInputFiles('#fileInput', [FILE]);
|
||||||
|
await page.waitForFunction(() => document.querySelectorAll('#resultsTable tr').length === 1, { timeout: 60000 });
|
||||||
|
await page.locator('#resultsTable tr').first().click();
|
||||||
|
await page.waitForTimeout(1200);
|
||||||
|
const panel = await page.evaluate(() => ({
|
||||||
|
visible: !document.getElementById('detailSection').classList.contains('hidden'),
|
||||||
|
basic: (document.getElementById('detailBasic') || {}).innerText?.replace(/\s+/g, ' ').trim(),
|
||||||
|
reason: (document.getElementById('detailReason') || {}).innerText?.replace(/\s+/g, ' ').trim().slice(0, 70),
|
||||||
|
images: document.querySelectorAll('#detailImage img, #detailMask img').length,
|
||||||
|
}));
|
||||||
|
console.log('panel:', panel);
|
||||||
|
|
||||||
|
console.log('\nвнешние запросы (заблокированы):', external.length ? external : 'нет');
|
||||||
|
console.log('проваленные запросы:', failed.length ? failed : 'нет');
|
||||||
|
|
||||||
|
const ok =
|
||||||
|
styling.dropZoneBorderRadius !== '0px' &&
|
||||||
|
panel.visible && panel.images === 2 &&
|
||||||
|
external.length === 0;
|
||||||
|
console.log('\nОФЛАЙН-ПРОВЕРКА:', ok ? 'ПРОЙДЕНА' : 'НЕ ПРОЙДЕНА');
|
||||||
|
|
||||||
|
await context.close();
|
||||||
|
fs.rmSync(userDataDir, { recursive: true, force: true });
|
||||||
|
process.exit(ok ? 0 : 1);
|
||||||
|
})().catch((e) => { console.error('FAILED:', e.message); process.exit(1); });
|
||||||
|
|
@ -0,0 +1,96 @@
|
||||||
|
/**
|
||||||
|
* Проверка панели деталей на снимках с нарушением качества.
|
||||||
|
*
|
||||||
|
* Все файлы из `Для теста` модель считает качественными, поэтому ветка
|
||||||
|
* «нарушение» (бейдж, overall_quality=POOR, severity=HIGH, текст заключения)
|
||||||
|
* иначе остаётся непроверенной.
|
||||||
|
*
|
||||||
|
* Требуется запущенный сервер:
|
||||||
|
* python -m uvicorn src.main:app --port 8123
|
||||||
|
* node tests/browser/ui_violation.js
|
||||||
|
*
|
||||||
|
* Переменные окружения: BASE_URL, PLAYWRIGHT (см. ui_check.js).
|
||||||
|
*/
|
||||||
|
const fs = require('fs');
|
||||||
|
const path = require('path');
|
||||||
|
const os = require('os');
|
||||||
|
|
||||||
|
function resolvePlaywright() {
|
||||||
|
if (process.env.PLAYWRIGHT) return process.env.PLAYWRIGHT;
|
||||||
|
const base = path.join(os.homedir(), '.qwen', 'updates', 'npm');
|
||||||
|
if (!fs.existsSync(base)) return 'playwright-core';
|
||||||
|
const versions = fs.readdirSync(path.join(base, fs.readdirSync(base)[0], 'versions')).sort();
|
||||||
|
for (const v of versions.reverse()) {
|
||||||
|
const candidate = path.join(
|
||||||
|
base, fs.readdirSync(base)[0], 'versions', v,
|
||||||
|
'node_modules/@qwen-code/qwen-code/bundled/browser-use/runtime/node_modules/playwright-core'
|
||||||
|
);
|
||||||
|
if (fs.existsSync(candidate)) return candidate;
|
||||||
|
}
|
||||||
|
return 'playwright-core';
|
||||||
|
}
|
||||||
|
|
||||||
|
const { chromium } = require(resolvePlaywright());
|
||||||
|
|
||||||
|
const BASE = process.env.BASE_URL || 'http://127.0.0.1:8123';
|
||||||
|
const REPO = path.resolve(__dirname, '..', '..');
|
||||||
|
const STUDY = 'dataset_hack/НД_для_обучения/Исследования';
|
||||||
|
// Два размеченных нарушения из обучающего набора (spine и hip_right).
|
||||||
|
const VIOLATIONS = [
|
||||||
|
path.join(REPO, STUDY, '2.25.11175860580562939441493697221497597640/series_003_2_CR/A2507431060 DXA/CR DXA/spine_01_bad.dcm'),
|
||||||
|
path.join(REPO, STUDY, '2.25.12798473087614376830819854616447908612/series_002_2_CR/A2507538847 DXA/CR DXA/r_hip_01_bad.dcm'),
|
||||||
|
];
|
||||||
|
|
||||||
|
(async () => {
|
||||||
|
const userDataDir = fs.mkdtempSync(path.join(require('os').tmpdir(), 'dxa-ui2-'));
|
||||||
|
const context = await chromium.launchPersistentContext(userDataDir, { channel: 'chrome', headless: true });
|
||||||
|
const page = context.pages()[0] || await context.newPage();
|
||||||
|
await page.setViewportSize({ width: 1400, height: 1200 });
|
||||||
|
|
||||||
|
const errors = [];
|
||||||
|
page.on('pageerror', (e) => errors.push('PAGEERROR ' + e.message));
|
||||||
|
await page.goto(BASE, { waitUntil: 'domcontentloaded' });
|
||||||
|
|
||||||
|
await page.setInputFiles('#fileInput', VIOLATIONS);
|
||||||
|
await page.waitForFunction(
|
||||||
|
() => document.querySelectorAll('#resultsTable tr').length === 2,
|
||||||
|
{ timeout: 60000 }
|
||||||
|
);
|
||||||
|
console.log('stats:', await page.evaluate(() => ({
|
||||||
|
total: document.getElementById('totalCount').innerText,
|
||||||
|
ok: document.getElementById('okCount').innerText,
|
||||||
|
violation: document.getElementById('violationCount').innerText,
|
||||||
|
})));
|
||||||
|
|
||||||
|
for (const idx of [0, 1]) {
|
||||||
|
await page.locator('#resultsTable tr').nth(idx).click();
|
||||||
|
await page.waitForTimeout(1200);
|
||||||
|
const s = await page.evaluate(() => {
|
||||||
|
const t = (id) => (document.getElementById(id) || {}).innerText?.replace(/\s+/g, ' ').trim();
|
||||||
|
const cards = (id) => Array.from((document.getElementById(id) || {}).children || [])
|
||||||
|
.map((c) => c.innerText.replace(/\s+/g, ' ').trim());
|
||||||
|
return {
|
||||||
|
badge: t('detailBadge'), quality: t('detailQuality'), overall: t('detailOverall'),
|
||||||
|
reason: t('detailReason'), roi: cards('detailArtifacts'), measurements: cards('detailMotion'),
|
||||||
|
spine: t('detailSpine'), hip: t('detailHip'),
|
||||||
|
spineHidden: document.getElementById('spineSection').classList.contains('hidden'),
|
||||||
|
hipHidden: document.getElementById('hipSection').classList.contains('hidden'),
|
||||||
|
};
|
||||||
|
});
|
||||||
|
console.log(`\n--- нарушение, строка ${idx} ---`);
|
||||||
|
console.log(' badge:', s.badge);
|
||||||
|
console.log(' quality:', s.quality);
|
||||||
|
console.log(' overall:', s.overall);
|
||||||
|
console.log(' roi:', s.roi);
|
||||||
|
console.log(' measurements:', s.measurements);
|
||||||
|
if (!s.spineHidden) console.log(' spine:', s.spine);
|
||||||
|
if (!s.hipHidden) console.log(' hip:', s.hip);
|
||||||
|
console.log(' reason:', s.reason);
|
||||||
|
await page.locator('#detailSection').scrollIntoViewIfNeeded();
|
||||||
|
await page.screenshot({ path: path.join(os.tmpdir(), `dxa_violation${idx}.png`) });
|
||||||
|
}
|
||||||
|
|
||||||
|
console.log('\npageerrors:', errors.length ? errors : 'none');
|
||||||
|
await context.close();
|
||||||
|
fs.rmSync(userDataDir, { recursive: true, force: true });
|
||||||
|
})().catch((e) => { console.error('FAILED:', e.message); process.exit(1); });
|
||||||
|
|
@ -0,0 +1,209 @@
|
||||||
|
"""
|
||||||
|
Тесты контракта API для веб-интерфейса.
|
||||||
|
|
||||||
|
Проверяют, что ответы эндпоинтов содержат поля, которые читает
|
||||||
|
`src/api/static/js/dxa-app.js`. Ранее эти ключи разошлись: панель деталей
|
||||||
|
показывала прочерки и одинаковые значения при любом клике, потому что API
|
||||||
|
отдавал `metrics` другой структуры. Тесты фиксируют контракт.
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||||
|
|
||||||
|
DATASET = Path("dataset_hack/Для теста")
|
||||||
|
CHECKPOINT = Path("models/dxa_model.pth")
|
||||||
|
|
||||||
|
pytest.importorskip("fastapi")
|
||||||
|
pytest.importorskip("httpx")
|
||||||
|
|
||||||
|
needs_assets = pytest.mark.skipif(
|
||||||
|
not (DATASET.is_dir() and CHECKPOINT.exists()),
|
||||||
|
reason="test DICOM files or trained checkpoint are not available",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Поля, которые панель деталей читает напрямую из ответа.
|
||||||
|
DETAIL_TOP_LEVEL = [
|
||||||
|
"anatomical_region", "quality_class", "quality_label", "violation_type",
|
||||||
|
"reason", "confidence", "confidence_per_class", "threshold_probability",
|
||||||
|
"region_confidence", "overall_quality", "severity", "metrics",
|
||||||
|
"view_quality", "reasons", "metrics_note",
|
||||||
|
]
|
||||||
|
# Ключи внутри metrics, которые панель использует.
|
||||||
|
METRICS_KEYS = ["sharpness_laplacian", "sharpness_fft", "roi", "calibration"]
|
||||||
|
ROI_KEYS = ["size", "bounding_box", "valid", "reason", "issues"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def client():
|
||||||
|
os.environ.setdefault("DXA_MODEL_PATH", str(CHECKPOINT))
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
from src.main import app
|
||||||
|
|
||||||
|
return TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
def _test_files():
|
||||||
|
return sorted(DATASET.glob("*.dcm"))
|
||||||
|
|
||||||
|
|
||||||
|
@needs_assets
|
||||||
|
class TestDetailedContract:
|
||||||
|
@pytest.fixture(scope="class")
|
||||||
|
def payloads(self, client):
|
||||||
|
out = []
|
||||||
|
for path in _test_files():
|
||||||
|
response = client.post(
|
||||||
|
"/api/v1/analyze/detailed",
|
||||||
|
files={"file": (path.name, path.read_bytes(), "application/dicom")},
|
||||||
|
params={"include_visualization": "true"},
|
||||||
|
)
|
||||||
|
assert response.status_code == 200, f"{path.name}: {response.text[:200]}"
|
||||||
|
out.append((path.name, response.json()))
|
||||||
|
return out
|
||||||
|
|
||||||
|
def test_all_documented_fields_present(self, payloads):
|
||||||
|
for name, data in payloads:
|
||||||
|
missing = [k for k in DETAIL_TOP_LEVEL if k not in data]
|
||||||
|
assert not missing, f"{name} missing {missing}"
|
||||||
|
|
||||||
|
def test_confidence_per_class_is_coherent(self, payloads):
|
||||||
|
for name, data in payloads:
|
||||||
|
cpc = data["confidence_per_class"]
|
||||||
|
assert set(cpc) == {"correct", "violation"}, name
|
||||||
|
assert cpc["correct"] + cpc["violation"] == pytest.approx(1.0, abs=1e-3), name
|
||||||
|
# Вероятность класса не должна расходиться с полем confidence.
|
||||||
|
assert data["confidence"] == pytest.approx(cpc["violation"], abs=1e-3), name
|
||||||
|
|
||||||
|
def test_metrics_have_expected_structure(self, payloads):
|
||||||
|
for name, data in payloads:
|
||||||
|
metrics = data["metrics"]
|
||||||
|
assert isinstance(metrics, dict) and metrics, f"{name}: metrics is empty"
|
||||||
|
missing = [k for k in METRICS_KEYS if k not in metrics]
|
||||||
|
assert not missing, f"{name} metrics missing {missing}"
|
||||||
|
for key in ROI_KEYS:
|
||||||
|
assert key in metrics["roi"], f"{name} roi missing {key}"
|
||||||
|
|
||||||
|
def test_metrics_differ_between_distinct_images(self, payloads):
|
||||||
|
"""
|
||||||
|
Метрики должны различаться для разных снимков.
|
||||||
|
|
||||||
|
Одинаковые значения были главным симптомом бага: пороги эвристик были
|
||||||
|
насыщены, и панель выглядела «не обновляющейся» при кликах.
|
||||||
|
"""
|
||||||
|
sharp = [d["metrics"]["sharpness_laplacian"] for _, d in payloads]
|
||||||
|
assert len(set(sharp)) > 1, "sharpness does not vary across images"
|
||||||
|
|
||||||
|
def test_duplicate_images_share_metrics(self, payloads):
|
||||||
|
"""Побайтные дубликаты должны давать идентичные метрики."""
|
||||||
|
by_study = {}
|
||||||
|
for name, data in payloads:
|
||||||
|
key = (data["study_uid"], data["image_uid"])
|
||||||
|
by_study.setdefault(key, []).append(data["metrics"]["sharpness_laplacian"])
|
||||||
|
for key, values in by_study.items():
|
||||||
|
if len(values) > 1:
|
||||||
|
assert max(values) == pytest.approx(min(values)), f"duplicates differ for {key}"
|
||||||
|
|
||||||
|
def test_region_specific_sections_match_region(self, payloads):
|
||||||
|
for name, data in payloads:
|
||||||
|
region = data["anatomical_region"]
|
||||||
|
if region == "spine":
|
||||||
|
assert "spine_completeness" in data, name
|
||||||
|
assert data["spine_completeness"], f"{name}: spine section empty"
|
||||||
|
if region.startswith("hip"):
|
||||||
|
assert data.get("hip_completeness") or data.get("hip_rotation"), name
|
||||||
|
|
||||||
|
def test_no_uncalibrated_verdicts_are_exposed(self, payloads):
|
||||||
|
"""
|
||||||
|
Эвристики не должны выглядеть как заключение.
|
||||||
|
|
||||||
|
Проверка фиксирует, что наружу не отдаются некалиброванные вердикты:
|
||||||
|
например число «позвонков» (эвристика выдавала 46) или тексты вида
|
||||||
|
«Позвонок 1 обрезан».
|
||||||
|
"""
|
||||||
|
for name, data in payloads:
|
||||||
|
spine = data.get("spine_completeness") or {}
|
||||||
|
assert "num_vertebrae" not in spine, f"{name}: raw vertebrae count exposed"
|
||||||
|
assert "issues" not in spine, f"{name}: raw issue texts exposed"
|
||||||
|
assert "reason" not in spine, f"{name}: raw verdict text exposed"
|
||||||
|
blob = str(data)
|
||||||
|
assert "Позвонок 1 обрезан" not in blob, f"{name}: verdict leaked into payload"
|
||||||
|
|
||||||
|
def test_quality_label_is_not_duplicated_in_badge(self, payloads):
|
||||||
|
"""Подпись класса не должна совпадать с текстом бейджа (было «OK OK»)."""
|
||||||
|
for name, data in payloads:
|
||||||
|
if data["quality_class"] == 0:
|
||||||
|
assert data["quality_label"] != "OK", name
|
||||||
|
|
||||||
|
def test_visualizations_are_valid_base64(self, payloads):
|
||||||
|
import base64
|
||||||
|
|
||||||
|
for name, data in payloads:
|
||||||
|
for field in ("image", "mask"):
|
||||||
|
raw = data.get(field)
|
||||||
|
assert raw, f"{name}: {field} is empty"
|
||||||
|
decoded = base64.b64decode(raw)
|
||||||
|
assert decoded[:8] == b"\x89PNG\r\n\x1a\n", f"{name}: {field} is not PNG"
|
||||||
|
|
||||||
|
def test_reasons_explain_the_decision(self, payloads):
|
||||||
|
for name, data in payloads:
|
||||||
|
reasons = data["reasons"]
|
||||||
|
assert isinstance(reasons, list) and reasons, name
|
||||||
|
text = " ".join(reasons)
|
||||||
|
assert data["anatomical_region"] in text or "область" in text, name
|
||||||
|
|
||||||
|
def test_metrics_note_warns_about_calibration(self, payloads):
|
||||||
|
for name, data in payloads:
|
||||||
|
assert "не калиброваны" in data["metrics_note"], name
|
||||||
|
|
||||||
|
|
||||||
|
@needs_assets
|
||||||
|
class TestBasicContract:
|
||||||
|
def test_analyze_returns_basic_fields(self, client):
|
||||||
|
path = _test_files()[0]
|
||||||
|
response = client.post(
|
||||||
|
"/api/v1/analyze",
|
||||||
|
files={"file": (path.name, path.read_bytes(), "application/dicom")},
|
||||||
|
)
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
for key in ("anatomical_region", "quality_class", "confidence", "processing_status"):
|
||||||
|
assert key in data
|
||||||
|
# Базовый эндпоинт не должен тянуть тяжёлые метрики
|
||||||
|
assert data["metrics"] == {}
|
||||||
|
|
||||||
|
|
||||||
|
@needs_assets
|
||||||
|
class TestExportContract:
|
||||||
|
def test_export_has_required_columns(self, client):
|
||||||
|
import io
|
||||||
|
|
||||||
|
import pandas as pd
|
||||||
|
|
||||||
|
files = [(f.name, f.read_bytes(), "application/dicom") for f in _test_files()]
|
||||||
|
response = client.post("/api/v1/export", files=[("files", f) for f in files])
|
||||||
|
assert response.status_code == 200
|
||||||
|
df = pd.read_excel(io.BytesIO(response.content))
|
||||||
|
required = [
|
||||||
|
"path_to_study", "study_uid", "image_uid", "anatomical_region",
|
||||||
|
"quality_class", "violation_type", "processing_status", "time_of_processing",
|
||||||
|
]
|
||||||
|
# Столбцы задания идут первыми и в заданном порядке.
|
||||||
|
assert list(df.columns)[: len(required)] == required
|
||||||
|
assert len(df) == len(_test_files())
|
||||||
|
assert (df["processing_status"] == "Success").all()
|
||||||
|
|
||||||
|
|
||||||
|
@needs_assets
|
||||||
|
class TestErrorHandling:
|
||||||
|
def test_garbage_upload_returns_500_not_crash(self, client):
|
||||||
|
response = client.post(
|
||||||
|
"/api/v1/analyze",
|
||||||
|
files={"file": ("junk.dcm", b"not a dicom", "application/dicom")},
|
||||||
|
)
|
||||||
|
assert response.status_code == 500
|
||||||
|
assert "error" in response.json()
|
||||||
|
|
@ -0,0 +1,172 @@
|
||||||
|
"""
|
||||||
|
Тесты разбора датасета 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
|
||||||
|
|
@ -0,0 +1,319 @@
|
||||||
|
"""
|
||||||
|
Тесты предобработки, метрик и обучения.
|
||||||
|
|
||||||
|
Метрики и порог проверяются на синтетических данных: важно, чтобы правило
|
||||||
|
решения работало при сильном дисбалансе классов и чтобы подбор порога не
|
||||||
|
деградировал, когда модель разделяет выборку почти идеально.
|
||||||
|
"""
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||||
|
|
||||||
|
from src.dxa.model import compute_metrics, create_model, select_threshold # noqa: E402
|
||||||
|
from src.dxa.preprocess import ( # noqa: E402
|
||||||
|
PreprocessConfig,
|
||||||
|
normalize_array,
|
||||||
|
preprocess_from_array,
|
||||||
|
to_model_tensor,
|
||||||
|
)
|
||||||
|
|
||||||
|
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 TestPreprocessConfig:
|
||||||
|
def test_defaults_are_valid(self):
|
||||||
|
cfg = PreprocessConfig()
|
||||||
|
assert cfg.norm in ("percentile", "minmax")
|
||||||
|
assert cfg.imagenet_norm is True
|
||||||
|
|
||||||
|
def test_rejects_unknown_norm(self):
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
PreprocessConfig(norm="zscore")
|
||||||
|
|
||||||
|
def test_roundtrip_through_dict(self):
|
||||||
|
cfg = PreprocessConfig(norm="minmax", imagenet_norm=False, input_size=128)
|
||||||
|
assert PreprocessConfig.from_dict(cfg.to_dict()) == cfg
|
||||||
|
|
||||||
|
def test_from_dict_ignores_unknown_keys(self):
|
||||||
|
cfg = PreprocessConfig.from_dict({"norm": "minmax", "legacy_field": 1})
|
||||||
|
assert cfg.norm == "minmax"
|
||||||
|
|
||||||
|
def test_legacy_config_disables_imagenet_norm(self):
|
||||||
|
assert PreprocessConfig.legacy().imagenet_norm is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestNormalizeArray:
|
||||||
|
def test_minmax_maps_to_unit_range(self):
|
||||||
|
arr = np.linspace(-100, 500, 1000).reshape(50, 20)
|
||||||
|
out = normalize_array(arr, PreprocessConfig(norm="minmax"))
|
||||||
|
assert out.min() == pytest.approx(0.0)
|
||||||
|
assert out.max() == pytest.approx(1.0)
|
||||||
|
assert out.dtype == np.float32
|
||||||
|
|
||||||
|
def test_percentile_clips_outliers(self):
|
||||||
|
arr = np.zeros((100, 100), dtype=np.float32)
|
||||||
|
arr[0, 0] = 10_000.0 # выброс
|
||||||
|
out = normalize_array(arr, PreprocessConfig(norm="percentile", p_low=1, p_high=99))
|
||||||
|
assert out.max() <= 1.0
|
||||||
|
assert out.min() >= 0.0
|
||||||
|
|
||||||
|
def test_constant_image_does_not_divide_by_zero(self):
|
||||||
|
out = normalize_array(np.full((10, 10), 7.0, dtype=np.float32), PreprocessConfig())
|
||||||
|
assert np.isfinite(out).all()
|
||||||
|
|
||||||
|
|
||||||
|
class TestToModelTensor:
|
||||||
|
def test_shape_and_channels(self):
|
||||||
|
cfg = PreprocessConfig(input_size=64)
|
||||||
|
out = to_model_tensor(np.zeros((30, 20), dtype=np.float32), cfg)
|
||||||
|
assert out.shape == (3, 64, 64)
|
||||||
|
assert out.dtype == np.float32
|
||||||
|
|
||||||
|
def test_imagenet_norm_changes_scale(self):
|
||||||
|
# Тёмный пиксель: при стандартизации по ImageNet он становится заметно
|
||||||
|
# отрицательным, поскольку среднее ImageNet для каналов ≈0.45.
|
||||||
|
arr = np.full((16, 16), 0.2, dtype=np.float32)
|
||||||
|
plain = to_model_tensor(arr, PreprocessConfig(input_size=16, imagenet_norm=False))
|
||||||
|
normed = to_model_tensor(arr, PreprocessConfig(input_size=16, imagenet_norm=True))
|
||||||
|
assert not np.allclose(plain, normed)
|
||||||
|
assert normed.mean() < 0.0
|
||||||
|
# Без стандартизации значение остаётся тёмным, но положительным.
|
||||||
|
assert 0.0 < plain.mean() < 0.25
|
||||||
|
|
||||||
|
def test_preprocess_from_array_matches_shape_contract(self):
|
||||||
|
out = preprocess_from_array(np.random.rand(40, 30).astype(np.float32),
|
||||||
|
PreprocessConfig(input_size=32))
|
||||||
|
assert out.shape == (3, 32, 32)
|
||||||
|
|
||||||
|
|
||||||
|
class TestComputeMetrics:
|
||||||
|
def test_perfect_separation(self):
|
||||||
|
logits = torch.tensor([-5.0, -4.0, 4.0, 5.0])
|
||||||
|
labels = torch.tensor([0, 0, 1, 1])
|
||||||
|
m = compute_metrics(logits, labels, threshold=0.0)
|
||||||
|
assert m["f1"] == pytest.approx(1.0)
|
||||||
|
assert m["roc_auc"] == pytest.approx(1.0)
|
||||||
|
assert m["tp"] == 2 and m["tn"] == 2 and m["fp"] == 0 and m["fn"] == 0
|
||||||
|
|
||||||
|
def test_inverted_predictions_give_zero_recall(self):
|
||||||
|
logits = torch.tensor([5.0, 4.0, -4.0, -5.0])
|
||||||
|
labels = torch.tensor([0, 0, 1, 1])
|
||||||
|
m = compute_metrics(logits, labels, threshold=0.0)
|
||||||
|
assert m["recall"] == 0.0
|
||||||
|
assert m["tp"] == 0
|
||||||
|
assert m["roc_auc"] == pytest.approx(0.0)
|
||||||
|
|
||||||
|
def test_single_class_labels_yield_none_auc(self):
|
||||||
|
m = compute_metrics(torch.tensor([-1.0, 1.0]), torch.tensor([0, 0]))
|
||||||
|
assert m["roc_auc"] is None and m["pr_auc"] is None
|
||||||
|
|
||||||
|
def test_threshold_is_reported_as_probability(self):
|
||||||
|
m = compute_metrics(torch.tensor([0.0]), torch.tensor([0]), threshold=0.0)
|
||||||
|
assert m["threshold_logit"] == 0.0
|
||||||
|
assert m["threshold_prob"] == pytest.approx(0.5)
|
||||||
|
|
||||||
|
def test_counts_sum_to_sample_size(self):
|
||||||
|
rng = np.random.default_rng(0)
|
||||||
|
logits = torch.tensor(rng.normal(size=200))
|
||||||
|
labels = torch.tensor((rng.random(200) > 0.7).astype(int))
|
||||||
|
m = compute_metrics(logits, labels, threshold=0.0)
|
||||||
|
assert m["tp"] + m["tn"] + m["fp"] + m["fn"] == 200
|
||||||
|
|
||||||
|
|
||||||
|
class TestSelectThreshold:
|
||||||
|
def test_finds_separating_threshold(self):
|
||||||
|
logits = torch.tensor([-3.0, -2.0, -1.0, 1.0, 2.0, 3.0])
|
||||||
|
labels = torch.tensor([0, 0, 0, 1, 1, 1])
|
||||||
|
threshold, m = select_threshold(logits, labels)
|
||||||
|
assert m["f1"] == pytest.approx(1.0)
|
||||||
|
assert -1.0 < threshold <= 1.0
|
||||||
|
|
||||||
|
def test_handles_imbalanced_data(self):
|
||||||
|
"""При ~10 % позитивов порог 0 даёт нулевой recall; подбор должен это исправить."""
|
||||||
|
rng = np.random.default_rng(1)
|
||||||
|
logits = torch.tensor(np.concatenate([rng.normal(-1, 1, 90), rng.normal(0.5, 1, 10)]))
|
||||||
|
labels = torch.tensor([0] * 90 + [1] * 10)
|
||||||
|
_, m = select_threshold(logits, labels)
|
||||||
|
assert m["recall"] > 0.5
|
||||||
|
|
||||||
|
def test_min_recall_constraint_is_respected(self):
|
||||||
|
rng = np.random.default_rng(2)
|
||||||
|
logits = torch.tensor(rng.normal(size=200))
|
||||||
|
labels = torch.tensor((rng.random(200) > 0.8).astype(int))
|
||||||
|
_, m = select_threshold(logits, labels, min_recall=0.9)
|
||||||
|
assert m["recall"] >= 0.9
|
||||||
|
|
||||||
|
def test_single_class_returns_default(self):
|
||||||
|
threshold, m = select_threshold(torch.tensor([1.0, 2.0]), torch.tensor([1, 1]))
|
||||||
|
assert threshold == 0.0
|
||||||
|
assert m["roc_auc"] is None
|
||||||
|
|
||||||
|
def test_unreachable_min_recall_falls_back(self):
|
||||||
|
"""Если требуемый recall недостижим, подбор не должен падать."""
|
||||||
|
logits = torch.tensor([5.0, 6.0, 7.0, 8.0])
|
||||||
|
labels = torch.tensor([0, 0, 1, 1])
|
||||||
|
_, m = select_threshold(logits, labels, min_recall=0.99)
|
||||||
|
assert m["threshold_logit"] is not None
|
||||||
|
|
||||||
|
|
||||||
|
class TestModelContract:
|
||||||
|
def test_forward_returns_expected_keys_and_shapes(self):
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
out = model.model(torch.zeros(2, 3, 224, 224))
|
||||||
|
assert set(out) == {"quality_logits", "region_logits"}
|
||||||
|
assert out["quality_logits"].shape == (2, 2)
|
||||||
|
assert out["region_logits"].shape == (2, 4)
|
||||||
|
|
||||||
|
def test_linear_head_has_single_layer(self):
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
linears = [m for m in model.model.quality_head if isinstance(m, torch.nn.Linear)]
|
||||||
|
assert len(linears) == 1
|
||||||
|
|
||||||
|
def test_save_load_roundtrip_preserves_predictions(self, tmp_path):
|
||||||
|
cfg = PreprocessConfig(input_size=224)
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu",
|
||||||
|
head="linear", learning_rate=1e-3)
|
||||||
|
x = torch.randn(2, 3, 224, 224)
|
||||||
|
with torch.no_grad():
|
||||||
|
before = model.model(x)["quality_logits"].clone()
|
||||||
|
|
||||||
|
path = tmp_path / "ckpt.pth"
|
||||||
|
model.save(path, preprocess=cfg, threshold=-1.5, epoch=3)
|
||||||
|
|
||||||
|
reloaded = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
metadata = reloaded.load(path)
|
||||||
|
with torch.no_grad():
|
||||||
|
after = reloaded.model(x)["quality_logits"]
|
||||||
|
|
||||||
|
assert torch.allclose(before, after, atol=1e-5)
|
||||||
|
assert metadata["threshold"] == -1.5
|
||||||
|
assert metadata["epoch"] == 3
|
||||||
|
assert metadata["backbone"] == "resnet18"
|
||||||
|
assert metadata["head"] == "linear"
|
||||||
|
assert metadata["preprocess"] == cfg.to_dict()
|
||||||
|
|
||||||
|
def test_feature_norm_buffers_survive_roundtrip(self, tmp_path):
|
||||||
|
"""Стандартизация признаков должна восстанавливаться вместе с весами."""
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
model.fit_feature_norm(_constant_loader(3))
|
||||||
|
x = torch.randn(2, 3, 224, 224)
|
||||||
|
model.model.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
before = model.model(x)["quality_logits"].clone()
|
||||||
|
|
||||||
|
path = tmp_path / "ckpt.pth"
|
||||||
|
model.save(path)
|
||||||
|
|
||||||
|
reloaded = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
reloaded.load(path)
|
||||||
|
reloaded.model.eval()
|
||||||
|
assert torch.allclose(reloaded.model.feat_mean, model.model.feat_mean)
|
||||||
|
assert torch.allclose(reloaded.model.feat_std, model.model.feat_std)
|
||||||
|
with torch.no_grad():
|
||||||
|
after = reloaded.model(x)["quality_logits"]
|
||||||
|
assert torch.allclose(before, after, atol=1e-5)
|
||||||
|
|
||||||
|
def test_frozen_backbone_keeps_batchnorm_in_eval(self):
|
||||||
|
"""
|
||||||
|
При заморозке backbone слои BatchNorm должны остаться в eval, иначе
|
||||||
|
бегущие статистики сдвигаются и признаки расходятся с теми, на которых
|
||||||
|
оценивалась стандартизация.
|
||||||
|
"""
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
model.set_backbone_trainable(False)
|
||||||
|
model.train_epoch(_constant_loader(4))
|
||||||
|
|
||||||
|
bn_modules = [m for m in model.model.features.modules()
|
||||||
|
if isinstance(m, torch.nn.modules.batchnorm._BatchNorm)]
|
||||||
|
assert bn_modules, "resnet18 must contain BatchNorm layers"
|
||||||
|
assert all(not m.training for m in bn_modules)
|
||||||
|
|
||||||
|
def test_unfrozen_backbone_enables_batchnorm_training(self):
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
model.set_backbone_trainable(True)
|
||||||
|
model.train_epoch(_constant_loader(4))
|
||||||
|
|
||||||
|
bn_modules = [m for m in model.model.features.modules()
|
||||||
|
if isinstance(m, torch.nn.modules.batchnorm._BatchNorm)]
|
||||||
|
assert all(m.training for m in bn_modules)
|
||||||
|
|
||||||
|
def test_mismatched_backbone_raises_on_load(self, tmp_path):
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
path = tmp_path / "ckpt.pth"
|
||||||
|
model.save(path)
|
||||||
|
other = create_model(backbone="resnet34", pretrained=False, device="cpu", head="linear")
|
||||||
|
with pytest.raises(RuntimeError):
|
||||||
|
other.load(path)
|
||||||
|
|
||||||
|
def test_raw_state_dict_is_rejected(self, tmp_path):
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
path = tmp_path / "raw.pth"
|
||||||
|
torch.save(model.model.state_dict(), path)
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
model.load(path)
|
||||||
|
|
||||||
|
def test_fit_feature_norm_produces_finite_std(self):
|
||||||
|
model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear")
|
||||||
|
model.fit_feature_norm(_constant_loader(4))
|
||||||
|
std = model.model.feat_std
|
||||||
|
assert torch.isfinite(std).all()
|
||||||
|
assert (std > 0).all(), "std must be strictly positive to avoid division by zero"
|
||||||
|
|
||||||
|
|
||||||
|
class _SingleBatchLoader:
|
||||||
|
"""Минимальный лоадер из одного батча: нужен для проверок без датасета."""
|
||||||
|
|
||||||
|
def __init__(self, n):
|
||||||
|
self.images = torch.zeros(n, 3, 224, 224)
|
||||||
|
self.meta = {
|
||||||
|
"label": torch.zeros(n, dtype=torch.long),
|
||||||
|
"region_id": torch.zeros(n, dtype=torch.long),
|
||||||
|
}
|
||||||
|
self.dataset = range(n)
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
yield self.images, self.meta
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
def _constant_loader(n):
|
||||||
|
"""Лоадер с постоянным входом: удобен для проверки статистик признаков."""
|
||||||
|
return _SingleBatchLoader(n)
|
||||||
|
|
||||||
|
|
||||||
|
@needs_dataset
|
||||||
|
class TestOnRealImages:
|
||||||
|
def test_preprocess_dicom_is_deterministic(self):
|
||||||
|
from src.dxa.labels import scan_dataset
|
||||||
|
from src.dxa.preprocess import preprocess_dicom
|
||||||
|
|
||||||
|
records = scan_dataset(DATASET_ROOT, with_pixel_dedup=True)
|
||||||
|
cfg = PreprocessConfig(input_size=224)
|
||||||
|
first = preprocess_dicom(records[0].path, cfg)
|
||||||
|
second = preprocess_dicom(records[0].path, cfg)
|
||||||
|
assert np.array_equal(first, second)
|
||||||
|
assert first.shape == (3, 224, 224)
|
||||||
|
assert np.isfinite(first).all()
|
||||||
|
|
||||||
|
def test_duplicate_files_produce_identical_tensors(self):
|
||||||
|
"""Файлы с одинаковым содержимым должны давать одинаковый вход модели."""
|
||||||
|
from src.dxa.labels import scan_dataset
|
||||||
|
from src.dxa.preprocess import preprocess_dicom
|
||||||
|
|
||||||
|
records = scan_dataset(DATASET_ROOT, with_pixel_dedup=True)
|
||||||
|
cfg = PreprocessConfig(input_size=224)
|
||||||
|
duplicated = [r for r in records if len(r.sources) > 1]
|
||||||
|
assert duplicated, "dataset should contain duplicates to test against"
|
||||||
|
rec = duplicated[0]
|
||||||
|
tensors = [preprocess_dicom(p, cfg) for p in rec.sources]
|
||||||
|
for other in tensors[1:]:
|
||||||
|
assert np.array_equal(tensors[0], other)
|
||||||
Loading…
Reference in New Issue