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