develop - hack_2026

This commit is contained in:
denis 2026-09-24 00:11:49 +03:00
parent 2f279d2816
commit ddfeb01174
12 changed files with 1808 additions and 0 deletions

9
src/api/static/vendor/fontawesome.css vendored Normal file

File diff suppressed because one or more lines are too long

83
src/api/static/vendor/tailwind.js vendored Normal file

File diff suppressed because one or more lines are too long

Binary file not shown.

176
src/dxa/discriminator.py Normal file
View File

@ -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())

329
src/dxa/labels.py Normal file
View File

@ -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"))

152
src/dxa/preprocess.py Normal file
View File

@ -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))

163
tests/browser/ui_check.js Normal file
View File

@ -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); });

100
tests/browser/ui_offline.js Normal file
View File

@ -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); });

View File

@ -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); });

209
tests/test_api_contract.py Normal file
View File

@ -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()

172
tests/test_labels.py Normal file
View File

@ -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

View File

@ -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)