bone_2026/src/dxa/inference.py

765 lines
34 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/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 base64
import io
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
def quality_metrics(arr: np.ndarray, region: str) -> Dict:
"""
Числовые характеристики снимка, показываемые в панели деталей.
ВАЖНО о статусе этих величин. Эвристики из `src/quality/detailed_assessment`
(пороги для «движения», «артефактов», «ROI») на этом датасете насыщены:
например `motion_detected` и `any_detected` оказываются True практически
для любого снимка, а проверка отступов ROI срабатывает для 227 из 252
изображений, потому что маска строится по 90-му перцентилю яркости и
касается границ кадра. Направление признака «резкости» к тому же
противоположно в позвоночнике и в бёдрах, поэтому единый порог для них
неприменим.
Выдавать такие значения как заключение («артефакт: Да») нельзя — это
дезинформирует врача. Поэтому здесь сохраняются числовые измерения, по
которым решение можно проверить, а флаги-вердикты не формируются.
Классификацию выполняет модель (`quality_class`), она проверена отдельно:
внутри областей её ROC-AUC 0.84–0.95, тогда как правило «позвоночник =
нарушение» даёт 0.50 (см. `python -m src.dxa.discriminator`).
"""
norm = (arr - arr.min()) / (arr.max() - arr.min() + 1e-6)
seg = (norm > np.percentile(norm, 90)).astype(np.uint8)
from src.quality.detailed_assessment import (
calculate_blur_fft,
calculate_blur_laplacian,
check_roi_boundaries,
)
laplacian = float(calculate_blur_laplacian(norm))
fft = float(calculate_blur_fft(norm))
roi = check_roi_boundaries(seg, norm.shape)
metrics: Dict = {
"sharpness_laplacian": laplacian,
"sharpness_fft": fft,
"roi": {
"size": roi.get("size"),
"bounding_box": roi.get("bounding_box"),
"valid": roi.get("valid"),
"reason": roi.get("reason"),
"issues": roi.get("issues") or [],
},
# Границы информативного поля: сравнивать значения между снимками
# корректно только при одинаковом размере кадра и одинаковой укладке.
"calibration": {
"frame_width": int(norm.shape[1]),
"frame_height": int(norm.shape[0]),
"mask_percentile": 90,
},
}
if region == "spine":
from src.quality.detailed_assessment import check_spine_completeness
completeness = check_spine_completeness(seg)
# Число «позвонков» и текстовые комментарии эвристики намеренно не
# передаются: на маске по яркости она выделяет десятки объектов
# (например «Позвонков 46») и выдаёт это как заключение. Показывать
# такое врачу нельзя. Остаётся только справочный флаг полноты.
metrics["spine_completeness"] = {
"valid": completeness.get("valid"),
"issue_count": len(completeness.get("issues") or []),
}
elif region in ("hip", "hip_left", "hip_right"):
from src.quality.detailed_assessment import check_hip_completeness, check_hip_rotation
completeness = check_hip_completeness(seg)
rotation = check_hip_rotation(seg, np.stack([norm] * 3, axis=2))
metrics["hip_completeness"] = {
"valid": completeness.get("valid"),
"aspect_ratio": completeness.get("aspect_ratio"),
"issue_count": len(completeness.get("issues") or []),
}
metrics["hip_rotation"] = {
"valid": rotation.get("valid"),
"rotation_angle": rotation.get("rotation_angle"),
"issue_count": len(rotation.get("issues") or []),
}
return metrics
def build_reasons(prediction: "Prediction", metrics: Dict) -> List[str]:
"""
Пояснения к решению для панели деталей.
Опираются только на проверенные величины: анатомическую область и
вероятность класса от модели. Числовые метрики эвристик в формулировки не
выносятся — их достоверность на этом датасете не подтверждена (см.
`quality_metrics`).
"""
region_label = {
"spine": "поясничный отдел позвоночника",
"hip_left": "проксимальный отдел левой бедренной кости",
"hip_right": "проксимальный отдел правой бедренной кости",
"hip": "проксимальный отдел бедренной кости",
}.get(prediction.anatomical_region, "область не определена")
reasons = [
f"Определена область: {region_label} "
f"(уверенность {prediction.region_confidence:.2f})."
]
if prediction.quality_class == 1:
reasons.append(
f"Модель оценила вероятность нарушения как {prediction.prob:.2f} "
f"(порог {prob_from_logit(prediction.threshold):.2f}); "
"изображение требует ручной проверки."
)
else:
reasons.append(
f"Вероятность нарушения {prediction.prob:.2f} ниже рабочего порога; "
"изображение отнесено к качественным."
)
if prediction.region_confidence < 0.5:
reasons.append(
"Анатомическая область определена неуверенно — проверьте укладку "
"и корректность выбранной области."
)
reasons.append(
"Числовые метрики (резкость, границы области) приведены для справки: "
"их пороги на этом наборе данных не калиброваны."
)
return reasons
@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]
metrics: Dict = field(default_factory=dict)
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 prob_from_logit(logit: float) -> float:
"""Перевести порог из логитов в вероятность (для отображения человеку)."""
return float(1.0 / (1.0 + np.exp(-logit)))
def mask_png_base64(arr: np.ndarray, outline: bool = True) -> str:
"""
PNG в base64: маска костной ткани (90-й перцентиль яркости).
При ``outline=True`` вместо сплошной маски рисуется её контур на исходном
изображении — так на снимке видно, какая зона попала в маску, и это
читаемее для врача, чем чёрно-белая заливка.
"""
norm = (arr - arr.min()) / (arr.max() - arr.min() + 1e-6)
mask = norm > np.percentile(norm, 90)
if not outline:
image = Image.fromarray((mask * 255).astype(np.uint8))
else:
rgb = np.stack([norm] * 3, axis=-1)
edges = np.zeros_like(mask)
edges[1:-1, 1:-1] = (
(mask[2:, 1:-1] ^ mask[:-2, 1:-1]) | (mask[1:-1, 2:] ^ mask[1:-1, :-2])
)
rgb[edges] = [1.0, 0.2, 0.2]
image = Image.fromarray((np.clip(rgb, 0, 1) * 255).astype(np.uint8))
buffer = io.BytesIO()
image.save(buffer, format="PNG")
return base64.b64encode(buffer.getvalue()).decode("utf-8")
def image_png_base64(arr: np.ndarray) -> str:
"""PNG в base64: исходное изображение, приведённое к диапазону [0, 255]."""
norm = (arr - arr.min()) / (arr.max() - arr.min() + 1e-6)
buffer = io.BytesIO()
Image.fromarray((norm * 255).astype(np.uint8)).save(buffer, format="PNG")
return base64.b64encode(buffer.getvalue()).decode("utf-8")
def predict_from_array(arr: np.ndarray, model: DXAQualityModel,
cfg: PreprocessConfig, device: str, threshold: float,
with_metrics: bool = False) -> Prediction:
"""
Предсказание качества по уже прочитанному массиву пикселей.
Общая точка входа для файлового инференса и HTTP API: гарантирует, что
предобработка и решающее правило совпадают во всех режимах.
Args:
with_metrics: если True, дополнительно считаются эвристические метрики
по критериям качества (структура как у ``generate_quality_report``)
и тип нарушения для снимков с нарушением.
"""
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)
metrics: Dict = {}
if with_metrics:
metrics = quality_metrics(arr, region)
# Тип нарушения: модель различает только «годно / нарушение», поэтому для
# снимков с нарушением указывается общая категория. Конкретный тип требует
# разметки типов на уровне снимка, которой в наборе нет.
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,
metrics=metrics,
)
def predict_from_bytes(dcm_bytes: bytes, model: DXAQualityModel,
cfg: PreprocessConfig, device: str, threshold: float,
with_metrics: bool = False) -> Tuple[Prediction, pydicom.Dataset]:
"""
Предсказание по байтам DICOM без записи на диск.
API принимает файлы потоком, поэтому чтение идёт через BytesIO — временные
файлы не создаются.
"""
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
prediction = predict_from_array(arr, model, cfg, device, threshold, with_metrics=with_metrics)
return prediction, 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())