Files
hypertower/v3/scripts/main/run_cv.py
T
2026-03-19 16:58:29 +01:00

88 lines
2.4 KiB
Python

#!/usr/bin/env python
"""
V3 cross-validation runner.
Outer/inner k-fold: test = current fold, val = next fold, train = rest.
No holdout. No checkpoint saving.
Usage (single fold-seed, 5-fold, binary, ensemble):
python -m v3.scripts.main.run_cv \
--run-name my_run \
--eval-mode binary \
--tower-mode ensemble \
--epochs 40 \
--augment \
--tune-binary-threshold \
--in-memory-cache
Usage (10x5 rep-CV, seeds 100..1000):
python -m v3.scripts.main.run_cv \
--run-name my_run_10x5 \
--reps 10 \
--rep-seed-start 100 \
--rep-seed-step 100 \
--eval-mode binary \
--tower-mode ensemble \
--epochs 40 \
--augment \
--tune-binary-threshold \
--in-memory-cache
"""
import argparse
import sys
from pathlib import Path
# Allow running as `python v3/scripts/main/run_cv.py` from repo root
sys.path.insert(0, str(Path(__file__).resolve().parents[3]))
from v3.classes.v3_hypertower import V3HyperTower
def build_parser() -> argparse.ArgumentParser:
ap = V3HyperTower.build_parser()
ap.description = __doc__
ap.formatter_class = argparse.RawDescriptionHelpFormatter
ap.add_argument(
"--reps", type=int, default=1,
help="Number of repetitions (each rep uses a different --fold-seed).",
)
ap.add_argument(
"--rep-seed-start", type=int, default=100,
help="fold-seed for rep 0 (default: 100).",
)
ap.add_argument(
"--rep-seed-step", type=int, default=100,
help="Increment between rep fold-seeds (default: 100; rep k uses seed start + k*step).",
)
return ap
def main():
ap = build_parser()
args = ap.parse_args()
reps = int(args.reps)
seed_start = int(args.rep_seed_start)
seed_step = int(args.rep_seed_step)
base_run_name = args.run_name or "v3_cv"
for rep in range(reps):
rep_seed = seed_start + rep * seed_step
args.fold_seed = rep_seed
if reps > 1:
args.run_name = f"{base_run_name}/rep{rep:02d}"
print(f"\n{'='*60}", flush=True)
print(f"Rep {rep+1}/{reps} fold_seed={rep_seed}", flush=True)
print(f"{'='*60}", flush=True)
else:
args.run_name = base_run_name
tower = V3HyperTower(args)
out_dir = tower.run()
print(f"\nRep {rep+1} output: {out_dir}", flush=True)
if __name__ == "__main__":
main()