summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/physical_grid_correlated_autozero_p9.py381
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()