765 lines
34 KiB
Python
765 lines
34 KiB
Python
#!/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())
|