AI_MRI / scripts /validation /analyze_cross_protocol.py
DeboraJ1's picture
include scripts for paper table generation, subgroup analysis and cross protocol evaluation
bd99b1b
Raw History Blame Contribute Delete
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()