#!/usr/bin/env python3 """ Create nested stratified subsets of per-fold training split files. For every fold directory /fold, the training file stratified_split_*_f-0.csv is subsampled into stratified_split_*_f-0_frac0.05.csv stratified_split_*_f-0_frac0.1.csv stratified_split_*_f-0_frac0.25.csv stratified_split_*_f-0_frac0.5.csv Nesting guarantee: 0.05 subset of 0.1 subset of 0.25 subset of 0.5 subset of full. This is achieved by shuffling each stratum exactly once and taking prefixes of increasing length -- a prefix of a prefix is always a subset, by construction. """ import argparse import sys from pathlib import Path import numpy as np import pandas as pd from auto_detect_breast_mri.config import resolve_path FRACTIONS = [0.05, 0.1, 0.25, 0.5] LABEL_CANDIDATES = ["Malign", "Malign_Left", "Lesion", "Malign_Right"] def fmt_frac(f: float) -> str: """0.05 -> '0.05', 0.1 -> '0.1', 0.25 -> '0.25', 0.5 -> '0.5' (no trailing zeros).""" return f"{f:g}" def pick_label_column(df: pd.DataFrame, explicit: str | None) -> str: if explicit: if explicit not in df.columns: sys.exit(f"ERROR: label column '{explicit}' not in file. Columns: {list(df.columns)}") return explicit for cand in LABEL_CANDIDATES: if cand in df.columns: return cand sys.exit( "ERROR: could not auto-detect a label column.\n" f" Columns present: {list(df.columns)}\n" " Pass one explicitly with --label-col." ) def nested_prefix_lengths(n: int, fractions: list[float], min_per_stratum: int) -> list[int]: """Prefix length for each fraction. Monotonic non-decreasing => nesting holds.""" lengths = [] prev = 0 for f in fractions: k = int(round(n * f)) k = max(k, min_per_stratum) # never empty a class k = min(k, n) # never exceed what exists k = max(k, prev) # enforce monotonicity lengths.append(k) prev = k return lengths def build_subsets(df: pd.DataFrame, label_col: str, group_col: str | None, seed: int, min_per_stratum: int) -> dict[float, pd.DataFrame]: """Return {fraction: dataframe}. Sampling unit is a row, or a group if group_col given.""" rng = np.random.RandomState(seed) if group_col: # One sampling unit per group. Group label = max over its rows, so a group # containing any positive counts as positive. units = df.groupby(group_col)[label_col].max().reset_index() unit_key = group_col else: units = df[[label_col]].copy() units["__row__"] = df.index unit_key = "__row__" selected_per_frac: dict[float, list] = {f: [] for f in FRACTIONS} for lab, block in units.groupby(label_col, sort=True): shuffled = block.sample(frac=1.0, random_state=rng) # ONE shuffle per stratum keys = shuffled[unit_key].to_numpy() for f, k in zip(FRACTIONS, nested_prefix_lengths(len(keys), FRACTIONS, min_per_stratum)): selected_per_frac[f].append(keys[:k]) # prefixes => nested out = {} for f in FRACTIONS: keys = np.concatenate(selected_per_frac[f]) if selected_per_frac[f] else np.array([]) if group_col: sub = df[df[group_col].isin(keys)] else: sub = df.loc[keys] out[f] = sub.sort_index() # keep original row order return out def process_fold(train_path: Path, label_col: str | None, group_col: str | None, seed: int, min_per_stratum: int, dry_run: bool) -> None: df = pd.read_csv(train_path, dtype=str, keep_default_na=False) lab = pick_label_column(df, label_col) # Labels read as str for safety; use a numeric view for grouping/max. df["__lab__"] = pd.to_numeric(df[lab], errors="coerce") if df["__lab__"].isna().any(): n_bad = int(df["__lab__"].isna().sum()) sys.exit(f"ERROR: {train_path.name}: {n_bad} rows have a non-numeric '{lab}' value.") subsets = build_subsets(df, "__lab__", group_col, seed, min_per_stratum) base = df["__lab__"].mean() unit = f"{group_col}s" if group_col else "rows" n_units = df[group_col].nunique() if group_col else len(df) print(f"\n{train_path.parent.name}/{train_path.name}") print(f" full: {len(df):>6} rows / {n_units:>6} {unit}, positive rate {base:.4f}") for f in FRACTIONS: sub = subsets[f].drop(columns="__lab__") out_path = train_path.with_name(f"{train_path.stem}_frac{fmt_frac(f)}.csv") n_sub_units = sub[group_col].nunique() if group_col else len(sub) rate = subsets[f]["__lab__"].mean() print(f" frac {fmt_frac(f):>4}: {len(sub):>6} rows / {n_sub_units:>6} {unit}" f", positive rate {rate:.4f} -> {out_path.name}") if not dry_run: sub.to_csv(out_path, index=False) # Verify nesting on the written selection, not on assumptions. idx = {f: set(subsets[f].index) for f in FRACTIONS} for small, big in zip(FRACTIONS, FRACTIONS[1:]): assert idx[small] <= idx[big], f"NESTING VIOLATED: {small} not subset of {big}" print(" nesting verified: 0.05 < 0.1 < 0.25 < 0.5") def main() -> None: p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) p.add_argument("split_root", nargs="?", default=None, type=Path, help="directory containing fold0 ... fold4") p.add_argument("--pattern", default="stratified_training_*-0.csv", help="glob for the training file inside each fold dir") p.add_argument("--label-col", default=None, help=f"label column; auto-detected from {LABEL_CANDIDATES} if omitted") p.add_argument("--group-col", default=None, help="e.g. 'id' -- sample whole patients instead of individual rows") p.add_argument("--seed", type=int, default=42) p.add_argument("--min-per-stratum", type=int, default=1, help="minimum units kept per class in the smallest fraction") p.add_argument("--dry-run", action="store_true", help="print only, write nothing") args = p.parse_args() split_root = Path(resolve_path(args.split_root, "split_root", "folder for the split files")) fold_dirs = sorted(d for d in split_root.glob("fold*") if d.is_dir()) if not fold_dirs: sys.exit(f"ERROR: no fold* directories under {split_root}") for fold_dir in fold_dirs: matches = [m for m in sorted(fold_dir.glob(args.pattern)) if "_frac" not in m.stem] # don't re-process our own output if len(matches) != 1: print(f"WARNING: {fold_dir.name}: expected 1 training file for " f"'{args.pattern}', found {len(matches)}: {[m.name for m in matches]}", file=sys.stderr) continue process_fold(matches[0], args.label_col, args.group_col, args.seed, args.min_per_stratum, args.dry_run) if args.dry_run: print("\n(dry run -- no files written)") if __name__ == "__main__": main()