File size: 7,199 Bytes
a23d562 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 | #!/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() |