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