""" Exploratory performance per indication subgroup on the out-of-fold predictions. python scripts/validation/analyze_subgroups.py 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__fold.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()