from __future__ import annotations import argparse import json import logging import random from collections import defaultdict from pathlib import Path logger = logging.getLogger(__name__) def classify_transformation(pair: dict) -> str: age_changed = pair.get("source_age") != pair.get("target_age") sex_changed = pair.get("source_sex") != pair.get("target_sex") if age_changed and sex_changed: return "intersectional" if age_changed: return "age_only" if sex_changed: return "sex_only" return "none" def _get_pathology_key(pair: dict) -> str: labels = pair.get("pathology_labels", []) if not labels: return "unlabeled" return ",".join(str(v) for v in sorted(labels)) def stratified_sample( pairs: list[dict], n_per_stratum: int, seed: int, ) -> list[int]: rng = random.Random(seed) by_stratum: dict[str, list[int]] = defaultdict(list) for i, pair in enumerate(pairs): stratum = classify_transformation(pair) if stratum != "none": by_stratum[stratum].append(i) selected: list[int] = [] for stratum in ["age_only", "sex_only", "intersectional"]: pool = by_stratum[stratum] if not pool: logger.warning("No pairs available for stratum %s", stratum) continue by_pathology: dict[str, list[int]] = defaultdict(list) for idx in pool: key = _get_pathology_key(pairs[idx]) by_pathology[key].append(idx) pathology_keys = sorted(by_pathology.keys()) n_categories = len(pathology_keys) per_category = n_per_stratum // n_categories remainder = n_per_stratum % n_categories stratum_selected: list[int] = [] overflow_pool: list[int] = [] for k_idx, key in enumerate(pathology_keys): candidates = by_pathology[key] rng.shuffle(candidates) target = per_category + (1 if k_idx < remainder else 0) taken = min(target, len(candidates)) stratum_selected.extend(candidates[:taken]) if taken < target: overflow_pool.extend([]) if taken < len(candidates): overflow_pool.extend(candidates[taken:]) if len(stratum_selected) < n_per_stratum and overflow_pool: rng.shuffle(overflow_pool) needed = n_per_stratum - len(stratum_selected) stratum_selected.extend(overflow_pool[:needed]) if len(stratum_selected) < n_per_stratum: logger.warning( "Only %d pairs available for %s (requested %d)", len(stratum_selected), stratum, n_per_stratum, ) selected.extend(stratum_selected) return selected def main() -> None: logging.basicConfig(level=logging.INFO) parser = argparse.ArgumentParser( description="Stratified sampling of pairs for radiologist validation" ) parser.add_argument("--pairs", type=str, required=True) parser.add_argument("--n-per-stratum", type=int, default=150) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--output", type=str, default="data/pairs.json") args = parser.parse_args() with open(args.pairs) as f: all_pairs = json.load(f) indices = stratified_sample(all_pairs, args.n_per_stratum, args.seed) sampled = [all_pairs[i] for i in indices] for pair in sampled: pair["_transform_type"] = classify_transformation(pair) counts: dict[str, int] = defaultdict(int) pathology_counts: dict[str, dict[str, int]] = defaultdict(lambda: defaultdict(int)) for pair in sampled: t = pair["_transform_type"] counts[t] += 1 pathology_counts[t][_get_pathology_key(pair)] += 1 output_path = Path(args.output) output_path.parent.mkdir(parents=True, exist_ok=True) with open(output_path, "w") as f: json.dump(sampled, f, indent=2) logger.info("Sampled %d pairs from %d total", len(sampled), len(all_pairs)) for stratum, count in sorted(counts.items()): logger.info(" %s: %d", stratum, count) for pkey, pcount in sorted(pathology_counts[stratum].items()): logger.info(" pathology [%s]: %d", pkey, pcount) logger.info("Saved to: %s", output_path) if __name__ == "__main__": main()