knee-model / checkpoint.py
leminhhung0101's picture
Create checkpoint.py
4e70112 verified
Raw
History Blame Contribute Delete
33.1 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
RSNA Knee MRI — Standalone Folder Inference
=============================================
Chạy model đã train (best_model.pth / last_checkpoint.pth) trên BẤT KỲ
folder DICOM nào, không cần cấu trúc dữ liệu Kaggle (train.csv,
test_series.csv, ...). Toàn bộ hàm cần thiết (đọc DICOM, tiền xử lý,
kiến trúc model, load checkpoint, suy luận) đều nằm trong 1 file này.
FOLDER STRUCTURE ĐƯỢC HỖ TRỢ (tự động phát hiện)
-------------------------------------------------
1) Một study duy nhất, các series là folder con trực tiếp:
my_study/
SeriesA/*.dcm
SeriesB/*.dcm
SeriesC/*.dcm
2) Nhiều study, mỗi study là 1 folder con chứa các series:
many_studies/
StudyInstanceUID_1/
SeriesA/*.dcm
SeriesB/*.dcm
StudyInstanceUID_2/
SeriesA/*.dcm
3) Lồng sâu hơn cũng được — bất kỳ folder nào chứa trực tiếp file .dcm
(hoặc file DICOM không có đuôi .dcm) sẽ được coi là 1 series; series
được gom theo folder cấp 1 ngay dưới --input-folder.
METADATA (Anatomical_Plane / Fluid_Sensitive / Fat_Suppression)
-----------------------------------------------------------------
Dữ liệu thi đấu gốc lấy 3 thông tin này từ train_series.csv /
test_series.csv. Với 1 folder bất kỳ, ta không có sẵn file đó, nên
script này TỰ SUY LUẬN (heuristic) từ DICOM header:
- Anatomical_Plane : từ ImageOrientationPatient (vector pháp tuyến
của mặt phẳng ảnh — trục nào chiếm ưu thế quyết định Sagittal /
Coronal / Axial).
- Fluid_Sensitive / Fat_Suppression : từ từ khoá trong
SeriesDescription / ProtocolName / SequenceName (vd "T2", "PD",
"STIR", "FS", "FATSAT" ...).
Đây CHỈ LÀ suy đoán — nếu bạn có metadata chính xác, hãy truyền vào
qua --metadata-csv với các cột:
SeriesFolder,Anatomical_Plane,Fluid_Sensitive,Fat_Suppression
(SeriesFolder = tên folder chứa trực tiếp các file .dcm của series đó).
Giá trị trong file này sẽ ĐÈ LÊN kết quả suy luận tự động.
CÁCH DÙNG
---------
Không dùng dòng lệnh / argparse. Sửa trực tiếp các đường dẫn trong khối
"USER CONFIG" ngay bên dưới (INPUT_FOLDER, CHECKPOINT_PATH, ...) rồi
chạy thẳng file:
python infer_folder.py
Yêu cầu thư viện: torch, timm, pydicom, pandas, numpy, (tuỳ chọn) opencv-python, tqdm.
"""
import os
import json
import math
import warnings
from pathlib import Path
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple
import numpy as np
import pandas as pd
try:
import cv2
except ImportError:
cv2 = None
import pydicom
import torch
import torch.nn as nn
import torch.nn.functional as F
try:
import timm
except ImportError as e:
raise ImportError(
"timm is required for this script. Install with: pip install timm"
) from e
try:
from tqdm.auto import tqdm
except ImportError:
def tqdm(x, **kwargs):
return x
# ============================================================
# USER CONFIG — sửa các đường dẫn dưới đây rồi chạy thẳng file này
# ============================================================
# Folder chứa DICOM cần test — 1 study (series là folder con trực tiếp)
# hoặc nhiều study (mỗi study là 1 folder con chứa các series).
INPUT_FOLDER = "/path/to/your/test_folder"
# Checkpoint đã train (best_model.pth hoặc last_checkpoint.pth).
CHECKPOINT_PATH = "/path/to/best_model.pth"
# Nơi lưu kết quả dự đoán.
OUTPUT_CSV = "/path/to/predictions.csv"
# (Tuỳ chọn) CSV ghi đè metadata tự suy luận — để None nếu không có.
# Cột bắt buộc: SeriesFolder,Anatomical_Plane,Fluid_Sensitive,Fat_Suppression
METADATA_CSV = None
# "cuda" / "cpu" / None (None = tự động chọn cuda nếu có).
DEVICE = None
# Lưu thêm attention weights (per-label) ra file .json cạnh OUTPUT_CSV.
SAVE_ATTENTION = False
# Chỉ chạy N study đầu tiên tìm được (None = chạy hết) — hữu ích để test nhanh.
LIMIT_STUDIES = None
# Các tham số kiến trúc — PHẢI khớp với lúc train, chỉ đổi nếu bạn train
# với cấu hình khác mặc định.
BACKBONE_NAME = "convnextv2_tiny.fcmae_ft_in22k_in1k"
IMG_SIZE = 224
EMBED_DIM = 384
MAX_SERIES = 8
CLIPS_PER_SERIES = 3
# ============================================================
# CONSTANTS — must match the values used during training
# ============================================================
LABELS = [
"ACL",
"MCL",
"Medial Meniscus",
"Lateral Meniscus",
"Medial OA",
"Lateral OA",
"PF OA",
"Effusion",
"Synovitis",
"Baker's",
"Contusion",
"Fracture",
]
N_LABELS = len(LABELS)
PLANES = ["Sagittal", "Coronal", "Axial"]
PLANE2IDX = {p: i for i, p in enumerate(PLANES)}
# Column order: [Sagittal, Coronal, Axial]
PLANE_PRIOR = np.array(
[
[1.00, 0.65, 0.40], # ACL
[0.65, 1.00, 0.45], # MCL
[1.00, 0.90, 0.60], # Medial Meniscus
[0.90, 1.00, 0.65], # Lateral Meniscus
[0.90, 1.00, 0.40], # Medial OA
[0.85, 1.00, 0.40], # Lateral OA
[0.90, 0.55, 1.00], # PF OA
[0.90, 0.85, 0.90], # Effusion
[0.90, 0.85, 0.90], # Synovitis
[0.95, 0.70, 1.00], # Baker's
[1.00, 0.90, 0.80], # Contusion
[1.00, 0.95, 0.80], # Fracture
],
dtype=np.float32,
)
FLUID_PRIOR = np.array(
[0.80, 0.75, 0.65, 0.65, 0.45, 0.45, 0.45, 0.95, 0.90, 0.90, 0.90, 0.85],
dtype=np.float32,
)
FAT_SUPPRESSION_PRIOR = np.array(
[0.40, 0.40, 0.35, 0.35, 0.25, 0.25, 0.25, 0.65, 0.65, 0.70, 0.80, 0.75],
dtype=np.float32,
)
IMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406], dtype=torch.float32).view(1, 3, 1, 1)
IMAGENET_STD = torch.tensor([0.229, 0.224, 0.225], dtype=torch.float32).view(1, 3, 1, 1)
# ============================================================
# CONFIG (inference-only)
# ============================================================
@dataclass
class InferCFG:
backbone_name: str = "convnextv2_tiny.fcmae_ft_in22k_in1k"
img_size: int = 224
embed_dim: int = 384
label_transformer_layers: int = 2
label_transformer_heads: int = 8
dropout: float = 0.15
max_series: int = 8
clips_per_series: int = 3
percentile_low: float = 1.0
percentile_high: float = 99.0
amp: bool = True
device: str = field(default_factory=lambda: "cuda" if torch.cuda.is_available() else "cpu")
# ============================================================
# DICOM READING (same numerical behavior as training)
# ============================================================
def _safe_float(x, default: float = 0.0) -> float:
try:
return float(x)
except Exception:
return default
def is_dicom_file(path: Path) -> bool:
"""Detect DICOM files that don't use the .dcm extension via the 'DICM' magic bytes."""
try:
with open(path, "rb") as f:
f.seek(128)
return f.read(4) == b"DICM"
except Exception:
return False
def find_dicom_files(folder: Path) -> List[Path]:
"""Return DICOM files directly inside `folder` (non-recursive)."""
if not folder.is_dir():
return []
files = [
p for p in folder.iterdir()
if p.is_file() and p.suffix.lower() == ".dcm"
]
if files:
return sorted(files)
# Fallback for DICOM files with no/other extension.
candidates = [
p for p in folder.iterdir()
if p.is_file() and not p.name.startswith(".")
]
return sorted(p for p in candidates if is_dicom_file(p))
def sort_dicom_paths(paths: List[str]) -> List[str]:
"""Sort slices by physical position when available, otherwise InstanceNumber."""
records = []
for p in paths:
try:
ds = pydicom.dcmread(p, stop_before_pixels=True, force=True)
ipp = getattr(ds, "ImagePositionPatient", None)
iop = getattr(ds, "ImageOrientationPatient", None)
instance = int(getattr(ds, "InstanceNumber", 0))
if ipp is not None and iop is not None and len(iop) >= 6 and len(ipp) >= 3:
row = np.asarray(iop[:3], dtype=np.float64)
col = np.asarray(iop[3:6], dtype=np.float64)
normal = np.cross(row, col)
position = float(np.dot(np.asarray(ipp[:3], dtype=np.float64), normal))
records.append((0, position, instance, p))
else:
records.append((1, instance, 0, p))
except Exception:
records.append((2, 0, 0, p))
records.sort(key=lambda x: (x[0], x[1], x[2], x[3]))
return [x[3] for x in records]
def read_dicom_pixels(path: str, plow: float, phigh: float) -> np.ndarray:
ds = pydicom.dcmread(path, force=True)
arr = ds.pixel_array.astype(np.float32)
if arr.ndim == 3:
arr = arr[..., 0]
slope = _safe_float(getattr(ds, "RescaleSlope", 1.0), 1.0)
intercept = _safe_float(getattr(ds, "RescaleIntercept", 0.0), 0.0)
arr = arr * slope + intercept
arr = np.nan_to_num(arr, nan=0.0, posinf=0.0, neginf=0.0)
photometric = str(getattr(ds, "PhotometricInterpretation", "")).upper()
if photometric == "MONOCHROME1":
arr = arr.max() - arr
lo, hi = np.percentile(arr, [plow, phigh])
if hi <= lo:
lo = float(arr.min())
hi = float(arr.max())
arr = np.clip(arr, lo, hi)
arr = (arr - lo) / (hi - lo + 1e-6)
return arr.astype(np.float32, copy=False)
def resize_slice(img: np.ndarray, size: int) -> np.ndarray:
if img.shape == (size, size):
return img.astype(np.float32, copy=False)
if cv2 is not None:
interpolation = (
cv2.INTER_AREA if (img.shape[0] >= size and img.shape[1] >= size)
else cv2.INTER_LINEAR
)
return cv2.resize(
img.astype(np.float32, copy=False), (size, size), interpolation=interpolation
).astype(np.float32, copy=False)
t = torch.from_numpy(img.astype(np.float32, copy=False))[None, None]
out = F.interpolate(t, size=(size, size), mode="bilinear", align_corners=False)
return out[0, 0].numpy().astype(np.float32, copy=False)
def preprocess_slice_float(path: str, cfg: InferCFG) -> np.ndarray:
"""Read + resize + per-slice z-score + clip to [-5, 5] (float32, no quantization)."""
img = read_dicom_pixels(path, cfg.percentile_low, cfg.percentile_high)
img = resize_slice(img, cfg.img_size)
mean = float(img.mean())
std = float(img.std())
img = (img - mean) / (std + 1e-6)
img = np.clip(img, -5.0, 5.0)
return img.astype(np.float32, copy=False)
def load_series_volume(series_dir: Path, cfg: InferCFG) -> np.ndarray:
"""Decode every slice of one series folder into shape [n_slices, H, W] float32."""
paths = [str(p) for p in find_dicom_files(series_dir)]
if not paths:
return np.zeros((0, cfg.img_size, cfg.img_size), dtype=np.float32)
paths = sort_dicom_paths(paths)
slices = []
for p in paths:
try:
slices.append(preprocess_slice_float(p, cfg))
except Exception as exc:
warnings.warn(f"Skipping unreadable slice {p}: {exc}")
if not slices:
return np.zeros((0, cfg.img_size, cfg.img_size), dtype=np.float32)
return np.stack(slices, axis=0)
# ============================================================
# METADATA INFERENCE (Anatomical_Plane / Fluid_Sensitive / Fat_Suppression)
# ============================================================
def infer_plane_from_orientation(iop) -> str:
"""Sagittal / Coronal / Axial from ImageOrientationPatient (LPS convention)."""
if iop is None or len(iop) < 6:
return "Sagittal"
row = np.asarray(iop[:3], dtype=np.float64)
col = np.asarray(iop[3:6], dtype=np.float64)
normal = np.cross(row, col)
axis = int(np.argmax(np.abs(normal)))
# normal dominant on X (L/R) -> Sagittal, Y (A/P) -> Coronal, Z (S/I) -> Axial
return {0: "Sagittal", 1: "Coronal", 2: "Axial"}[axis]
def infer_fluid_fat_from_text(text: str) -> Tuple[int, int]:
"""Best-effort keyword heuristic on SeriesDescription/ProtocolName/SequenceName."""
t = (text or "").lower()
fat_keywords = [
"fs", "fat sat", "fatsat", "fat-sat", "fatsupp",
"stir", "spair", "spir", "tirm",
]
fluid_keywords = [
"t2", "pd", "proton density", "stir", "fse t2", "frfse",
]
fat = int(any(k in t for k in fat_keywords))
fluid = int(any(k in t for k in fluid_keywords))
return fluid, fat
def read_series_metadata(sample_dcm_path: str) -> dict:
ds = pydicom.dcmread(sample_dcm_path, stop_before_pixels=True, force=True)
iop = getattr(ds, "ImageOrientationPatient", None)
plane = infer_plane_from_orientation(iop)
description = " ".join(
str(getattr(ds, tag, ""))
for tag in ("SeriesDescription", "ProtocolName", "SequenceName")
)
fluid, fat = infer_fluid_fat_from_text(description)
return {"plane": plane, "fluid": fluid, "fat": fat, "description": description.strip()}
def scan_series_metadata(
series_dir: Path,
metadata_overrides: Optional[Dict[str, dict]] = None,
) -> Optional[dict]:
"""Cheap metadata scan (reads only ONE header per series, no pixel decode)."""
paths = find_dicom_files(series_dir)
n_slices = len(paths)
if n_slices == 0:
return None
key = series_dir.name
if metadata_overrides and key in metadata_overrides:
row = metadata_overrides[key]
plane = str(row.get("Anatomical_Plane", "Sagittal")).strip().title()
if plane not in PLANES:
plane = "Sagittal"
fluid = int(pd.to_numeric(row.get("Fluid_Sensitive", 0), errors="coerce") or 0)
fat = int(pd.to_numeric(row.get("Fat_Suppression", 0), errors="coerce") or 0)
else:
try:
info = read_series_metadata(str(paths[0]))
plane = info["plane"] if info["plane"] in PLANES else "Sagittal"
fluid, fat = info["fluid"], info["fat"]
except Exception as exc:
warnings.warn(f"Could not read metadata for {series_dir}: {exc}")
plane, fluid, fat = "Sagittal", 0, 0
return {
"series_dir": series_dir,
"id": key,
"plane": plane,
"fluid": fluid,
"fat": fat,
"n_slices": n_slices,
}
# ============================================================
# FOLDER DISCOVERY
# ============================================================
def discover_series_dirs(root: Path) -> List[Path]:
"""Every directory (anywhere under root) that directly contains DICOM files."""
series_dirs = []
for dirpath, _dirnames, _filenames in os.walk(root):
d = Path(dirpath)
if find_dicom_files(d):
series_dirs.append(d)
return sorted(series_dirs)
def group_series_by_study(root: Path, series_dirs: List[Path]) -> Dict[str, List[Path]]:
"""
Group series folders into studies:
- series_dir directly under root -> single study named after root
- series_dir nested one level deeper (study/series/...) -> grouped by
the first path component below root
"""
groups: Dict[str, List[Path]] = {}
for s in series_dirs:
try:
rel_parts = s.relative_to(root).parts
except ValueError:
rel_parts = s.parts
if len(rel_parts) <= 1:
study_id = root.name or "study"
else:
study_id = rel_parts[0]
groups.setdefault(study_id, []).append(s)
return groups
# ============================================================
# SOFT SERIES SELECTION (same policy as training)
# ============================================================
def initial_series_score(s: dict) -> float:
score = 0.10 * math.log1p(s["n_slices"])
if s["plane"] in PLANES:
score += 0.30
if s["fluid"]:
score += 0.25
if s["fat"]:
score += 0.10
if s["fluid"] and s["fat"]:
score += 0.20
return float(score)
def select_series_for_study(series_list: List[dict], max_series: int) -> List[dict]:
if len(series_list) <= max_series:
return series_list
selected: List[dict] = []
used = set()
# Preserve plane diversity.
for plane in PLANES:
candidates = [s for s in series_list if s["plane"] == plane]
if candidates:
best = max(candidates, key=initial_series_score)
selected.append(best)
used.add(best["id"])
# Preserve sequence diversity.
for fluid, fat in [(1, 1), (1, 0), (0, 1), (0, 0)]:
if len(selected) >= max_series:
break
candidates = [
s for s in series_list
if s["fluid"] == fluid and s["fat"] == fat and s["id"] not in used
]
if candidates:
best = max(candidates, key=initial_series_score)
selected.append(best)
used.add(best["id"])
# Fill remaining slots by score.
remaining = [s for s in series_list if s["id"] not in used]
remaining.sort(key=initial_series_score, reverse=True)
for s in remaining:
if len(selected) >= max_series:
break
selected.append(s)
return selected[:max_series]
# ============================================================
# CLIP SAMPLING
# ============================================================
def sample_clip_centers(n_slices: int, n_clips: int) -> List[int]:
if n_slices <= 0:
return [0] * n_clips
if n_slices == 1:
return [0] * n_clips
if n_clips == 1:
return [n_slices // 2]
return np.linspace(0, n_slices - 1, n_clips).astype(int).tolist()
def make_clip(volume: np.ndarray, center: int, img_size: int) -> np.ndarray:
n = volume.shape[0]
if n == 0:
return np.zeros((3, img_size, img_size), dtype=np.float32)
center = int(np.clip(center, 0, n - 1))
idxs = [max(0, center - 1), center, min(n - 1, center + 1)]
return volume[idxs].astype(np.float32, copy=False)
def build_study_batch(
series_infos: List[dict], cfg: InferCFG
) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Decode pixels for the SELECTED series only and assemble one study's clips."""
images, planes, fluids, fats = [], [], [], []
for s in series_infos:
volume = load_series_volume(s["series_dir"], cfg)
centers = sample_clip_centers(volume.shape[0], cfg.clips_per_series)
for c in centers:
images.append(make_clip(volume, c, cfg.img_size))
planes.append(PLANE2IDX.get(s["plane"], 0))
fluids.append(float(s["fluid"]))
fats.append(float(s["fat"]))
if not images:
images = [np.zeros((3, cfg.img_size, cfg.img_size), dtype=np.float32)]
planes, fluids, fats = [0], [0.0], [0.0]
return (
np.stack(images, axis=0),
np.array(planes, dtype=np.int64),
np.array(fluids, dtype=np.float32),
np.array(fats, dtype=np.float32),
)
# ============================================================
# MODEL (identical architecture to training — required for checkpoint compatibility)
# ============================================================
class LabelAwareSeriesAttention(nn.Module):
def __init__(self, dim: int):
super().__init__()
self.dim = dim
self.query = nn.Parameter(torch.randn(N_LABELS, dim) * 0.02)
self.key = nn.Linear(dim, dim, bias=False)
self.value = nn.Linear(dim, dim, bias=False)
self.plane_embed = nn.Embedding(len(PLANES), dim)
self.fluid_proj = nn.Linear(1, dim)
self.fat_proj = nn.Linear(1, dim)
self.plane_bias = nn.Parameter(torch.zeros(N_LABELS, len(PLANES)))
self.fluid_bias = nn.Parameter(torch.zeros(N_LABELS))
self.fat_bias = nn.Parameter(torch.zeros(N_LABELS))
self.register_buffer("plane_prior", torch.tensor(PLANE_PRIOR, dtype=torch.float32))
self.register_buffer("fluid_prior", torch.tensor(FLUID_PRIOR, dtype=torch.float32))
self.register_buffer("fat_prior", torch.tensor(FAT_SUPPRESSION_PRIOR, dtype=torch.float32))
self.norm = nn.LayerNorm(dim)
def forward(self, series_feat, plane, fluid, fat):
_, d = series_feat.shape
meta = self.plane_embed(plane) + self.fluid_proj(fluid[:, None]) + self.fat_proj(fat[:, None])
z = self.norm(series_feat + meta)
k = self.key(z)
v = self.value(z)
scores = torch.matmul(self.query, k.transpose(0, 1)) / math.sqrt(d)
plane_soft = self.plane_prior[:, plane]
plane_learn = self.plane_bias[:, plane]
fluid_soft = self.fluid_prior[:, None] * fluid[None, :]
fluid_learn = self.fluid_bias[:, None] * fluid[None, :]
fat_soft = self.fat_prior[:, None] * fat[None, :]
fat_learn = self.fat_bias[:, None] * fat[None, :]
bias = plane_soft + plane_learn + fluid_soft + fluid_learn + fat_soft + fat_learn
scores = scores + 0.25 * bias
attn = F.softmax(scores, dim=-1)
pooled = torch.matmul(attn, v)
return pooled, attn
class LabelTokenTransformer(nn.Module):
def __init__(self, dim: int, heads: int, layers: int, dropout: float):
super().__init__()
self.label_embedding = nn.Parameter(torch.randn(N_LABELS, dim) * 0.02)
encoder_layer = nn.TransformerEncoderLayer(
d_model=dim, nhead=heads, dim_feedforward=dim * 4,
dropout=dropout, activation="gelu", batch_first=True, norm_first=True,
)
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=layers)
self.norm = nn.LayerNorm(dim)
def forward(self, label_features):
x = (label_features + self.label_embedding).unsqueeze(0)
x = self.transformer(x)
return self.norm(x.squeeze(0))
class VectorizedLabelHeads(nn.Module):
"""Equivalent to 12 independent heads, evaluated as batched tensor ops."""
def __init__(self, dim: int, hidden: int, dropout: float):
super().__init__()
self.eps = 1e-5
self.ln_weight = nn.Parameter(torch.ones(N_LABELS, dim))
self.ln_bias = nn.Parameter(torch.zeros(N_LABELS, dim))
self.fc1_weight = nn.Parameter(torch.empty(N_LABELS, hidden, dim))
self.fc1_bias = nn.Parameter(torch.zeros(N_LABELS, hidden))
self.fc2_weight = nn.Parameter(torch.empty(N_LABELS, hidden))
self.fc2_bias = nn.Parameter(torch.zeros(N_LABELS))
self.dropout = nn.Dropout(dropout)
nn.init.trunc_normal_(self.fc1_weight, std=0.02)
nn.init.trunc_normal_(self.fc2_weight, std=0.02)
def forward(self, x):
mean = x.mean(dim=-1, keepdim=True)
var = x.var(dim=-1, unbiased=False, keepdim=True)
z = (x - mean) / torch.sqrt(var + self.eps)
z = z * self.ln_weight + self.ln_bias
h = torch.einsum("nd,nhd->nh", z, self.fc1_weight) + self.fc1_bias
h = F.gelu(h)
h = self.dropout(h)
logits = torch.einsum("nh,nh->n", h, self.fc2_weight) + self.fc2_bias
return logits
class RSNAKneeModel(nn.Module):
def __init__(self, cfg: InferCFG, pretrained: bool = False):
super().__init__()
# pretrained=False by default: weights are restored from the checkpoint
# right after construction, so downloading ImageNet weights is unnecessary.
self.backbone = timm.create_model(
cfg.backbone_name, pretrained=pretrained, num_classes=0, global_pool="avg",
)
backbone_dim = self.backbone.num_features
self.proj = nn.Sequential(
nn.Linear(backbone_dim, cfg.embed_dim),
nn.LayerNorm(cfg.embed_dim),
nn.GELU(),
nn.Dropout(cfg.dropout),
)
self.series_attention = LabelAwareSeriesAttention(cfg.embed_dim)
self.label_transformer = LabelTokenTransformer(
cfg.embed_dim, cfg.label_transformer_heads, cfg.label_transformer_layers, cfg.dropout,
)
self.label_heads = VectorizedLabelHeads(cfg.embed_dim, cfg.embed_dim // 2, cfg.dropout)
def encode_images(self, images):
return self.proj(self.backbone(images))
def forward(self, batch):
images, plane, fluid, fat, offsets = (
batch["images"], batch["plane"], batch["fluid"], batch["fat"], batch["offsets"]
)
series_feat_all = self.encode_images(images)
all_logits, all_attn = [], []
for b in range(len(offsets) - 1):
s0, s1 = int(offsets[b]), int(offsets[b + 1])
label_feat, attn = self.series_attention(
series_feat_all[s0:s1], plane[s0:s1], fluid[s0:s1], fat[s0:s1],
)
label_feat = self.label_transformer(label_feat)
logits = self.label_heads(label_feat)
all_logits.append(logits)
all_attn.append(attn)
return torch.stack(all_logits, dim=0), all_attn
# ============================================================
# CHECKPOINT LOADING
# ============================================================
def _convert_legacy_label_heads(state: dict) -> dict:
"""Convert an old 12-ModuleList label-head checkpoint into the vectorized-head format."""
if "label_heads.0.0.weight" not in state or "label_heads.ln_weight" in state:
return state
state = dict(state)
try:
state["label_heads.ln_weight"] = torch.stack(
[state.pop(f"label_heads.{i}.0.weight") for i in range(N_LABELS)]
)
state["label_heads.ln_bias"] = torch.stack(
[state.pop(f"label_heads.{i}.0.bias") for i in range(N_LABELS)]
)
state["label_heads.fc1_weight"] = torch.stack(
[state.pop(f"label_heads.{i}.1.weight") for i in range(N_LABELS)]
)
state["label_heads.fc1_bias"] = torch.stack(
[state.pop(f"label_heads.{i}.1.bias") for i in range(N_LABELS)]
)
state["label_heads.fc2_weight"] = torch.stack(
[state.pop(f"label_heads.{i}.4.weight").squeeze(0) for i in range(N_LABELS)]
)
state["label_heads.fc2_bias"] = torch.stack(
[state.pop(f"label_heads.{i}.4.bias").squeeze(0) for i in range(N_LABELS)]
)
print("[CHECKPOINT] Converted legacy ModuleList label heads -> vectorized label heads.")
except Exception as exc:
print(f"[WARNING] Legacy label-head conversion failed: {exc}")
return state
def load_model_from_checkpoint(checkpoint_path: str, cfg: InferCFG, device: torch.device) -> nn.Module:
checkpoint_path = Path(checkpoint_path)
if not checkpoint_path.exists():
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
model = RSNAKneeModel(cfg, pretrained=False).to(device)
checkpoint = torch.load(checkpoint_path, map_location=device)
state = checkpoint.get("model", checkpoint) if isinstance(checkpoint, dict) else checkpoint
state = _convert_legacy_label_heads(state)
missing, unexpected = model.load_state_dict(state, strict=False)
if missing:
print(f"[CHECKPOINT] Missing keys: {len(missing)}")
if unexpected:
print(f"[CHECKPOINT] Unexpected keys: {len(unexpected)}")
if isinstance(checkpoint, dict) and "ema" in checkpoint:
named_params = dict(model.named_parameters())
for name, value in checkpoint["ema"].items():
if name in named_params:
named_params[name].data.copy_(value.to(device))
print("[CHECKPOINT] Applied EMA weights.")
model.eval()
return model
# ============================================================
# INFERENCE
# ============================================================
@torch.no_grad()
def predict_study(
model: nn.Module, series_infos: List[dict], cfg: InferCFG, device: torch.device,
) -> Tuple[Dict[str, float], np.ndarray]:
images, planes, fluids, fats = build_study_batch(series_infos, cfg)
x = torch.from_numpy(images)
x = (x - IMAGENET_MEAN[0]) / IMAGENET_STD[0]
x = x.to(device)
plane_t = torch.from_numpy(planes).to(device)
fluid_t = torch.from_numpy(fluids).to(device)
fat_t = torch.from_numpy(fats).to(device)
offsets = torch.tensor([0, x.shape[0]], dtype=torch.long)
with torch.amp.autocast(device_type=device.type, enabled=(cfg.amp and device.type == "cuda")):
logits, attn = model({
"images": x, "plane": plane_t, "fluid": fluid_t, "fat": fat_t, "offsets": offsets,
})
probs = torch.sigmoid(logits)[0].float().cpu().numpy()
result = {label: float(np.clip(p, 1e-5, 1.0 - 1e-5)) for label, p in zip(LABELS, probs)}
return result, attn[0].float().cpu().numpy()
def _load_metadata_overrides(metadata_csv: Optional[str]) -> Optional[Dict[str, dict]]:
if not metadata_csv:
return None
meta_df = pd.read_csv(metadata_csv)
if "SeriesFolder" not in meta_df.columns:
warnings.warn("--metadata-csv is missing a 'SeriesFolder' column; ignoring overrides.")
return None
return {str(row["SeriesFolder"]): row.to_dict() for _, row in meta_df.iterrows()}
def run_folder_inference(
input_folder: str,
checkpoint_path: str,
cfg: Optional[InferCFG] = None,
output_csv: Optional[str] = "predictions.csv",
metadata_csv: Optional[str] = None,
device_str: Optional[str] = None,
save_attention: bool = False,
limit_studies: Optional[int] = None,
) -> pd.DataFrame:
"""Discover every study under `input_folder`, run the model, return a predictions DataFrame."""
cfg = cfg or InferCFG()
root = Path(input_folder)
if not root.exists():
raise FileNotFoundError(f"Input folder not found: {root}")
series_dirs = discover_series_dirs(root)
if not series_dirs:
raise RuntimeError(f"No DICOM series found under: {root}")
study_groups = group_series_by_study(root, series_dirs)
print(f"Found {len(study_groups)} study(ies), {len(series_dirs)} total series under: {root}")
metadata_overrides = _load_metadata_overrides(metadata_csv)
device = torch.device(device_str or cfg.device)
model = load_model_from_checkpoint(checkpoint_path, cfg, device)
study_items = list(study_groups.items())
if limit_studies:
study_items = study_items[:limit_studies]
rows = []
attn_dump = {}
for study_id, s_dirs in tqdm(study_items, desc="studies"):
infos = []
for sd in s_dirs:
meta = scan_series_metadata(sd, metadata_overrides)
if meta:
infos.append(meta)
if not infos:
warnings.warn(f"Study '{study_id}': no readable series, skipping.")
continue
selected = select_series_for_study(infos, cfg.max_series)
print(
f"\nStudy '{study_id}': using {len(selected)}/{len(infos)} series "
f"(planes={[s['plane'] for s in selected]}, "
f"fluid={[s['fluid'] for s in selected]}, fat={[s['fat'] for s in selected]})"
)
probs, attn = predict_study(model, selected, cfg, device)
row = {"study_id": study_id}
row.update(probs)
rows.append(row)
if save_attention:
attn_dump[study_id] = attn.tolist()
result_df = (
pd.DataFrame(rows)[["study_id"] + LABELS]
if rows else pd.DataFrame(columns=["study_id"] + LABELS)
)
if output_csv:
result_df.to_csv(output_csv, index=False)
print(f"\nSaved predictions: {output_csv}")
if save_attention and attn_dump:
attn_path = str(Path(output_csv).with_suffix(".attention.json")) if output_csv else "attention.json"
with open(attn_path, "w", encoding="utf-8") as f:
json.dump(attn_dump, f)
print(f"Saved attention maps: {attn_path}")
return result_df
# ============================================================
# MAIN — chạy thẳng bằng các biến trong khối "USER CONFIG" phía trên
# ============================================================
def main():
cfg = InferCFG(
backbone_name=BACKBONE_NAME,
img_size=IMG_SIZE,
embed_dim=EMBED_DIM,
max_series=MAX_SERIES,
clips_per_series=CLIPS_PER_SERIES,
)
df = run_folder_inference(
input_folder=INPUT_FOLDER,
checkpoint_path=CHECKPOINT_PATH,
cfg=cfg,
output_csv=OUTPUT_CSV,
metadata_csv=METADATA_CSV,
device_str=DEVICE,
save_attention=SAVE_ATTENTION,
limit_studies=LIMIT_STUDIES,
)
print("\n" + "=" * 90)
print("PREDICTIONS")
print("=" * 90)
print(df.to_string(index=False))
if __name__ == "__main__":
main()