Download scripts/validation/cross_protocol_evaluation.py from deboraJ23/AI_MRI: direct link, hf CLI and curl.
- Browser
- Download file 23.5 kB
-
https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/cross_protocol_evaluation.py
- Command line
-
hf download hf://deboraJ23/AI_MRI/scripts/validation/cross_protocol_evaluation.py
-
curl -L -o cross_protocol_evaluation.py https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/cross_protocol_evaluation.py
23.5 kB
| """ | |
| 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_<arch>_fold<k>[_frac=<f>].csv | |
| python scripts/validation/cross_protocol_evaluation.py predict resnet18 <data> <metadata> \\ | |
| <splits> -c 0 -m '<checkpoint pattern>' | |
| # once, over all folds | |
| python scripts/validation/cross_protocol_evaluation.py analyse <score_dir> <metadata> | |
| 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_(?P<architecture>resnet\d+)_fold(?P<fold>\d+)' | |
| r'(?:_frac=(?P<fraction>[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 | |
| 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<k>/stratified_test_set-f<k>.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() | |