AI_MRI / scripts /data_utils /make_split_fractions.py
DeboraJ1's picture
init
a23d562
Raw History Blame Contribute Delete
7.2 kB
#!/usr/bin/env python3
"""
Create nested stratified subsets of per-fold training split files.
For every fold directory <split_root>/fold<k>, the training file
stratified_split_*_f<k>-0.csv
is subsampled into
stratified_split_*_f<k>-0_frac0.05.csv
stratified_split_*_f<k>-0_frac0.1.csv
stratified_split_*_f<k>-0_frac0.25.csv
stratified_split_*_f<k>-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()