Files
hypertower/scripts/main/v1/run_multifold_grid.py
T
2026-02-24 10:39:48 +01:00

432 lines
14 KiB
Python
Executable File

#!/usr/bin/env python3
"""
Grid-search runner for run_multifold experiments.
Features:
* Enumerates the requested configuration grid and writes grid_plan.csv.
* Picks the next incomplete run, marks it running, executes run_multifold.py.
* Records AUC/accuracy metrics per fold into grid_report.csv.
* Removes model checkpoints for runs dominated (80%+ metrics worse) by others.
"""
from __future__ import annotations
import argparse
import csv
import json
import math
import os
import shutil
import subprocess
import sys
from contextlib import contextmanager
from datetime import datetime
from pathlib import Path
from typing import Dict, List, Optional
import fcntl
REPO_ROOT = Path(__file__).resolve().parents[2]
RUN_SCRIPT = REPO_ROOT / "scripts" / "run_multifold.py"
MANIFEST = REPO_ROOT / "manifest.csv"
GRID_DIR = REPO_ROOT / "analysis_data" / "grid_search"
PLAN_PATH = GRID_DIR / "grid_plan.csv"
REPORT_PATH = GRID_DIR / "grid_report.csv"
LOCK_PATH = GRID_DIR / ".grid_lock"
MODELS_ROOT = REPO_ROOT / "models" / "grid_search"
def parse_args() -> argparse.Namespace:
ap = argparse.ArgumentParser(description="Grid-search orchestrator for run_multifold.")
ap.add_argument("--plan-date", default=datetime.now().strftime("%Y%m%d"),
help="Date prefix used when generating run IDs (default: today).")
ap.add_argument("--regen-plan", action="store_true",
help="Rebuild the grid plan from scratch (overwrites existing plan).")
ap.add_argument("--manifest", type=Path, default=MANIFEST,
help="UNet manifest CSV for cropper.")
ap.add_argument("--weights-dir", type=Path, default=REPO_ROOT / "models" / "unet_segmenter",
help="Directory containing norm_* subfolders with best.pt.")
ap.add_argument("--dry-run", action="store_true", help="Enumerate next run without executing.")
ap.add_argument("--max-runs", type=int, default=1,
help="Maximum runs to execute in this invocation (default: 1).")
ap.add_argument("--run-all", action="store_true",
help="Execute runs sequentially until plan is exhausted (overrides --max-runs).")
return ap.parse_args()
@contextmanager
def file_lock(lock_path: Path):
lock_path.parent.mkdir(parents=True, exist_ok=True)
with open(lock_path, "w") as lock_file:
fcntl.flock(lock_file, fcntl.LOCK_EX)
try:
yield
finally:
fcntl.flock(lock_file, fcntl.LOCK_UN)
def read_csv(path: Path) -> List[Dict[str, str]]:
if not path.exists():
return []
with path.open(newline="") as fh:
reader = csv.DictReader(fh)
return list(reader)
def write_csv(path: Path, rows: List[Dict[str, str]], headers: List[str]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", newline="") as fh:
writer = csv.DictWriter(fh, fieldnames=headers)
writer.writeheader()
for row in rows:
writer.writerow(row)
def grid_configs(base_date: str, weights_dir: Path) -> List[Dict[str, str]]:
eval_modes = ["binary", "multiclass"]
crop_variants = [
("norm_imagenet", "imagenet"),
("normalize_none", "none"),
("norm_per_image", "per_image"),
]
tta_opts = [False, True]
loss_modes = ["focal", "balanced", "none"]
thaw_modes = ["none", "gradual"]
se_configs = []
# none
se_configs.append(("none", {"se_enabled": False}))
# bridge only
for pre in (True, False):
se_configs.append((
"bridge",
{"se_enabled": True, "se_where": "bridge", "bridge_pre_norm": pre, "tower_pre_norm": None},
))
# tower only
for pre in (True, False):
se_configs.append((
"tower",
{"se_enabled": True, "se_where": "tower", "bridge_pre_norm": None, "tower_pre_norm": pre},
))
# both (four combos)
for b_pre in (True, False):
for t_pre in (True, False):
se_configs.append((
"both",
{
"se_enabled": True,
"se_where": "both",
"bridge_pre_norm": b_pre,
"tower_pre_norm": t_pre,
},
))
combos = []
idx = 0
for eval_mode in eval_modes:
for variant, norm in crop_variants:
weights_path = weights_dir / variant / "best.pt"
for tta in tta_opts:
for loss in loss_modes:
for thaw in thaw_modes:
for se_name, se_opts in se_configs:
run_id = f"{base_date}-{idx:04d}"
combos.append({
"run_id": run_id,
"status": "incomplete",
"eval_mode": eval_mode,
"crop_variant": variant,
"crop_normalize": norm,
"crop_weights": str(weights_path),
"crop_tta": str(tta),
"loss_mode": loss,
"thaw_mode": thaw,
"se_mode": se_name,
"se_bridge_pre_norm": str(se_opts.get("bridge_pre_norm")),
"se_tower_pre_norm": str(se_opts.get("tower_pre_norm")),
})
idx += 1
return combos
PLAN_HEADERS = [
"run_id",
"status",
"eval_mode",
"crop_variant",
"crop_normalize",
"crop_weights",
"crop_tta",
"loss_mode",
"thaw_mode",
"se_mode",
"se_bridge_pre_norm",
"se_tower_pre_norm",
]
def ensure_plan(args: argparse.Namespace) -> None:
if args.regen_plan or not PLAN_PATH.exists():
combos = grid_configs(args.plan_date, args.weights_dir)
write_csv(PLAN_PATH, combos, PLAN_HEADERS)
print(f"[grid] Plan created with {len(combos)} runs at {PLAN_PATH}")
def select_next_run() -> Optional[Dict[str, str]]:
rows = read_csv(PLAN_PATH)
for row in rows:
if row["status"] == "incomplete":
row["status"] = "running"
write_csv(PLAN_PATH, rows, PLAN_HEADERS)
return row
return None
def update_run_status(run_id: str, new_status: str) -> None:
rows = read_csv(PLAN_PATH)
for row in rows:
if row["run_id"] == run_id:
row["status"] = new_status
break
write_csv(PLAN_PATH, rows, PLAN_HEADERS)
def build_run_command(row: Dict[str, str], manifest: Path) -> List[str]:
cmd = [
sys.executable,
str(RUN_SCRIPT),
"--backbone",
"resnet50",
"--fusion-mode",
"fused",
"--epochs",
"40",
"--batch-size",
"8",
"--img-crop-manifest",
str(manifest),
"--img-crop-weights",
row["crop_weights"],
"--img-crop-normalize",
row["crop_normalize"],
"--eval_mode",
row["eval_mode"],
"--holdout-per-class",
"12",
"--run-id",
row["run_id"],
"--shortname",
"grid_search",
]
if row["crop_tta"] == "True":
cmd.append("--img-crop-tta")
# Loss/balancing modes
if row["loss_mode"] == "focal":
cmd.extend(["--focal-gamma", "2.0"])
elif row["loss_mode"] == "balanced":
cmd.append("--balanced-sampler")
# Thaw schedule
if row["thaw_mode"] == "gradual":
cmd.append("--gradual-thaw")
cmd.extend(["--thaw-ratio", "0.33"])
cmd.extend(["--thaw-start-epoch", "10"])
cmd.extend(["--thaw-target", "image"])
# SE settings
if row["se_mode"] == "none":
cmd.append("--no-se")
else:
cmd.extend(["--se-reduction", "16"])
cmd.extend(["--se-reduction-tower", "16"])
cmd.extend(["--se-where", row["se_mode"]])
bridge_pre = row["se_bridge_pre_norm"]
tower_pre = row["se_tower_pre_norm"]
if bridge_pre == "True":
cmd.append("--se-pre-norm")
elif bridge_pre == "False":
cmd.append("--no-se-pre-norm")
if tower_pre == "True":
cmd.append("--se-pre-norm-tower")
elif tower_pre == "False":
cmd.append("--no-se-pre-norm-tower")
return cmd
def run_command(cmd: List[str]) -> None:
print("[grid] Launching:", " ".join(cmd))
subprocess.run(cmd, check=True)
METRIC_KEYS = ["auc_fused", "auc_img", "auc_md", "acc_fused", "acc_img", "acc_md"]
def extract_metrics(run_id: str) -> Dict[str, str]:
summary_path = REPO_ROOT / "analysis_data" / "grid_search" / run_id / "summary.json"
if not summary_path.exists():
raise FileNotFoundError(f"Missing summary.json for run {run_id}")
with summary_path.open() as fh:
summary = json.load(fh)
rows = {}
for fold in summary.get("fold_metrics", []):
if not isinstance(fold, dict):
continue
f_idx = fold.get("fold")
stats = fold.get("stats") or {}
if not isinstance(stats, dict):
continue
for key in METRIC_KEYS:
val = stats.get(key)
if val is None:
continue
rows[f"metric_fold{f_idx}_{key}"] = str(val)
best_mean = summary.get("best_metric_mean")
if best_mean is not None:
rows["metric_best_mean"] = str(best_mean)
return rows
def update_report(row: Dict[str, str], metrics: Dict[str, str]) -> None:
existing = read_csv(REPORT_PATH)
# Remove existing entry for run_id
existing = [r for r in existing if r.get("run_id") != row["run_id"]]
record = {**row, **metrics}
existing.append(record)
headers = sorted({key for r in existing for key in r.keys()})
write_csv(REPORT_PATH, existing, headers)
def load_report_rows() -> List[Dict[str, str]]:
return read_csv(REPORT_PATH)
def metric_columns(rows: List[Dict[str, str]]) -> List[str]:
keys = set()
for row in rows:
for key in row:
if key.startswith("metric_"):
keys.add(key)
return sorted(keys)
def _to_float(val: str) -> Optional[float]:
try:
f = float(val)
if math.isnan(f):
return None
return f
except Exception:
return None
def prune_dominated(rows: List[Dict[str, str]]) -> None:
"""
Remove model directories for runs that are clearly dominated by another run.
A run is dominated if:
* Another run has a strictly higher metric_best_mean, OR
* Another run is >= on >=80% of overlapping metrics and strictly better on at least one.
"""
metrics = metric_columns(rows)
if not metrics:
return
dominated = set()
for row in rows:
run_id = row["run_id"]
row_vals = {m: row.get(m) for m in metrics}
row_best = _to_float(row_vals.get("metric_best_mean"))
for other in rows:
if other["run_id"] == run_id:
continue
other_vals = {m: other.get(m) for m in metrics}
other_best = _to_float(other_vals.get("metric_best_mean"))
# Fast path: compare aggregate best mean if both have it
if row_best is not None and other_best is not None and other_best > row_best:
dominated.add(run_id)
break
# Fallback: overlap-wise dominance
comparisons = []
better = 0
for key in metrics:
v1 = _to_float(row_vals.get(key))
v2 = _to_float(other_vals.get(key))
if v1 is None or v2 is None:
continue
comparisons.append(v2 >= v1)
if v2 > v1:
better += 1
if not comparisons:
continue
fraction = sum(comparisons) / len(comparisons)
if fraction >= 0.8 and better > 0:
dominated.add(run_id)
break
for run_id in dominated:
model_dir = MODELS_ROOT / run_id
if model_dir.exists():
print(f"[grid] Removing dominated model artifacts for {run_id}")
try:
shutil.rmtree(model_dir)
except OSError as exc:
# Don't fail the grid run if cleanup isn't permitted (e.g., locked SMB dirs).
print(f"[grid] Warning: could not remove {model_dir}: {exc}")
def main():
args = parse_args()
ensure_plan(args)
if args.dry_run:
with file_lock(LOCK_PATH):
next_run = select_next_run()
if next_run is None:
print("[grid] No incomplete runs remaining.")
return
update_run_status(next_run["run_id"], "incomplete")
print("[grid] Next run:", next_run)
return
max_runs = None if args.run_all else args.max_runs
runs_done = 0
while True:
with file_lock(LOCK_PATH):
next_run = select_next_run()
if next_run is None:
if runs_done == 0:
print("[grid] All runs completed.")
else:
print(f"[grid] No more runs remaining after {runs_done} run(s).")
return
run_id = next_run["run_id"]
try:
cmd = build_run_command(next_run, args.manifest)
run_command(cmd)
metrics = extract_metrics(run_id)
with file_lock(LOCK_PATH):
update_run_status(run_id, "completed")
update_report(next_run, metrics)
report_rows = load_report_rows()
prune_dominated(report_rows)
print(f"[grid] Run {run_id} completed.")
except Exception as exc:
with file_lock(LOCK_PATH):
update_run_status(run_id, "incomplete")
raise SystemExit(f"[grid] Run {run_id} failed: {exc}") from exc
runs_done += 1
if max_runs is not None and runs_done >= max_runs:
print(f"[grid] Reached run limit ({max_runs}); stopping.")
return
if __name__ == "__main__":
main()