diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 16:42:03 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 16:42:03 -0500 |
| commit | 611ff73e78a0e10345393eaaf3b5a3cf075a6d8b (patch) | |
| tree | 3baf8d447923d9834d40f6f6e9fa187e9ef5c012 /experiments/physical_grid_correlated_autozero_p9.py | |
| parent | bd2167df3b069703f2e12368dbaaa8694add2786 (diff) | |
exp: add correlated autozero robustness matrix
Diffstat (limited to 'experiments/physical_grid_correlated_autozero_p9.py')
| -rw-r--r-- | experiments/physical_grid_correlated_autozero_p9.py | 381 |
1 files changed, 381 insertions, 0 deletions
diff --git a/experiments/physical_grid_correlated_autozero_p9.py b/experiments/physical_grid_correlated_autozero_p9.py new file mode 100644 index 0000000..4a33779 --- /dev/null +++ b/experiments/physical_grid_correlated_autozero_p9.py @@ -0,0 +1,381 @@ +#!/usr/bin/env python3 +"""Evaluate correlated auto-zero sampling on the physical CLLN grid.""" + +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 select_tasks # noqa: E402 +from sdil.physical_grid import ( # noqa: E402 + CorrelatedDoubleSampleHold, + GridCircuit, + GridSquareLawImperfection, + RingClassificationDataset, + train_grid_classifier, +) + + +def setting( + name: str, + *, + common_pedestal: float = 0.0, + pedestal_mismatch: float = 0.0, + gain_mismatch: float = 0.0, + noise: float = 0.0, + refresh: int = 1, + method: str = "cds_autozero_sdil", +) -> dict: + return { + "name": name, + "common_pedestal_standard_deviation_v_per_s": common_pedestal, + "pedestal_mismatch_standard_deviation_v_per_s": pedestal_mismatch, + "sample_gain_mismatch_standard_deviation": gain_mismatch, + "sample_noise_standard_deviation_v_per_s": noise, + "refresh_interval_updates": refresh, + "method": method, + } + + +def conditions() -> list[dict]: + settings = [setting("ideal_cds")] + for common_pedestal in (2.3, 10.0): + settings.append(setting( + f"common_pedestal_{common_pedestal:g}", + common_pedestal=common_pedestal, + )) + for mismatch in (0.01, 0.025, 0.05, 0.1, 0.25, 0.5): + settings.append(setting( + f"pedestal_mismatch_{mismatch:g}", + common_pedestal=2.3, + pedestal_mismatch=mismatch, + )) + for mismatch in (0.001, 0.005, 0.01, 0.025, 0.05): + settings.append(setting( + f"gain_mismatch_{mismatch:g}", + common_pedestal=2.3, + gain_mismatch=mismatch, + )) + for noise in (0.1, 0.25, 0.5, 1.0): + settings.append(setting( + f"sample_noise_{noise:g}", + common_pedestal=2.3, + noise=noise, + )) + for refresh in (2, 4, 8): + settings.append(setting( + f"refresh_every_{refresh}", + common_pedestal=2.3, + refresh=refresh, + )) + settings.extend(( + setting( + "combined_mild", + common_pedestal=2.3, + pedestal_mismatch=0.025, + gain_mismatch=0.005, + noise=0.1, + ), + setting( + "combined_mild_refresh4", + common_pedestal=2.3, + pedestal_mismatch=0.025, + gain_mismatch=0.005, + noise=0.1, + refresh=4, + ), + setting( + "combined_strong", + common_pedestal=2.3, + pedestal_mismatch=0.1, + gain_mismatch=0.01, + noise=0.25, + ), + setting( + "overclamp_plus_combined_mild", + common_pedestal=2.3, + pedestal_mismatch=0.025, + gain_mismatch=0.005, + noise=0.1, + method="overclamp_cds_autozero_sdil", + ), + )) + return settings + + +def fixed_edge_errors( + circuit: GridCircuit, task_index: int, device_seed: int, condition: dict +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + common_rng = np.random.default_rng( + device_seed + 6_000_003 + task_index) + pedestal_rng = np.random.default_rng( + device_seed + 7_000_003 + task_index) + gain_rng = np.random.default_rng( + device_seed + 8_000_003 + task_index) + return ( + common_rng.normal( + 0.0, + condition["common_pedestal_standard_deviation_v_per_s"], + circuit.edge_count, + ), + pedestal_rng.normal( + 0.0, + condition["pedestal_mismatch_standard_deviation_v_per_s"], + circuit.edge_count, + ), + gain_rng.normal( + 0.0, + condition["sample_gain_mismatch_standard_deviation"], + circuit.edge_count, + ), + ) + + +def run_job(job: dict) -> dict: + circuit = GridCircuit() + task = job["task"] + condition = job["condition"] + task_index = job["task_index"] + device_seed = job["device_seed"] + dataset = RingClassificationDataset( + inputs_v=np.asarray(task["inputs_v"], dtype=float).T, + labels_v=( + 2.0 * np.asarray(task["classes"], dtype=float) - 1.0 + ) * 0.018, + ) + imperfection = GridSquareLawImperfection.sample_appendix_c( + circuit.edge_count, device_seed) + common_pedestal, pedestal_mismatch, gain_mismatch = fixed_edge_errors( + circuit, task_index, device_seed, condition) + sampler = CorrelatedDoubleSampleHold( + common_sample_gain=1.0, + sample_gain_mismatch=gain_mismatch, + common_pedestal_offset_v_per_s=common_pedestal, + pedestal_mismatch_v_per_s=pedestal_mismatch, + sample_noise_standard_deviation_v_per_s=( + condition["sample_noise_standard_deviation_v_per_s"]), + refresh_interval_updates=condition["refresh_interval_updates"], + ) + kwargs = dict( + circuit=circuit, + initial_gates=np.asarray(task["initial_gates_v"], dtype=float), + dataset=dataset, + imperfection=imperfection, + method=condition["method"], + correlated_sample_hold=sampler, + autozero_seed=device_seed + 9_000_003 + task_index, + ) + try: + if condition["method"].startswith("overclamp"): + result = train_grid_classifier( + **kwargs, + epochs=job["overclamp_epochs"], + overclamp_time_seconds_per_v=0.0025, + record_every=10, + early_stop_perfect_checkpoints=3, + ) + else: + result = train_grid_classifier( + **kwargs, + epochs=job["standard_epochs"], + standard_learning_time_seconds=1e-3, + record_every=50, + ) + except RuntimeError as error: + return { + "condition": condition["name"], + "task_index": task_index, + "device_seed": device_seed, + "input_diameter_v": task["input_diameter_v"], + "status": "circuit_solver_failure", + "failure_message": str(error), + "classification_error": 1.0, + "zero_error": False, + "neutral_samples": None, + "active_samples": None, + "local_updates": None, + "applied_rate_rmse_v_per_s": None, + "max_abs_clamp_displacement_v": None, + } + return { + "condition": condition["name"], + "task_index": task_index, + "device_seed": device_seed, + "input_diameter_v": task["input_diameter_v"], + "status": "completed", + "classification_error": result["classification_error"], + "zero_error": result["classification_error"] == 0.0, + "hinge_loss_v2": result["hinge_loss_v2"], + "neutral_samples": result["autozero_samples"], + "active_samples": result["autozero_active_samples"], + "local_updates": result["local_updates"], + "neutral_sample_fraction_per_update": ( + result["autozero_sample_fraction_per_update"]), + "applied_rate_rmse_v_per_s": ( + result["autozero_applied_rate_rmse_v_per_s"]), + "max_abs_clamp_displacement_v": ( + result["max_abs_clamp_displacement_v"]), + } + + +def reference_records(path: Path) -> dict[tuple[int, int], dict]: + report = json.loads(path.read_text()) + return { + (record["task_index"], record["device_seed"]): record + for record in report["records"] + } + + +def summarize(records: list[dict], settings: list[dict], reference: dict) -> dict: + output = {} + for condition in settings: + selected = [ + record for record in records + if record["condition"] == condition["name"] + ] + completed = [ + record for record in selected if record["status"] == "completed" + ] + errors = np.asarray([ + record["classification_error"] for record in selected + ]) + overclamp_errors = np.asarray([ + reference[(record["task_index"], record["device_seed"])] + ["methods"]["overclamp"]["classification_error"] + for record in selected + ]) + output[condition["name"]] = { + **condition, + "trials": len(selected), + "mean_classification_error": float(np.mean(errors)), + "median_classification_error": float(np.median(errors)), + "zero_error_fraction": float(np.mean(errors == 0.0)), + "solver_failure_fraction": float(np.mean([ + record["status"] != "completed" for record in selected + ])), + "lower_error_than_overclamp_fraction": float(np.mean( + errors < overclamp_errors)), + "equal_error_to_overclamp_fraction": float(np.mean( + errors == overclamp_errors)), + "higher_error_than_overclamp_fraction": float(np.mean( + errors > overclamp_errors)), + "mean_neutral_sample_fraction_per_update": float(np.mean([ + record["neutral_sample_fraction_per_update"] + for record in completed + ])), + "median_applied_rate_rmse_v_per_s": float(np.median([ + record["applied_rate_rmse_v_per_s"] + for record in completed + ])), + "mean_max_abs_clamp_displacement_v": float(np.mean([ + record["max_abs_clamp_displacement_v"] + for record in completed + ])), + } + return output + + +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( + "--reference", type=Path, + default=Path( + "results/physical_bias/p5_full_grid_bias_crossover.json")) + parser.add_argument( + "--output", type=Path, + default=Path( + "results/physical_bias/p9_grid_correlated_autozero.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("--overclamp-epochs", type=int, default=1000) + parser.add_argument("--workers", type=int, default=16) + parser.add_argument("--conditions", default="all") + 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(",")) + all_settings = conditions() + if args.conditions == "all": + settings = all_settings + else: + requested = set(args.conditions.split(",")) + settings = [ + condition for condition in all_settings + if condition["name"] in requested + ] + missing = requested - {condition["name"] for condition in settings} + if missing: + raise ValueError(f"unknown conditions: {sorted(missing)}") + jobs = [{ + "condition": condition, + "task_index": task_index, + "task": task, + "device_seed": device_seed, + "standard_epochs": args.standard_epochs, + "overclamp_epochs": args.overclamp_epochs, + } for condition in settings + for task_index, task in enumerate(tasks) + for device_seed in device_seeds] + 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"])) + references = reference_records(args.reference) + report = { + "analysis": "physical_grid_correlated_autozero_p9", + "confirmatory": False, + "autodiff_used": False, + "source_protocol": str(args.protocol), + "reference_results": str(args.reference), + "protocol": { + "task_count": len(tasks), + "rotations_per_input_diameter": args.rotations, + "device_seeds": device_seeds, + "trials_per_condition": len(tasks) * len(device_seeds), + "component_imperfection": { + "measurement_gain_standard_deviation": 0.01, + "twin_input_mismatch_standard_deviation_v": 0.001, + "multiplier_output_offset_standard_deviation_v_per_s": 2.3, + }, + "sampling_operation": ( + "Each edge samples neutral and active outputs through matched " + "local paths and applies their difference. Common sample-path " + "pedestal cancels without reading component parameters."), + "standard_epochs": args.standard_epochs, + "overclamp_epochs_maximum": args.overclamp_epochs, + "conditions": settings, + }, + "records": records, + "summary": summarize(records, settings, references), + } + 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() |
