File size: 6,001 Bytes
bd99b1b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 | """
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()
|