import argparse
import json
import math
import os
import statistics
import subprocess
import tempfile
import time
from collections import defaultdict
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import r5_common as _r5
import r6_common as _r6
from r5_common import BIN, EXP, FIRES, HOLDOUT, VAL
__all__ = [
"BIN", "EXP", "FIRES", "HOLDOUT", "VAL", "REPO", "PRIOR", "ARM_B", "NICE_PREFIX",
"arg_parser", "command_for", "dry_run", "run", "run_all",
"check_binary_git", "load_1min", "wait_for_load",
"fire_stats", "verdict", "summary_table", "e33_baseline", "arm_b_baseline",
]
REPO = Path(__file__).resolve().parents[3]
PRIOR = VAL / "scripts" / "experiments" / "priors" / "arrival_x4.json"
NICE_PREFIX = ["nice", "-n", "10"]
ARM_B = {
"SMC_SPREAD": "arrival",
"SMC_WIND_LAW": "rear_focus",
"SMC_STEPS_SCALE": "4",
"SMC_PRIOR": str(PRIOR),
"SMC_WIND_ROT_GENE": "90",
}
_E33_FALLBACK = {
"Bear_2020": (0.479, 0.015),
"Brattain_2020": (0.416, 0.004),
"Buck_2017": (0.590, 0.039),
"Chimney_2016": (0.434, 0.012),
"Ferguson_2018": (0.344, 0.007),
"Pier_2017": (0.535, 0.003),
}
def arg_parser(description):
p = argparse.ArgumentParser(description=description)
p.add_argument("--dry-run", action="store_true",
help="print every job's command and env; launch nothing")
p.add_argument("--workers", type=int, default=2,
help="max parallel wildfire_smc processes (default 2, shared machine)")
return p
def command_for(out_dir, fire, label, env, members, mode):
rep = out_dir / f"{fire}_{label}.json"
return [*NICE_PREFIX, str(BIN), str(VAL / "data" / "scenarios" / fire), str(members), mode, str(rep)]
def dry_run(jobs, out_json):
out_dir = EXP / out_json.replace(".json", "")
print(f"[dry-run] {len(jobs)} job(s) -> {out_json} (raw reports under {out_dir})")
for fire, label, env, members, mode in jobs:
full_env = {**_r5.BASE_ENV, **env}
env_str = " ".join(f"{k}={v}" for k, v in full_env.items())
argv = command_for(out_dir, fire, label, env, members, mode)
print(f" {fire:14s} {label:24s} mode={mode:6s} members={members:3d}")
print(f" env: {env_str}")
print(f" cmd: {' '.join(argv)}")
def _git(*args):
return subprocess.run(["git", *args], cwd=REPO, capture_output=True, text=True, check=True).stdout
def check_binary_git():
head = _git("rev-parse", "--short", "HEAD").strip()
dirty_tracked = [
line for line in _git("status", "--porcelain").splitlines()
if not line.startswith("??")
]
if dirty_tracked:
raise RuntimeError(
"refusing to run: tracked working-tree changes present:\n" + "\n".join(dirty_tracked)
)
with tempfile.TemporaryDirectory(prefix="r7_binary_check_") as tmp:
report = _r6.run_nulls(Path(tmp), "Bear_2020", argv_prefix=NICE_PREFIX)
stamp = report.get("binary_git", "unknown")
stamp_sha = stamp[: -len("-dirty")] if stamp.endswith("-dirty") else stamp
if stamp == "unknown" or stamp.endswith("-dirty") or stamp_sha != head:
raise RuntimeError(
f"refusing to run: wildfire_smc's binary_git={stamp!r} does not match a clean "
f"HEAD={head}; rebuild wildfire_smc (cargo build --release --example wildfire_smc)"
)
return stamp, head
def load_1min():
out = subprocess.run(["uptime"], capture_output=True, text=True, check=True).stdout
la = out.split("load average:")[-1]
return float(la.split(",")[0].strip())
def wait_for_load(threshold=8.0, sleep_s=600):
load = load_1min()
while load > threshold:
print(f"[r7] load {load:.2f} > {threshold}, waiting {sleep_s}s before re-checking...", flush=True)
time.sleep(sleep_s)
load = load_1min()
return load
def run(out_dir, fire, label, env, members=32, mode="assim"):
if mode == "map":
return _run_map(out_dir, fire, label, env, members)
return _r5.run(out_dir, fire, label, env, members, mode, argv_prefix=NICE_PREFIX)
def _run_map(out_dir, fire, label, env, members=32):
out_dir.mkdir(parents=True, exist_ok=True)
rep = out_dir / f"{fire}_{label}.json"
full_env = {**os.environ, **_r5.BASE_ENV, **env}
argv = [*NICE_PREFIX, str(BIN), str(VAL / "data" / "scenarios" / fire), str(members), "map", str(rep)]
subprocess.run(argv, check=True, capture_output=True, env=full_env)
r = json.loads(rep.read_text())
a = r["archive"]
growths = [e["descriptor"][0] for e in a["elites"]]
elongs = [e["descriptor"][1] for e in a["elites"]]
observed = r.get("observed", [])
obs_growth_day5 = observed[-1][1] if observed else float("nan")
obs_elong_day5 = observed[-1][2] if observed else float("nan")
at_size = [el for g, el in zip(growths, elongs) if g >= obs_growth_day5]
row = {
"fire": fire, "config": label, "mode": "map", "members": members,
"holdout": fire in HOLDOUT,
"binary_git": r.get("binary_git", "unknown"), "binary_built_utc": r.get("binary_built_utc", "unknown"),
"elites": a["stats"]["elites"], "coverage": a["stats"]["coverage"],
"labels": a["labels"], "ranges": a["ranges"],
"max_elongation_any_size": max(elongs) if elongs else None,
"max_elongation_at_size": max(at_size) if at_size else None,
"observed_growth_day5": obs_growth_day5, "observed_elongation_day5": obs_elong_day5,
"observed": observed,
}
at_size_str = "None" if row["max_elongation_at_size"] is None else f"{row['max_elongation_at_size']:.2f}"
any_size_str = "None" if row["max_elongation_any_size"] is None else f"{row['max_elongation_any_size']:.2f}"
print(f"{fire:14s} {label:24s} map elites {row['elites']:3d} coverage {row['coverage']:.2f} "
f"max elong any {any_size_str} at-size {at_size_str} "
f"(observed g={obs_growth_day5:.3f} e={obs_elong_day5:.2f})", flush=True)
return row
def run_all(jobs, out_json, workers=2, load_gate=8.0, gate_sleep=600, skip_binary_check=False):
out_dir = EXP / out_json.replace(".json", "")
load_start = wait_for_load(load_gate, gate_sleep)
if not skip_binary_check:
check_binary_git()
print(f"[r7] batch start: {len(jobs)} job(s), workers={workers}, load(1m)={load_start:.2f}", flush=True)
t0 = time.time()
with ThreadPoolExecutor(max_workers=workers) as ex:
rows = list(ex.map(lambda j: run(out_dir, j[0], j[1], j[2], j[3], j[4]), jobs))
wall_s = time.time() - t0
load_end = load_1min()
(EXP / out_json).write_text(json.dumps(rows, indent=1))
summary = {
"out_json": out_json, "n_jobs": len(jobs), "workers": workers,
"wall_s": round(wall_s, 1),
"load_1min_start": load_start, "load_1min_end": load_end,
"binary_git": rows[0]["binary_git"] if rows else "unknown",
}
(EXP / out_json.replace(".json", "_summary.json")).write_text(json.dumps(summary, indent=1))
print(f"[r7] batch done: {len(jobs)} job(s), {wall_s:.1f}s, load(1m) {load_start:.2f} -> {load_end:.2f}",
flush=True)
return rows, summary
def fire_stats(rows, value_key="mean_consensus_iou"):
by_fire = defaultdict(list)
for row in rows:
by_fire[row["fire"]].append(row[value_key])
out = {}
for fire, vals in by_fire.items():
mean = statistics.mean(vals)
sd = statistics.stdev(vals) if len(vals) > 1 else 0.0
out[fire] = (mean, sd, len(vals))
return out
def e33_baseline(value_key="mean_consensus_iou"):
path = EXP / "exp33_noise.json"
if path.exists():
rows = json.loads(path.read_text())
stats = fire_stats(rows, value_key)
return {fire: (mean, sd) for fire, (mean, sd, _n) in stats.items()}
return dict(_E33_FALLBACK)
_ARM_B_FALLBACK = {
"Bear_2020": (0.47285919911915464, 0.0053476183233077705),
"Brattain_2020": (0.4374735867318865, 0.03392424393899961),
"Buck_2017": (0.6397036656372039, 0.00538794378529568),
"Chimney_2016": (0.489335016243905, 0.03077702875196609),
"Ferguson_2018": (0.3859720996337302, 0.013614755532273325),
"Pier_2017": (0.5207486354935142, 0.022770117809419003),
}
def arm_b_baseline():
path = EXP / "exp44_arm_b_5seed_summary.json"
if path.exists():
summary = json.loads(path.read_text())
block = summary.get("arm_b_sd")
if block:
return {fire: (v["mean"], v["sd"]) for fire, v in block.items()}
return dict(_ARM_B_FALLBACK)
def verdict(delta_sd):
a = abs(delta_sd)
if a < 1.0:
return "tie"
direction = "gain" if delta_sd > 0 else "loss"
label = f"beyond {2 if a >= 2.0 else 1} sd ({direction})"
return f"**{label}**" if a >= 2.0 else label
def summary_table(arm_stats_by_label, baseline, fires=FIRES, show_se=False, baseline_n=5):
labels = list(arm_stats_by_label)
header = "| Fire | Baseline mean | Baseline sd |" + "".join(
f" {l} mean | {l} sd | Delta {l} (sd) | verdict {l} |"
+ (f" {l} SE of difference | Delta {l} (SE) |" if show_se else "")
for l in labels
)
sep = "|---|---|---|" + "".join(
"---|---|---|---|" + ("---|---|" if show_se else "") for _ in labels
)
lines = [header, sep]
for fire in fires:
b_mean, b_sd = baseline[fire]
name = fire.split("_")[0] + ("*" if fire in HOLDOUT else "")
cells = [name, f"{b_mean:.3f}", f"{b_sd:.3f}"]
for label in labels:
mean, sd, n = arm_stats_by_label[label].get(fire, (float("nan"), float("nan"), 0))
if n == 0:
cells += ["missing", "missing", "missing", "missing"]
if show_se:
cells += ["missing", "missing"]
continue
delta = mean - b_mean
if b_sd:
delta_str = f"{delta:+.3f} ({delta / b_sd:+.2f} sd)"
v = verdict(delta / b_sd)
else:
delta_str = f"{delta:+.3f} (sd=0)"
v = "undefined (zero sd)"
cells += [f"{mean:.3f}", f"{sd:.3f}", delta_str, v]
if show_se:
se = math.sqrt(sd ** 2 / n + b_sd ** 2 / baseline_n)
if se:
cells += [f"{se:.4f}", f"{delta / se:+.2f}"]
else:
cells += ["0.0000", "undefined (zero SE)"]
lines.append("| " + " | ".join(cells) + " |")
return "\n".join(lines)