""" Cross-protocol evaluation: a model trained on one protocol applied to the other. The deployment case this answers is a model developed on archived full protocol (FDP) examinations and then applied to abbreviated (AP) acquisitions, where Sub_2, Sub_3, Sub_4 and T2 simply do not exist. The reverse direction is not symmetric: abbreviated = [Dyn_0, Sub_1] full = [Dyn_0, Sub_1, Sub_2, Sub_3, Sub_4, T2] AP is a strict subset of FDP, so an AP model applied to a full acquisition just reads the two sequences it was trained on and nothing has to be substituted -- that direction is reported for completeness but is a no-op by construction. An FDP model applied to an abbreviated acquisition is the real question, because four of its six input channels are missing and something has to stand in for them. Which substitute is used is a modelling decision, not a detail, so it is explicit: zero the missing channels are set to 0, i.e. the mean of the z-normalised image repeat Sub_2/3/4 are filled with Sub_1 and T2 with Dyn_0, a naive carry forward mean every missing channel is filled with the voxelwise mean of the two available ones Two steps, because the first needs a GPU and the second does not: # once per (architecture, fold), writes cross_protocol__fold[_frac=].csv python scripts/validation/cross_protocol_evaluation.py predict resnet18 \\ -c 0 -m '' # once, over all folds python scripts/validation/cross_protocol_evaluation.py analyse The scores of every setting are produced in one pass over the same loaded volumes, so the comparison is paired case by case and free of crop or augmentation differences between settings. Note that AP-native scores here are recomputed from the full protocol tensor rather than taken from predict_oof.py, so they can differ marginally through the random noise in the inference transform. """ import argparse import csv import re from collections import defaultdict from pathlib import Path import numpy as np import pandas as pd 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.breast_mri_dataset import protocol_mappings, split_patient_key from auto_detect_breast_mri.data.metadata import get_uka_metatensor from auto_detect_breast_mri.data.transforms import default_transform from auto_detect_breast_mri.evaluation.cluster_bootstrap import (FoldClusterBootstrap, percentile_interval, weighted_fold_difference) from auto_detect_breast_mri.evaluation.oof_table import load_patient_map from auto_detect_breast_mri.evaluation.subgroups import bootstrap_metrics, weighted_fold_metric from auto_detect_breast_mri.models.resnets import model_names CPU, GPU = "cpu", "cuda" FULL_SEQUENCES = protocol_mappings['full'] ABRV_SEQUENCES = protocol_mappings['abbreviated'] # index of every abbreviated sequence inside the full channel stack, and of the ones it lacks KEPT = [FULL_SEQUENCES.index(name) for name in ABRV_SEQUENCES] MISSING = [index for index in range(len(FULL_SEQUENCES)) if index not in KEPT] SUBSTITUTES = ('zero', 'repeat', 'mean') FILE_PATTERN = re.compile(r'^cross_protocol_(?Presnet\d+)_fold(?P\d+)' r'(?:_frac=(?P[0-9.]+))?\.csv$') # (model, the protocol actually available at inference): whether it is the native or the cross case SETTINGS = [('full', 'FDP', 'native'), ('full', 'AP', 'cross'), ('abrv', 'AP', 'native'), ('abrv', 'FDP', 'native-by-subset')] def load_checkpoint(model, model_path, device): """ Load a checkpoint that is either a bare state dict or wrapped in a 'state_dict' entry. Same loader as predict_oof.py: models.checkpoints.load_pretrained_model returns a state dict rather than a model and only handles the wrapped form, so it cannot be used here. """ 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: {sorted(unexpected)[:5]}") return model def abbreviate(volume, substitute): """ Turn a full protocol tensor into what an abbreviated acquisition would have delivered. :param volume: (B, C, ...) with C = len(FULL_SEQUENCES) :param substitute: what stands in for the sequences an abbreviated exam does not contain :return: a copy with the missing channels replaced """ reduced = volume.clone() if substitute == 'zero': reduced[:, MISSING] = 0.0 elif substitute == 'repeat': # Sub_2/3/4 carry Sub_1 forward, T2 falls back to the native scan source = {FULL_SEQUENCES.index('Sub_2'): FULL_SEQUENCES.index('Sub_1'), FULL_SEQUENCES.index('Sub_3'): FULL_SEQUENCES.index('Sub_1'), FULL_SEQUENCES.index('Sub_4'): FULL_SEQUENCES.index('Sub_1'), FULL_SEQUENCES.index('T2'): FULL_SEQUENCES.index('Dyn_0')} for target, origin in source.items(): reduced[:, target] = volume[:, origin] elif substitute == 'mean': reduced[:, MISSING] = volume[:, KEPT].mean(dim=1, keepdim=True) else: raise ValueError(f"Unknown substitute '{substitute}'. Use one of {SUBSTITUTES}.") return reduced @torch.no_grad() def predict_settings(models, data_loader, device, substitute): """ Score every case under every setting, from one pass over the loaded volumes. :param models: dict protocol suffix ('abrv'/'full') -> model, already on `device` :return: list of row dicts """ for model in models.values(): model.eval() rows = [] for batch_id, subject in enumerate(data_loader): volume = subject['image']['data'].to(device) labels = subject['label'].cpu().tolist() keys = subject['path'] scored = {} with torch.autocast(device_type=device, dtype=torch.float16): if 'full' in models: scored[('full', 'FDP')] = models['full'](volume)[:, 0] scored[('full', 'AP')] = models['full'](abbreviate(volume, substitute))[:, 0] if 'abrv' in models: # the abbreviated model only ever sees the sequences it was trained on, so the # full acquisition gives it exactly the same tensor as an abbreviated one available = models['abrv'](volume[:, KEPT])[:, 0] scored[('abrv', 'AP')] = available scored[('abrv', 'FDP')] = available for (suffix, available_protocol), logits in scored.items(): role = next(role for model, protocol, role in SETTINGS if model == suffix and protocol == available_protocol) for key, label, score in zip(keys, labels, logits.float().cpu().tolist()): examination_id, side = split_patient_key(key) rows.append({'examination_id': examination_id, 'side': side, 'label': int(label), 'model_protocol': suffix, 'input_protocol': available_protocol, 'role': role, 'score': float(score)}) if batch_id % 20 == 0: print(f" batch {batch_id}/{len(data_loader)}", flush=True) return rows def run_prediction(args): """Score one fold under every setting and write the per case CSV.""" device = GPU if torch.cuda.is_available() else CPU print(f"Use device: {device}") features = get_uka_metatensor(0, args.feature_path) fraction = float(args.fraction) if args.fraction else 1.0 frac_suffix = f'_frac={fraction}' output_dir = Path(resolve_path(args.output_path, "output_root", "output folder")) / "cross_protocol" output_dir.mkdir(parents=True, exist_ok=True) output_file = output_dir / f"cross_protocol_{args.architecture}_fold{args.fold}{frac_suffix}.csv" models = {} for suffix in ('abrv', 'full'): 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)}") checkpoint = args.model_path_pattern.format(model_key=model_key, fold=args.fold, frac_suffix=frac_suffix) print(f"Load {checkpoint}") models[suffix] = load_checkpoint(model, checkpoint, device).to(device) # always load the full protocol: every setting is derived from that one tensor loader = loaders.get_evaluation_dataloader( args.data_path, features, tuple(args.image_shape), 'full', 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"fold {args.fold}: {len(loader.dataset)} cases, substitute '{args.substitute}'") rows = predict_settings(models, loader, device, args.substitute) for row in rows: row['outer_fold'] = args.fold row['substitute'] = args.substitute fields = ['outer_fold', 'examination_id', 'side', 'label', 'model_protocol', 'input_protocol', 'role', 'substitute', 'score'] with open(output_file, 'w', newline='') as handle: writer = csv.DictWriter(handle, fieldnames=fields) writer.writeheader() writer.writerows(rows) print(f"Wrote {len(rows)} rows to {output_file}") def read_scores(score_dirs, metadata_file, architecture=None, fraction=None): """ The per case scores of every fold, with the woman each examination belongs to attached. :param score_dirs: one folder, or several when the architectures were predicted into separate output roots; all of them are read into one table :return: (DataFrame, dict describing what was read) """ if isinstance(score_dirs, (str, Path)): score_dirs = [score_dirs] paths = [] for score_dir in score_dirs: score_dir = Path(score_dir) if not score_dir.is_dir(): raise FileNotFoundError(f"{score_dir} is not a directory.") paths.extend(sorted(score_dir.iterdir())) frames, report = [], {} for path in paths: parsed = FILE_PATTERN.match(path.name) if parsed is None: continue if architecture and parsed['architecture'] != architecture: continue found = float(parsed['fraction']) if parsed['fraction'] else 1.0 if fraction is not None and found != float(fraction): continue frame = pd.read_csv(path, dtype={'examination_id': str, 'side': str}) frame['architecture'] = parsed['architecture'] frame['fraction'] = found frames.append(frame) if not frames: raise FileNotFoundError(f"No cross_protocol_*.csv in " f"{', '.join(str(d) for d in score_dirs)} for the requested " f"architecture/fraction. Run the 'predict' step first.") table = pd.concat(frames, ignore_index=True) for column in ('architecture', 'fraction', 'substitute'): report[column] = sorted(table[column].unique().tolist()) if len(report['fraction']) > 1: raise ValueError(f"The given folders mix training fractions {report['fraction']}; those " f"are different models. Pass --fraction to pick one.") if len(report['substitute']) > 1: raise ValueError(f"The given folders mix substitutes {report['substitute']}; keep one " f"substitute per analysis run.") patient_map = load_patient_map(metadata_file, table['examination_id'].unique()) table['patient_id'] = table['examination_id'].map(patient_map) report['dropped_without_patient_id'] = int(table['patient_id'].isna().sum()) table = table[table['patient_id'].notna()].reset_index(drop=True) report['breasts'] = int(len(table) / table.groupby(['model_protocol', 'input_protocol']).ngroups) return table, report def paired_difference(wide, column_a, column_b, replications, weighting, seed, stratify): """ Interval for the paired difference column_a - column_b, both scored on the same women. :param wide: one row per breast, holding both score columns :return: dict with the observed difference and its interval """ bootstrap = FoldClusterBootstrap(wide['outer_fold'].to_numpy(), wide['patient_id'].to_numpy(), wide['label'].to_numpy(), stratify_by_outcome=stratify) weights = bootstrap.fold_weights(weighting) labels = wide['label'].to_numpy() scores_a, scores_b = wide[column_a].to_numpy(), wide[column_b].to_numpy() observed, _ = weighted_fold_difference(bootstrap.observed_rows(), labels, scores_a, scores_b, weights) rng = np.random.default_rng(seed) replicated = np.full(replications, np.nan) for replication in range(replications): drawn = bootstrap.resample_rows(rng) replicated[replication], _ = weighted_fold_difference(drawn, labels, scores_a, scores_b, weights) low, high = percentile_interval(replicated, 0.95) return {'difference': observed, 'ci_low': low, 'ci_high': high} def analyse(table, target=0.9, replications=2000, weighting='pairs', seed=0, stratify=True): """ AUC of every setting with its interval, plus the paired cost of the missing sequences. :return: (list of per setting rows, list of paired comparison rows) """ settings, paired = [], [] for architecture in sorted(table['architecture'].unique()): part = table[table['architecture'] == architecture] for model_protocol, input_protocol, role in SETTINGS: rows = part[(part['model_protocol'] == model_protocol) & (part['input_protocol'] == input_protocol)] if rows.empty: continue print(f" {architecture} {model_protocol} model on {input_protocol} input " f"({len(rows)} breasts)", flush=True) statistics = bootstrap_metrics(rows['outer_fold'].to_numpy(), rows['patient_id'].to_numpy(), rows['label'].to_numpy(), rows['score'].to_numpy(), target=target, replications=replications, weighting=weighting, seed=seed, stratify=stratify) settings.append({'architecture': architecture, 'model_protocol': model_protocol, 'input_protocol': input_protocol, 'role': role, 'breasts': len(rows), **{f'{metric}_{field}': value for metric, values in statistics.items() for field, value in values.items()}}) # the cost of losing the four sequences, paired case by case on the FDP model wide = part.pivot_table(index=['outer_fold', 'examination_id', 'side', 'label', 'patient_id'], columns=['model_protocol', 'input_protocol'], values='score').reset_index() wide.columns = ['_'.join(part for part in column if part).strip('_') if isinstance(column, tuple) else column for column in wide.columns] comparisons = [('full_AP', 'full_FDP', 'FDP model: abbreviated input minus full input'), ('full_AP', 'abrv_AP', 'abbreviated input: FDP model minus AP model')] for column_a, column_b, description in comparisons: if column_a not in wide.columns or column_b not in wide.columns: continue print(f" {architecture} paired: {description}", flush=True) result = paired_difference(wide, column_a, column_b, replications, weighting, seed, stratify) paired.append({'architecture': architecture, 'comparison': description, **result}) return settings, paired def format_report(settings, paired, report, target, replications, decimals=3): lines = [] add = lines.append add("=" * 100) add("Cross-protocol evaluation: each model applied to the other protocol's input") add("=" * 100) add("") add(f"{'architectures':<28}{', '.join(report['architecture'])}") add(f"{'training fraction':<28}{', '.join(f'{value:g}' for value in report['fraction'])}") add(f"{'substitute for missing seq.':<28}{', '.join(report['substitute'])}") add(f"{'breasts per setting':<28}{report['breasts']}") add(f"{'bootstrap replications':<28}{replications}") if report.get('dropped_without_patient_id'): add(f"{'dropped, no patient ID':<28}{report['dropped_without_patient_id']}") add("") add("-" * 100) add("1. Every setting") add("-" * 100) add(f"{'architecture':<12} {'model':>6} {'input':>6} {'role':>18} {'AUC':>7} " f"{'95% CI':>20} {f'sens@{target:.0%}spec':>15}") for row in settings: interval = f"[{row['auc_ci_low']:.{decimals}f}, {row['auc_ci_high']:.{decimals}f}]" add(f"{row['architecture']:<12} {row['model_protocol']:>6} {row['input_protocol']:>6} " f"{row['role']:>18} {row['auc_value']:>7.{decimals}f} {interval:>20} " f"{row['sens_at_spec_value']:>15.{decimals}f}") add("") add("'native-by-subset' is not a separate experiment: the abbreviated model reads only Dyn_0") add("and Sub_1, which a full acquisition also contains, so its score cannot change.") add("") add("-" * 100) add("2. Paired differences, same women in every bootstrap replication") add("-" * 100) for row in paired: add(f"{row['architecture']:<12} {row['comparison']:<48} " f"{row['difference']:>+.{decimals}f} " f"[{row['ci_low']:+.{decimals}f}, {row['ci_high']:+.{decimals}f}]") add("") add("-" * 100) add("How to read this") add("-" * 100) add("The first paired row is the price of deploying an FDP trained model on abbreviated data,") add("with the missing sequences filled in as stated above. The second asks whether, given only") add("an abbreviated acquisition, one is better off with the FDP model plus substitution or with") add("a model actually trained on AP. A negative value favours the second term of the pair.") add("The substitute is a modelling choice: rerun with --substitute to see how much it matters.") return "\n".join(lines) def build_parser(): parser = argparse.ArgumentParser(prog="cross_protocol_evaluation", description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) subparsers = parser.add_subparsers(dest="step", required=True) predict = subparsers.add_parser("predict", help="score one fold under every setting (GPU)") predict.add_argument("architecture", choices=["resnet18", "resnet50"]) predict.add_argument("data_path", help="Root folder(s) of the cropped MRIs, comma separated.") predict.add_argument("feature_path", help="CSV/XLSX holding the label information.") predict.add_argument("split_files_folder", help="Folder holding fold/stratified_test_set-f.csv") predict.add_argument("-c", "--fold", type=int, required=True) predict.add_argument("-m", "--model_path_pattern", required=True, help="Checkpoint path with {model_key}, {fold} and {frac_suffix}.") predict.add_argument("-f", "--fraction", type=float, default=None, help="Training fraction of the checkpoints. Default: 1.0") predict.add_argument("-b", "--batch_size", type=int, default=16) predict.add_argument("-o", "--output_path", default=None) predict.add_argument("-s", "--image_shape", type=int, nargs=3, default=(256, 256, 32)) predict.add_argument("--substitute", choices=SUBSTITUTES, default='zero', help="What stands in for the sequences an abbreviated exam lacks. " "Default: zero") analyse_parser = subparsers.add_parser("analyse", help="aggregate the per fold score files") analyse_parser.add_argument("score_dirs", nargs='+', metavar="SCORE_DIR", help="Folder(s) holding cross_protocol_*.csv. Give several when " "the architectures were predicted into separate output roots.") analyse_parser.add_argument("metadata_file", help="Metadata export, for the patient clustering.") analyse_parser.add_argument("-o", "--output_path", default=None) analyse_parser.add_argument("--architecture", default=None, choices=["resnet18", "resnet50"]) analyse_parser.add_argument("--fraction", type=float, default=None) analyse_parser.add_argument("--operating_point", type=float, default=0.9) analyse_parser.add_argument("-r", "--replications", type=int, default=2000) analyse_parser.add_argument("-w", "--weighting", choices=["pairs", "cases", "equal"], default="pairs") analyse_parser.add_argument("--seed", type=int, default=0) analyse_parser.add_argument("--no_stratify", action="store_true") analyse_parser.add_argument("-d", "--decimals", type=int, default=3) return parser def main(): args = build_parser().parse_args() if args.step == "predict": run_prediction(args) return table, report = read_scores(args.score_dirs, args.metadata_file, args.architecture, args.fraction) print(f"Bootstrapping with {args.replications} replications ...") settings, paired = analyse(table, target=args.operating_point, replications=args.replications, weighting=args.weighting, seed=args.seed, stratify=not args.no_stratify) output_dir = Path(resolve_path(args.output_path, "output_root", "output folder")) / "cross_protocol" output_dir.mkdir(parents=True, exist_ok=True) pd.DataFrame(settings).to_csv(output_dir / "cross_protocol_settings.csv", index=False) pd.DataFrame(paired).to_csv(output_dir / "cross_protocol_paired.csv", index=False) text = format_report(settings, paired, report, args.operating_point, args.replications, args.decimals) (output_dir / "cross_protocol_report.txt").write_text(text) print() print(text) print(f"\nWrote the report and two CSVs to {output_dir}") if __name__ == '__main__': main()