Download scripts/validation/analyze_subgroups.py from deboraJ23/AI_MRI: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/analyze_subgroups.py
- Command line
-
hf download hf://deboraJ23/AI_MRI/scripts/validation/analyze_subgroups.py
-
curl -L -o analyze_subgroups.py https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/analyze_subgroups.py
10.3 kB
| """ | |
| 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() | |