File size: 10,280 Bytes
6524cf7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bd99b1b
 
 
6524cf7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bd99b1b
 
6524cf7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
"""
Exploratory performance per indication subgroup on the out-of-fold predictions.

    python scripts/validation/analyze_subgroups.py <prediction_dir> <metadata_file>

Reads the same out-of-fold table as analyze_noninferiority.py, attaches the per-side indication code
from the metadata export, and reports for every indication: how many breasts and women it holds, the
AUC of the abbreviated (AP) and the full (FDP) model of every architecture, and the paired AP minus
FDP difference with a patient-clustered bootstrap interval.

This is exploratory. The subgroups are small, the intervals are not adjusted for multiplicity across
them, and no non-inferiority verdict is derived: a wide interval means the subgroup cannot resolve
the difference, not that AP is inferior in it. The primary analysis stays the one over all breasts.

The indication codes are reported as they stand in the metadata. Add an `indication_labels:` block
to the site config to give them readable names:

    indication_labels:
      0: screening
      1: diagnostic
"""

import argparse
import json
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.oof_table import (architectures_of, attach_indication,
                                                         check_integrity, load_oof_table)
from auto_detect_breast_mri.evaluation.subgroups import MIN_BREASTS, MIN_PER_CLASS, analyse_by_group, to_frame

UNKNOWN = 'unknown'


def format_report(results, architectures, load_report, indication_report, replications, weighting,
                  dropped_unknown=False):
    lines = []

    def add(text=""):
        lines.append(text)

    add("=" * 100)
    add("Performance per indication subgroup (exploratory)")
    add("=" * 100)
    add("")
    add(f"{'bootstrap replications':<34}{replications}")
    add(f"{'fold weighting':<34}{weighting} (recomputed within each subgroup)")
    add(f"{'architectures':<34}{', '.join(architectures)}")
    add(f"{'breasts in the table':<34}{indication_report['breasts']}")
    add(f"{'without an indication':<34}{indication_report['without_indication']}"
        f"{' (dropped)' if dropped_unknown else ' (own subgroup)'}")
    add(f"{'examination ID matching':<34}{indication_report['id_matching']}")
    if load_report.get('dropped_without_patient_id'):
        add(f"{'dropped, no patient ID':<34}{load_report['dropped_without_patient_id']}")
    add("")

    add("-" * 100)
    add("1. Subgroup sizes")
    add("-" * 100)
    add(f"{'subgroup':<22} {'breasts':>8} {'women':>7} {'malignant':>10} {'benign':>8} "
        f"{'folds +/-':>10}")
    for key, entry in results.items():
        counts = entry['counts']
        add(f"{str(entry.get('label', key))[:22]:<22} {counts['breasts']:>8} {counts['women']:>7} "
            f"{counts['malignant']:>10} {counts['benign']:>8} {counts['folds_with_both_classes']:>10}")
    add("")
    add(f"'folds +/-' counts the outer folds holding both classes; only those carry a fold level AUC.")
    add("")

    add("-" * 100)
    add("2. AUC and the paired AP minus FDP difference per subgroup")
    add("-" * 100)
    add("A negative difference means the abbreviated protocol scores worse than the full one.")
    add("")
    for architecture in architectures:
        add(f"{architecture}")
        add(f"{'  subgroup':<22} {'breasts':>8} {'AUC(AP)':>9} {'AUC(FDP)':>9} {'delta':>9} "
            f"{'95% interval':>22} {'folds':>6}")
        for key, entry in results.items():
            label = str(entry.get('label', key))[:20]
            counts = entry['counts']
            values = entry['architectures'].get(architecture)
            if values is None:
                add(f"  {label:<20} {counts['breasts']:>8} {'-':>9} {'-':>9} {'-':>9} "
                    f"{'not analysed':>22} {'-':>6}")
                continue
            interval = f"[{values['ci_95'][0]:+.4f}, {values['ci_95'][1]:+.4f}]"
            add(f"  {label:<20} {counts['breasts']:>8} {values['auc_ap']:>9.4f} "
                f"{values['auc_fdp']:>9.4f} {values['observed_delta']:>+9.4f} {interval:>22} "
                f"{values['folds_used']:>6}")
        add("")

    skipped = {key: entry['skipped'] for key, entry in results.items() if 'skipped' in entry}
    if skipped:
        add("-" * 100)
        add("3. Subgroups reported without a bootstrap")
        add("-" * 100)
        add(f"A subgroup needs at least {MIN_BREASTS} breasts and {MIN_PER_CLASS} of each class; "
            f"below that an interval")
        add("would describe a handful of cases rather than the model.")
        add("")
        for key, reason in skipped.items():
            add(f"  {str(results[key].get('label', key)):<22} {reason}")
        add("")

    add("-" * 100)
    add("How to read this")
    add("-" * 100)
    add("The fold weights are recomputed inside each subgroup, so a subgroup delta is not a")
    add("decomposition of the overall delta and the subgroup deltas do not average to it.")
    add("No multiplicity adjustment is applied across subgroups and no margin is tested here.")
    add("Subgroup differences should be described, not used to claim non-inferiority or")
    add("inferiority within an indication.")
    return "\n".join(lines)


def json_ready(value):
    """
    Recursively turn numpy scalars and non-string dict keys (the fold IDs are numpy ints) into
    plain Python, so json.dump accepts the nested result structure.
    """
    if isinstance(value, dict):
        return {str(key): json_ready(item) for key, item in value.items()}
    if isinstance(value, (list, tuple)):
        return [json_ready(item) for item in value]
    if isinstance(value, np.ndarray):
        return json_ready(value.tolist())
    if isinstance(value, np.integer):
        return int(value)
    if isinstance(value, (np.floating, float)):
        return None if np.isnan(value) else float(value)
    return value


def build_parser():
    parser = argparse.ArgumentParser(prog="analyse_subgroups", description=__doc__,
                                     formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("prediction_dir", help="Folder holding oof_<architecture>_fold<k>.csv")
    parser.add_argument("metadata_file",
                        help="Metadata export holding the per-side indication columns.")
    parser.add_argument("--fraction", type=float, default=None,
                        help="Training fraction to read when the prediction folder holds a whole "
                             "sweep. Default: the largest one present, i.e. the full data run.")
    parser.add_argument("-o", "--output_path", default=None,
                        help="Folder the per subgroup table, report and json are written to.")
    parser.add_argument("-r", "--replications", type=int, default=2000,
                        help="Bootstrap replications per subgroup. Default: 2000")
    parser.add_argument("-w", "--weighting", choices=["pairs", "cases", "equal"], default="pairs",
                        help="How fold level AUCs are combined. Default: pairs")
    parser.add_argument("--seed", type=int, default=0)
    parser.add_argument("--no_stratify", action="store_true",
                        help="Do not stratify the resampling by patient level outcome.")
    parser.add_argument("--drop_unknown", action="store_true",
                        help="Leave breasts without an indication out instead of reporting them "
                             "as their own 'unknown' subgroup.")
    return parser


def main():
    args = build_parser().parse_args()

    output_dir = Path(resolve_path(args.output_path, "output_root", "output folder")) / "subgroups"
    output_dir.mkdir(parents=True, exist_ok=True)

    print("Assembling the out-of-fold table ...")
    table, load_report = load_oof_table(args.prediction_dir, args.metadata_file,
                                       fraction=args.fraction)
    architectures = architectures_of(table)
    print(f"Architectures in the table: {', '.join(architectures)}")

    table, indication_report = attach_indication(table, args.metadata_file)
    print(f"Indications: {indication_report['per_code']}")
    if args.drop_unknown:
        before = len(table)
        table = table[table['indication'] != UNKNOWN].reset_index(drop=True)
        print(f"Dropped {before - len(table)} breasts without an indication")

    checks, _ = check_integrity(table)
    failed = [name for name, (passed, _) in checks.items() if not passed]
    if failed:
        print(f"WARNING: integrity checks failed: {failed}. See analyze_noninferiority.py.")

    print(f"Bootstrapping each subgroup with {args.replications} replications ...")
    results = analyse_by_group(table, 'indication', architectures, replications=args.replications,
                               weighting=args.weighting, seed=args.seed,
                               stratify=not args.no_stratify, label_column='indication_label')

    frame = pd.DataFrame(to_frame(results, architectures))
    frame.to_csv(output_dir / "subgroup_performance.csv", index=False)

    report = format_report(results, architectures, load_report, indication_report,
                           args.replications, args.weighting, dropped_unknown=args.drop_unknown)
    (output_dir / "subgroup_report.txt").write_text(report)

    serialisable = {str(key): {'label': entry.get('label'), 'counts': entry['counts'],
                               'skipped': entry.get('skipped'),
                               'architectures': {
                                   architecture: {name: value for name, value in values.items()
                                                  if name != 'replications'}
                                   for architecture, values in entry['architectures'].items()}}
                    for key, entry in results.items()}
    with open(output_dir / "subgroup_results.json", 'w') as handle:
        json.dump(json_ready(serialisable), handle, indent=2)

    table.to_csv(output_dir / "oof_table_with_indication.csv", index=False)
    print(report)
    print(f"\nWrote the report, the per subgroup CSV and the json to {output_dir}")


if __name__ == '__main__':
    main()