File size: 7,285 Bytes
a23d562 6524cf7 a23d562 6524cf7 a23d562 6524cf7 a23d562 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | """
Produce out-of-fold predictions with case identifiers for the AP vs FDP non-inferiority analysis.
For one outer fold this runs the abbreviated (AP) and the full (FDP) model of one architecture over that
fold's test set and writes one row per breast:
outer_fold, examination_id, side, label, protocol, score
`score` is the raw logit of the malignancy output, so it is monotonically related to the predicted
probability and can be used for AUC and for thresholding without further transformation.
Run this once per (architecture, fold); merge_oof_predictions.py then assembles the wide table.
Unlike test_training.py this records the case identifier of every sample and keeps the last, incomplete
batch, so no case is silently dropped or padded with a zero label.
"""
import argparse
import csv
from pathlib import Path
import torch
from auto_detect_breast_mri.config import resolve_path
from auto_detect_breast_mri.data import loaders
from auto_detect_breast_mri.data.metadata import get_uka_metatensor
from auto_detect_breast_mri.data.breast_mri_dataset import split_patient_key
from auto_detect_breast_mri.data.transforms import default_transform
from auto_detect_breast_mri.models.resnets import model_names
CPU = "cpu"
GPU = "cuda"
PROTOCOL_OF_SUFFIX = {'abrv': 'abbreviated', 'full': 'full'}
def load_pretrained_model(model, model_path, device):
"""Load a checkpoint that is either a bare state dict or wrapped in a 'state_dict' entry."""
checkpoint = torch.load(model_path, map_location=device, weights_only=False)
state_dict = checkpoint['state_dict'] if 'state_dict' in checkpoint else checkpoint
state_dict = {key.replace('module.', ''): value for key, value in state_dict.items()}
missing, unexpected = model.load_state_dict(state_dict, strict=False)
if missing or unexpected:
raise RuntimeError(f"Checkpoint {model_path} does not match the architecture. "
f"Missing keys: {sorted(missing)[:5]}, unexpected keys: {sorted(unexpected)[:5]}")
return model
@torch.no_grad()
def predict(model, data_loader, device):
"""
Run the model over every sample of data_loader.
:return: list of (patient_key, label, score) with score being the raw logit of the malignancy output
"""
model.eval()
predictions = []
for batch_id, subject in enumerate(data_loader):
data = subject["image"]["data"].to(device)
with torch.autocast(device_type=device, dtype=torch.float16):
logits = model(data)[:, 0]
logits = logits.float().cpu()
ground_truth = subject["label"].cpu()
for key, label, score in zip(subject["path"], ground_truth.tolist(), logits.tolist()):
predictions.append((key, int(label), float(score)))
if batch_id % 20 == 0:
print(f" batch {batch_id}/{len(data_loader)}: {len(predictions)} cases predicted", flush=True)
return predictions
def main():
parser = argparse.ArgumentParser(
prog="predict_oof",
description="Write per case out-of-fold predictions of the AP and FDP model of one architecture.")
parser.add_argument("architecture", choices=["resnet18", "resnet50"])
parser.add_argument("data_path",
help="Root folder(s) of the cropped MRIs. Several roots may be given comma separated.")
parser.add_argument("feature_path", help="CSV/XLSX file holding the label information, "
"read with get_uka_metatensor mode 0 as during training.")
parser.add_argument("split_files_folder",
help="Folder holding fold<k>/stratified_test_set-f<k>.csv")
parser.add_argument("-c", "--fold", type=int, required=True, help="Outer fold to predict.")
parser.add_argument("-m", "--model_path_pattern", required=True,
help="Checkpoint path with {model_key} and {fold} placeholders, e.g. "
"'<model_root>/{model_key}_b=32_l=0.0001_n=60_t=basic_pretrain=False_split{fold}-0.pth'")
parser.add_argument("-b", "--batch_size", type=int, default=16)
parser.add_argument("-f", "--fraction", type=float)
parser.add_argument("-o", "--output_path", default=None,
help="Folder the per fold prediction CSV is written to.")
parser.add_argument("-s", "--image_shape", type=int, nargs=3, default=(256, 256, 32),
help="Must match the shape the checkpoints were trained with (training.py uses 256 256 32).")
args = parser.parse_args()
device = GPU if torch.cuda.is_available() else CPU
print(f"Use device: {device}")
features = get_uka_metatensor(0, args.feature_path)
output_dir = Path(resolve_path(args.output_path, "output_root", "output folder")) / "oof"
output_dir.mkdir(parents=True, exist_ok=True)
# train.py writes an explicit fraction into every checkpoint name, full data included
# (frac_string turns the internal -1.0 sentinel back into 1.0), so the suffix is always
# present. Keep the float form: '_frac=1' would not match a '_frac=1.0_FINAL.pth'.
fraction = float(args.fraction) if args.fraction else 1.0
frac_suffix = f'_frac={fraction}'
# The fraction belongs in the file name. Without it every fraction of a sweep writes to the same
# five files, the last job silently wins, and the resulting folder cannot be told apart from a
# full-data run afterwards -- which is how a folder of weak predictions ends up in the analysis.
output_file = output_dir / f"oof_{args.architecture}_fold{args.fold}{frac_suffix}.csv"
rows = []
for suffix, protocol in PROTOCOL_OF_SUFFIX.items():
model_key = f"{args.architecture}_{suffix}"
model = model_names.get(model_key)
if model is None:
raise ValueError(f"Unknown model key {model_key}. Available: {sorted(model_names)}")
model.to(device)
checkpoint = args.model_path_pattern.format(model_key=model_key, fold=args.fold, frac_suffix=frac_suffix)
print(f"Load {checkpoint}")
model = load_pretrained_model(model, checkpoint, device)
loader = loaders.get_evaluation_dataloader(
args.data_path, features, tuple(args.image_shape), protocol,
data_selection_file_sceleton=str(Path(args.split_files_folder) / f"fold{args.fold}" / "stratified_test_set"),
batch_size=args.batch_size, fold=args.fold,
transform=default_transform(tuple(args.image_shape)))
print(f"{model_key} fold {args.fold}: {len(loader.dataset)} cases")
for patient_key, label, score in predict(model, loader, device):
examination_id, side = split_patient_key(patient_key)
rows.append({'outer_fold': args.fold, 'examination_id': examination_id, 'side': side,
'label': label, 'protocol': suffix, 'score': score})
model.cpu()
torch.cuda.empty_cache()
with open(output_file, 'w', newline='') as handle:
writer = csv.DictWriter(handle, fieldnames=['outer_fold', 'examination_id', 'side', 'label', 'protocol', 'score'])
writer.writeheader()
writer.writerows(rows)
print(f"Wrote {len(rows)} rows to {output_file}")
if __name__ == '__main__':
main()
|