#!/usr/bin/env python3 """ Пакетный инференс классификатора качества DXA ============================================= Обрабатывает DICOM-исследования и формирует таблицу в формате требований: `path_to_study, study_uid, image_uid, anatomical_region, quality_class, violation_type, processing_status, time_of_processing`. Анатомическая область определяется по содержимому изображения (`determine_region_from_image`); имя файла не используется, чтобы не зависеть от соглашения о именах на закрытых данных. Дополнительно обученная вспомогательная голова предсказывает область, что при низкой уверенности основного метода служит уточнением. Решение принимается по логиту: порог хранится в чекпоинте и подобран по F1 на валидации при обучении. Если чекпоинт старый и порога не содержит, берётся 0. Использование: python -m src.dxa.inference --input-path <файл или каталог> --output-path results.xlsx python -m src.dxa.inference --input-path dataset_hack --output-path report.csv --zip-out masks.zip """ from __future__ import annotations import argparse import logging import sys import warnings import zipfile from dataclasses import dataclass, field from datetime import datetime from pathlib import Path from typing import Dict, List, Optional, Sequence, Tuple warnings.filterwarnings("ignore") import numpy as np import pandas as pd import pydicom import torch from PIL import Image from tqdm import tqdm sys.path.insert(0, str(Path(__file__).resolve().parents[2])) from src.dxa.labels import REGIONS, iter_dicom_files from src.dxa.model import DXAQualityModel, create_model from src.dxa.preprocess import ( PreprocessConfig, load_dicom_array, preprocess_from_array, ) logger = logging.getLogger("dxa.inference") # Предсказание вспомогательной головы -> анатомическая область REGION_BY_ID = {1: "spine", 2: "hip_right", 3: "hip_left"} # Человекочитаемые причины для типа нарушения REASON_BY_TYPE = { "position_error": "Геометрия или укладка области исследования нарушены", "artifact_motion": "Признаки артефактов движения (размытие, раздвоение контуров)", "artifact_other": "Посторонние включения или артефакты в зоне интереса", "incomplete_view": "Нужная анатомическая область видна не полностью", "roi_error": "Границы области интереса не совпадают с анатомическими", "rotation": "Выраженная ротация, искажающая анатомические границы", "quality_violation_detected": "Выявлено нарушение качества изображения", } OUTPUT_COLUMNS = [ "path_to_study", "study_uid", "image_uid", "anatomical_region", "quality_class", "violation_type", "processing_status", "time_of_processing", ] @dataclass class HeuristicSignals: """Дешёвые признаки изображения для определения анатомической области.""" bbox_aspect: float = 1.0 symmetry: float = 1.0 left_right_ratio: float = 1.0 width: int = 0 height: int = 0 @dataclass class ImageResult: """Результат обработки одного изображения.""" path_to_study: str study_uid: str image_uid: str anatomical_region: str quality_class: int violation_type: str processing_status: str time_of_processing: float confidence: float = 0.0 violation_reason: str = "" region_confidence: float = 0.0 dcm_path: str = "" metrics: Dict = field(default_factory=dict) @dataclass class Prediction: """Предсказание по одному изображению, независимое от источника (файл/поток).""" quality_class: int prob: float logit: float threshold: float anatomical_region: str region_confidence: float violation_type: str violation_reason: str samples: Dict[str, float] def get_device(prefer: Optional[str] = None) -> str: """Выбрать устройство: явно заданное или лучшее из доступных.""" if prefer: return prefer if torch.cuda.is_available(): return "cuda" if torch.backends.mps.is_available(): return "mps" return "cpu" @dataclass class LoadedCheckpoint: """Загруженный чекпоинт: модель, предобработка, порог и метаданные.""" model: DXAQualityModel preprocess: PreprocessConfig threshold: float metadata: Dict def load_model( model_path: str, backbone: Optional[str] = None, head: Optional[str] = None, device: str = "cpu", input_size: Optional[int] = None, ) -> LoadedCheckpoint: """ Загрузить модель и восстановить параметры её обработки из чекпоинта. Архитектура (backbone, тип головы) и параметры предобработки берутся из самого чекпоинта, поэтому вызывающей стороне не нужно их дублировать и невозможно рассинхронизировать обучение и инференс. Явно переданные `backbone`/`head` проверяются на совместимость с сохранёнными. Returns: LoadedCheckpoint с моделью, конфигом предобработки, порогом и метаданными. """ # Сначала читаем метаданные лёгким способом, чтобы не создавать лишние сети. try: meta_only = torch.load(model_path, map_location="cpu", weights_only=False) except Exception as exc: raise ValueError(f"Cannot read checkpoint {model_path}: {exc}") from exc saved_backbone = (meta_only or {}).get("backbone", "resnet18") saved_head = (meta_only or {}).get("head", "mlp") if backbone and backbone != saved_backbone: raise ValueError( f"Checkpoint was trained with backbone={saved_backbone!r}, but {backbone!r} was requested" ) if head and head != saved_head: raise ValueError( f"Checkpoint was trained with head={saved_head!r}, but {head!r} was requested" ) model = create_model(backbone=saved_backbone, head=saved_head, pretrained=False, device=device) metadata = model.load(model_path) model.model.eval() cfg = load_preprocess(metadata, input_size) threshold = float(metadata.get("threshold_logit", metadata.get("threshold", 0.0))) return LoadedCheckpoint(model=model, preprocess=cfg, threshold=threshold, metadata=metadata) def load_preprocess(metadata: Dict, input_size: Optional[int] = None) -> PreprocessConfig: """Восстановить параметры предобработки из чекпоинта.""" cfg = PreprocessConfig.from_dict(metadata.get("preprocess") or {}) if input_size: cfg = PreprocessConfig.from_dict({**cfg.to_dict(), "input_size": input_size}) return cfg def heuristic_signals(img: np.ndarray) -> HeuristicSignals: """ Геометрические признаки изображения. Порог яркой области берётся по 95-му перцентилю. Размеры кадра входят в набор сигналов, потому что у аппарата позвоночные и бедренные снимки имеют разную ширину кадра (300 против 280 пикселей), и это самый надёжный признак области на данном оборудовании. """ norm = (img - img.min()) / (img.max() - img.min() + 1e-6) h, w = norm.shape binary = norm > np.percentile(norm, 95) if not binary.any(): return HeuristicSignals(width=w, height=h) rows, cols = np.any(binary, axis=1), np.any(binary, axis=0) rmin, rmax = np.where(rows)[0][[0, -1]] cmin, cmax = np.where(cols)[0][[0, -1]] bbox_aspect = (rmax - rmin) / ((cmax - cmin) + 1e-6) left = binary[:, :w // 2].sum() right = binary[:, w // 2:].sum() ratio = left / (right + 1e-6) left_half = norm[:, :w // 2] right_half = np.fliplr(norm[:, w // 2:]) m = min(left_half.shape[1], right_half.shape[1]) symmetry = 1 - np.abs(left_half[:, :m] - right_half[:, :m]).mean() / (norm.std() + 1e-6) return HeuristicSignals(float(bbox_aspect), float(symmetry), float(ratio), int(w), int(h)) # Ширина кадра в пикселях, разделяющая области на этом оборудовании. SPINE_MIN_WIDTH = 295 def resolve_region( predicted_region: Optional[str], region_confidence: float, signals: HeuristicSignals, min_confidence: float = 0.5, ) -> Tuple[str, float]: """ Определить анатомическую область. Порядок решений: 1. Ширина кадра: у позвоночника кадр шире (300 px против 280 px у бёдер). На этом оборудовании признак разделяет области безошибочно, поэтому используется первым. 2. При нетипичной ширине — предсказание обученной головы области. 3. Если голова неуверена — форма яркой области и перевес светимости. Важно, что область НЕ берётся из имени файла: на закрытом наборе имена могут не содержать разметки региона. """ if signals.width: if signals.width >= SPINE_MIN_WIDTH: return "spine", 0.8 # Бедро: ширину кадра делят левый и правый снимки, поэтому сторону # определяем по перевесу светимости яркой области. Голова обучена на # обе стороны, но различает их хуже, чем асимметрия. if signals.left_right_ratio > 1.3: return "hip_right", 0.6 if signals.left_right_ratio < 0.7: return "hip_left", 0.6 if predicted_region in ("hip_left", "hip_right") and region_confidence >= min_confidence: return predicted_region, region_confidence return "hip", 0.4 if predicted_region in REGIONS and region_confidence >= min_confidence: return predicted_region, region_confidence aspect = signals.bbox_aspect if aspect < 1.5 or (aspect < 1.8 and signals.symmetry > 0.35): return "spine", 0.4 if signals.left_right_ratio > 1.3: return "hip_right", 0.3 if signals.left_right_ratio < 0.7: return "hip_left", 0.3 return "hip", 0.3 def classify_violation_type( region: Optional[str], metrics: Dict, samples: Optional[Dict[str, float]], ) -> Tuple[str, str]: """ Определить тип нарушения по эвристическим метрикам изображения. Возвращает (тип, пояснение). Тип выбирается по наиболее выраженному признаку; при отсутствии сигналов возвращается общая категория. """ motion = samples.get("laplacian_variance") if samples else None bright_frac = samples.get("bright_fraction") if samples else None if motion is not None and motion < metrics.get("motion_threshold", 0.0): return "artifact_motion", REASON_BY_TYPE["artifact_motion"] if bright_frac is not None and bright_frac > metrics.get("artifact_threshold", 1.0): return "artifact_other", REASON_BY_TYPE["artifact_other"] if region in ("hip_left", "hip_right") and samples: aspect = samples.get("bbox_aspect", 1.0) if aspect < 0.4 or aspect > 3.0: return "rotation", REASON_BY_TYPE["rotation"] return "quality_violation_detected", REASON_BY_TYPE["quality_violation_detected"] def image_samples(img: np.ndarray) -> Dict[str, float]: """Числовые характеристики изображения для отчёта и выбора типа нарушения.""" norm = (img - img.min()) / (img.max() - img.min() + 1e-6) lap = np.abs(np.diff(norm, axis=0)).mean() + np.abs(np.diff(norm, axis=1)).mean() return { "laplacian_variance": float(((norm - norm.mean()) ** 2).mean() * lap), "bright_fraction": float((norm > np.percentile(norm, 99)).mean()), "bbox_aspect": heuristic_signals(img).bbox_aspect, } def predict_from_array(arr: np.ndarray, model: DXAQualityModel, cfg: PreprocessConfig, device: str, threshold: float) -> Prediction: """ Предсказание качества по уже прочитанному массиву пикселей. Общая точка входа для файлового инференса и HTTP API: гарантирует, что предобработка и решающее правило совпадают во всех режимах. """ tensor = torch.from_numpy(preprocess_from_array(arr, cfg)).float().unsqueeze(0).to(device) with torch.no_grad(): out = model.model(tensor) logit = float(out["quality_logits"][0, 1]) prob = float(torch.sigmoid(out["quality_logits"][0, 1])) region_probs = torch.softmax(out["region_logits"][0], dim=0) region_conf, region_id = float(region_probs.max()), int(region_probs.argmax()) signals = heuristic_signals(arr) region, region_conf = resolve_region(REGION_BY_ID.get(region_id), region_conf, signals) quality_class = 1 if logit >= threshold else 0 samples = image_samples(arr) if quality_class == 1: violation_type, reason = classify_violation_type(region, {}, samples) else: violation_type, reason = "", "" return Prediction( quality_class=quality_class, prob=prob, logit=logit, threshold=threshold, anatomical_region=region, region_confidence=region_conf, violation_type=violation_type, violation_reason=reason, samples=samples, ) def predict_from_bytes(dcm_bytes: bytes, model: DXAQualityModel, cfg: PreprocessConfig, device: str, threshold: float) -> Tuple[Prediction, pydicom.Dataset]: """ Предсказание по байтам DICOM без записи на диск. API принимает файлы потоком, поэтому чтение идёт через BytesIO — временные файлы не создаются. """ import io as _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 str(getattr(ds, "PhotometricInterpretation", "")).upper() == "MONOCHROME1": arr = arr.max() - arr return predict_from_array(arr, model, cfg, device, threshold), ds def study_uid_for(ds: pydicom.Dataset, dcm_path: Path) -> str: """StudyInstanceUID из DICOM; при отсутствии — имя каталога исследования.""" uid = str(getattr(ds, "StudyInstanceUID", "") or "").strip() if uid: return uid return dcm_path.parent.name def process_one( dcm_path: Path, model: DXAQualityModel, cfg: PreprocessConfig, device: str, threshold: float, ) -> ImageResult: """Полная обработка одного DICOM файла.""" started = datetime.now() path_to_study = str(dcm_path.parent) try: arr = load_dicom_array(dcm_path) prediction = predict_from_array(arr, model, cfg, device, threshold) ds = pydicom.dcmread(str(dcm_path), stop_before_pixels=True) return ImageResult( path_to_study=path_to_study, study_uid=study_uid_for(ds, dcm_path), image_uid=str(getattr(ds, "SOPInstanceUID", "") or ""), anatomical_region=prediction.anatomical_region, quality_class=prediction.quality_class, violation_type=prediction.violation_type, processing_status="Success", time_of_processing=(datetime.now() - started).total_seconds(), confidence=round(prediction.prob, 4), violation_reason=prediction.violation_reason, region_confidence=round(prediction.region_confidence, 4), dcm_path=str(dcm_path), metrics={"logit": prediction.logit, "threshold": threshold, **prediction.samples}, ) except Exception as exc: # Требование: необработанных исключений быть не должно, все ошибки # фиксируются в отчёте со статусом Failure. logger.warning("Failed to process %s: %s", dcm_path, exc) return ImageResult( path_to_study=path_to_study, study_uid="", image_uid="", anatomical_region="unknown", quality_class=-1, violation_type="", processing_status=f"Failure: {type(exc).__name__}: {str(exc)[:120]}", time_of_processing=(datetime.now() - started).total_seconds(), dcm_path=str(dcm_path), ) def write_visualizations(results: Sequence[ImageResult], cfg: PreprocessConfig, zip_path: Path, enabled: bool) -> None: """ Дополнительный функционал: zip-архив с изображениями, где выделена зона интереса (порог по 90-му перцентилю яркости). """ if not enabled: return with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as archive: for res in results: if res.processing_status != "Success" or not res.dcm_path: continue try: arr = load_dicom_array(res.dcm_path) norm = (arr - arr.min()) / (arr.max() - arr.min() + 1e-6) mask = norm > np.percentile(norm, 90) rgb = np.stack([norm] * 3, axis=-1) edge = np.zeros_like(mask) vertical = mask[2:, 1:-1] ^ mask[:-2, 1:-1] horizontal = mask[1:-1, 2:] ^ mask[1:-1, :-2] edge[1:-1, 1:-1] = vertical | horizontal rgb[edge] = [1.0, 0.2, 0.2] name = f"{res.study_uid or 'study'}_{res.image_uid or Path(res.dcm_path).stem}.png" archive.writestr(name, _to_png_bytes(rgb)) except Exception as exc: logger.warning("Visualization failed for %s: %s", res.dcm_path, exc) def _to_png_bytes(rgb: np.ndarray) -> bytes: import io image = Image.fromarray((np.clip(rgb, 0, 1) * 255).astype(np.uint8)) buffer = io.BytesIO() image.save(buffer, format="PNG") return buffer.getvalue() def results_to_dataframe(results: Sequence[ImageResult]) -> pd.DataFrame: """Таблица строго в формате требований; пояснения добавляются справа.""" rows = [] for r in results: rows.append({ "path_to_study": r.path_to_study, "study_uid": r.study_uid, "image_uid": r.image_uid, "anatomical_region": r.anatomical_region, "quality_class": r.quality_class, "violation_type": r.violation_type, "processing_status": r.processing_status, "time_of_processing": round(r.time_of_processing, 4), "confidence": r.confidence, "violation_reason": r.violation_reason, "region_confidence": r.region_confidence, }) return pd.DataFrame(rows, columns=OUTPUT_COLUMNS + ["confidence", "violation_reason", "region_confidence"]) def run(args: argparse.Namespace) -> pd.DataFrame: """Пакетная обработка входного пути и запись отчёта.""" device = get_device(args.device) logger.info("Device: %s", device) model_path = Path(args.model_path) if not model_path.exists(): raise FileNotFoundError( f"Model checkpoint not found: {model_path}. Train it first: python -m src.dxa.train" ) checkpoint = load_model( str(model_path), backbone=args.backbone, head=args.head, device=device, input_size=args.input_size, ) logger.info( "Model: backbone=%s head=%s, threshold(logit)=%.4f, preprocess=%s", checkpoint.metadata.get("backbone"), checkpoint.metadata.get("head"), checkpoint.threshold, checkpoint.preprocess.to_dict(), ) dcm_files = list(iter_dicom_files(args.input_path)) if not dcm_files: raise FileNotFoundError(f"No DICOM files found under {args.input_path}") logger.info("Found %d DICOM files", len(dcm_files)) results: List[ImageResult] = [] for dcm_path in tqdm(dcm_files, desc="Processing"): results.append(process_one( dcm_path, checkpoint.model, checkpoint.preprocess, device, checkpoint.threshold )) df = results_to_dataframe(results) output_path = Path(args.output_path) output_path.parent.mkdir(parents=True, exist_ok=True) if output_path.suffix.lower() == ".csv": df.to_csv(output_path, index=False) else: df.to_excel(output_path, index=False) if args.zip_out: write_visualizations(results, checkpoint.preprocess, Path(args.zip_out), enabled=True) logger.info("Visualizations: %s", args.zip_out) successful = int((df["processing_status"] == "Success").sum()) logger.info( "Done. %d/%d processed (%.1f%%), quality_class=1 in %d rows, median time %.3fs -> %s", successful, len(df), 100 * successful / max(len(df), 1), int((df["quality_class"] == 1).sum()), float(df["time_of_processing"].median()) if len(df) else 0.0, output_path, ) return df def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( description="Пакетный инференс классификатора качества DXA", formatter_class=argparse.ArgumentDefaultsHelpFormatter, ) parser.add_argument("--input-path", required=True, help="DICOM файл или каталог") parser.add_argument("--output-path", required=True, help="Путь к .xlsx или .csv") parser.add_argument("--model-path", default="models/dxa_model.pth", help="Чекпоинт модели") parser.add_argument("--backbone", default=None, choices=["resnet18", "resnet34"], help="Проверить соответствие backbone в чекпоинте (по умолчанию — из чекпоинта)") parser.add_argument("--head", default=None, choices=["linear", "mlp"], help="Проверить соответствие головы в чекпоинте (по умолчанию — из чекпоинта)") parser.add_argument("--input-size", type=int, default=None, help="Переопределить размер входа (по умолчанию — из чекпоинта)") parser.add_argument("--device", default=None, help="cpu / cuda / mps") parser.add_argument("--zip-out", default=None, help="Zip-архив с визуализацией зоны интереса (дополнительный функционал)") return parser def main(argv: Optional[List[str]] = None) -> int: logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") args = build_parser().parse_args(argv) try: run(args) except (FileNotFoundError, ValueError) as exc: logger.error("%s", exc) return 1 return 0 if __name__ == "__main__": raise SystemExit(main())