""" Analyse cross-protocol predictions and write a compact paper table. The input files are written by cross_protocol_evaluation.py. This table keeps only the two cross-protocol deployments: * an AP-trained model evaluated on FDP input; * an FDP-trained model evaluated on AP input. Usage: python scripts/validation/analyze_cross_protocol.py """ import argparse from pathlib import Path import numpy as np import pandas as pd from auto_detect_breast_mri.config import resolve_path from auto_detect_breast_mri.evaluation.subgroups import bootstrap_metrics from scripts.validation.cross_protocol_evaluation import read_scores CROSS_SETTINGS = ( ('abrv', 'FDP'), ('full', 'AP'), ) MODEL_LABELS = {'abrv': 'AP', 'full': 'FDP'} ARCHITECTURE_LABELS = {'resnet18': 'ResNet18', 'resnet50': 'ResNet50'} def format_interval(value, low, high, decimals): """Format one metric as ``mean (lower-upper)``.""" if any(value is None or np.isnan(value) for value in (value, low, high)): return '-' return f'{value:.{decimals}f} ({low:.{decimals}f}-{high:.{decimals}f})' def generate_cross_protocol_table(table, target=0.9, replications=2000, weighting='pairs', seed=0, stratify=True, decimals=3): """Return the cross-protocol paper table and notes. :param table: long-form table returned by ``cross_protocol_evaluation.read_scores`` :param target: specificity at which sensitivity is reported :return: ``(DataFrame, notes)`` """ headers = ['Model', 'validation protocol', 'AUC (95% CI)', 'Sensitivity (95% CI)', 'Specificity (95% CI)'] rows, notes = [], [] for architecture in sorted(table['architecture'].unique()): part = table[table['architecture'] == architecture] architecture_label = ARCHITECTURE_LABELS.get(architecture, architecture) for model_protocol, input_protocol in CROSS_SETTINGS: selected = part[(part['model_protocol'] == model_protocol) & (part['input_protocol'] == input_protocol)] if selected.empty: notes.append(f'{architecture_label}: no scores for {MODEL_LABELS[model_protocol]} ' f'model on {input_protocol} input') continue print(f' {architecture_label}: {MODEL_LABELS[model_protocol]} model on ' f'{input_protocol} input ({len(selected)} breasts)', flush=True) statistics = bootstrap_metrics( selected['outer_fold'].to_numpy(), selected['patient_id'].to_numpy(), selected['label'].to_numpy(), selected['score'].to_numpy(), target=target, replications=replications, weighting=weighting, seed=seed, stratify=stratify) row = { 'Model': architecture_label, 'validation protocol': input_protocol, 'AUC (95% CI)': format_interval(statistics['auc']['value'], statistics['auc']['ci_low'], statistics['auc']['ci_high'], decimals), 'Sensitivity (95% CI)': format_interval( statistics['sens_at_spec']['value'], statistics['sens_at_spec']['ci_low'], statistics['sens_at_spec']['ci_high'], decimals), 'Specificity (95% CI)': format_interval( statistics['spec_at_sens']['value'], statistics['spec_at_sens']['ci_low'], statistics['spec_at_sens']['ci_high'], decimals), } rows.append(row) if statistics['auc']['folds_used'] < 5: notes.append(f'{architecture_label} / {input_protocol}: AUC uses only ' f"{statistics['auc']['folds_used']} folds") return pd.DataFrame(rows, columns=headers), notes def build_parser(): parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) parser.add_argument('score_dir', help='Folder holding cross_protocol_*.csv files.') parser.add_argument('metadata_file', help='Metadata export used for patient clustering.') parser.add_argument('-o', '--output', default=None, help='Output CSV path. Default: cross_protocol_results.csv beside score_dir.') parser.add_argument('--architecture', choices=['resnet18', 'resnet50'], default=None) parser.add_argument('--fraction', type=float, default=None) parser.add_argument('--operating_point', type=float, default=0.9, help='Specificity target for sensitivity. Default: %(default)s') parser.add_argument('-r', '--replications', type=int, default=2000) parser.add_argument('-w', '--weighting', choices=['pairs', 'cases', 'equal'], default='pairs') parser.add_argument('--seed', type=int, default=0) parser.add_argument('--no_stratify', action='store_true') parser.add_argument('-d', '--decimals', type=int, default=3) return parser def main(): args = build_parser().parse_args() table, report = read_scores(args.score_dir, args.metadata_file, args.architecture, args.fraction) print(f'Bootstrapping with {args.replications} replications ...') result, notes = generate_cross_protocol_table( table, target=args.operating_point, replications=args.replications, weighting=args.weighting, seed=args.seed, stratify=not args.no_stratify, decimals=args.decimals) output = Path(args.output) if args.output else Path(args.score_dir) / 'cross_protocol_results.csv' output.parent.mkdir(parents=True, exist_ok=True) result.to_csv(output, index=False) print(result.to_string(index=False)) for note in notes: print(f'NOTE: {note}') print(f'Wrote {output}') if __name__ == '__main__': main()