#!/usr/bin/env python3 """Stress-test SDIL with limited and noisy neutral observations.""" from __future__ import annotations import argparse from concurrent.futures import ProcessPoolExecutor, as_completed import json from pathlib import Path import sys import numpy as np ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) sys.path.insert(0, str(Path(__file__).resolve().parent)) from physical_grid_bias_p5 import run_trial, select_tasks # noqa: E402 def conditions() -> list[dict]: settings = [ {"name": "degree1_n16_clean", "degree": 1, "observations": 16, "noise": 0.0}, {"name": "degree2_n3_clean", "degree": 2, "observations": 3, "noise": 0.0}, {"name": "degree2_n4_clean", "degree": 2, "observations": 4, "noise": 0.0}, {"name": "degree2_n8_clean", "degree": 2, "observations": 8, "noise": 0.0}, {"name": "degree2_n16_clean", "degree": 2, "observations": 16, "noise": 0.0}, ] for noise in (0.05, 0.1, 0.25, 0.5, 1.0): settings.append({ "name": f"degree2_n16_noise{noise:g}", "degree": 2, "observations": 16, "noise": noise, }) return settings def run_job(job: dict) -> dict: result = run_trial(job) method = result["methods"]["sdil"] return { "condition": job["condition"], "task_index": result["task_index"], "input_diameter_v": result["input_diameter_v"], "device_seed": result["device_seed"], "classification_error": method["classification_error"], "hinge_loss_v2": method["hinge_loss_v2"], "max_abs_clamp_displacement_v": ( method["max_abs_clamp_displacement_v"]), "heldout_bias_rmse_v_per_s": ( result["calibration"]["sdil_heldout_rmse_v_per_s"]), } def summarize(records: list[dict], settings: list[dict]) -> dict: result = {} for setting in settings: selected = [ record for record in records if record["condition"] == setting["name"] ] errors = np.asarray([ record["classification_error"] for record in selected ]) result[setting["name"]] = { **setting, "trials": len(selected), "mean_classification_error": float(np.mean(errors)), "zero_error_fraction": float(np.mean(errors == 0.0)), "median_heldout_bias_rmse_v_per_s": float(np.median([ record["heldout_bias_rmse_v_per_s"] for record in selected ])), } return result def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument( "--protocol", type=Path, default=Path("results/physical_bias/dillavou_fig5_protocol.json")) parser.add_argument( "--output", type=Path, default=Path("results/physical_bias/p6_grid_calibration_robustness.json")) parser.add_argument("--rotations", type=int, default=8) parser.add_argument( "--device-seeds", default="20260829,20260830,20260831,20260832") parser.add_argument("--standard-epochs", type=int, default=600) parser.add_argument("--workers", type=int, default=16) return parser.parse_args() def main() -> None: args = parse_args() protocol = json.loads(args.protocol.read_text()) tasks = select_tasks(protocol, args.rotations) device_seeds = tuple(int(seed) for seed in args.device_seeds.split(",")) settings = conditions() jobs = [] for setting in settings: for task_index, task in enumerate(tasks): for device_seed in device_seeds: jobs.append({ "condition": setting["name"], "task_index": task_index, "task": task, "device_seed": device_seed, "calibration_observations": setting["observations"], "heldout_observations": 64, "calibration_noise_standard_deviation_v_per_s": setting["noise"], "sdil_degree": setting["degree"], "standard_epochs": args.standard_epochs, "overclamp_epochs": 1, "methods": ("sdil",), }) records = [] with ProcessPoolExecutor(max_workers=args.workers) as executor: futures = [executor.submit(run_job, job) for job in jobs] for completed, future in enumerate(as_completed(futures), start=1): records.append(future.result()) if completed % 80 == 0 or completed == len(jobs): print(f"completed {completed}/{len(jobs)}", flush=True) records.sort(key=lambda record: ( record["condition"], record["task_index"], record["device_seed"])) report = { "analysis": "physical_grid_sdil_calibration_robustness_p6", "confirmatory": False, "autodiff_used": False, "source_protocol": str(args.protocol), "protocol": { "task_count": len(tasks), "device_seeds": device_seeds, "training_method": "sdil", "standard_epochs": args.standard_epochs, "heldout_neutral_observations": 64, "calibration_noise_units": "V/s added independently per edge and observation", "conditions": settings, }, "records": records, "summary": summarize(records, settings), } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(report, indent=2) + "\n") print(json.dumps(report["summary"], indent=2)) print(f"wrote {args.output}") if __name__ == "__main__": main()