Download scripts/validation/predict_oof.py from deboraJ23/AI_MRI: direct link, hf CLI and curl.
- Browser
- Download file 7.29 kB
-
https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/predict_oof.py
- Command line
-
hf download hf://deboraJ23/AI_MRI/scripts/validation/predict_oof.py
-
curl -L -o predict_oof.py https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/predict_oof.py
7.29 kB
| """ | |
| 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 | |
| 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() | |