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()