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()