""" 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/stratified_test_set-f.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_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()