summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 13:20:53 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 13:20:53 -0500
commit591b4c833f8b36be80f258a164b4737319a54526 (patch)
treea0a938613da0b170090a3d74ecbc3d974abe7ac9 /experiments
parent6aaa8d54d45f67a08ffcc943412eddc3e66ed4d6 (diff)
exp: add physical grid bias crossover
Diffstat (limited to 'experiments')
-rw-r--r--experiments/physical_grid_bias_p5.py378
1 files changed, 378 insertions, 0 deletions
diff --git a/experiments/physical_grid_bias_p5.py b/experiments/physical_grid_bias_p5.py
new file mode 100644
index 0000000..8757ad2
--- /dev/null
+++ b/experiments/physical_grid_bias_p5.py
@@ -0,0 +1,378 @@
+#!/usr/bin/env python3
+"""Compare SDIL and overclamping on the reconstructed Figure-5 grid.
+
+Each device draw uses the component-imperfection model from Appendix C of
+Dillavou et al. A frozen per-edge SDIL predictor is fitted from neutral
+observations, where the free and clamped circuit states are identical.
+"""
+
+from __future__ import annotations
+
+import argparse
+from concurrent.futures import ProcessPoolExecutor, as_completed
+import json
+from pathlib import Path
+import sys
+import time
+
+import numpy as np
+
+ROOT = Path(__file__).resolve().parents[1]
+sys.path.insert(0, str(ROOT))
+
+from sdil.physical_grid import ( # noqa: E402
+ EdgePolynomialPredictor,
+ GridCircuit,
+ GridSquareLawImperfection,
+ RingClassificationDataset,
+ edge_voltage_drops,
+ fit_edge_predictor,
+ solve_grid_state,
+ train_grid_classifier,
+)
+
+
+METHODS = (
+ "clean",
+ "raw",
+ "constant",
+ "sdil",
+ "oracle_neutral",
+ "overclamp",
+ "overclamp_sdil",
+)
+
+
+def select_tasks(protocol: dict, rotations: int) -> list[dict]:
+ standard = [
+ record for record in protocol["experiments"]
+ if record["method"] == "standard"
+ ]
+ diameters = protocol["protocol_checks"]["input_diameters_v"]
+ selected = []
+ for diameter in diameters:
+ candidates = sorted(
+ (
+ record for record in standard
+ if abs(record["input_diameter_v"] - diameter) < 1e-12
+ ),
+ key=lambda record: record["classes"],
+ )
+ selected.extend(candidates[:rotations])
+ return selected
+
+
+def collect_neutral_observations(
+ circuit: GridCircuit,
+ gates: np.ndarray,
+ imperfection: GridSquareLawImperfection,
+ *,
+ count: int,
+ seed: int,
+) -> tuple[np.ndarray, np.ndarray]:
+ rng = np.random.default_rng(seed)
+ local_states = []
+ neutral_measurements = []
+ previous_state = None
+ for _ in range(count):
+ inputs = rng.uniform(
+ circuit.low_voltage, circuit.high_voltage, size=2)
+ state = solve_grid_state(
+ circuit,
+ gates,
+ circuit.source_values(*inputs),
+ initial_state=previous_state,
+ )
+ previous_state = state
+ drops = edge_voltage_drops(circuit, state)
+ local_states.append(drops)
+ neutral_measurements.append(imperfection.neutral_bias(
+ circuit.measured_learning_rate, drops))
+ return np.asarray(local_states), np.asarray(neutral_measurements)
+
+
+def make_predictor(
+ local_states: np.ndarray,
+ neutral_measurements: np.ndarray,
+ *,
+ degree: int,
+) -> EdgePolynomialPredictor:
+ center = np.mean(local_states, axis=0)
+ scale = np.maximum(np.std(local_states, axis=0), 1e-6)
+ predictor = EdgePolynomialPredictor.zeros(
+ center, scale, degree=degree)
+ fit_edge_predictor(predictor, local_states, neutral_measurements)
+ return predictor
+
+
+def calibration_error(
+ predictor: EdgePolynomialPredictor,
+ local_states: np.ndarray,
+ neutral_measurements: np.ndarray,
+) -> float:
+ predictions = np.asarray([
+ predictor.predict(state) for state in local_states
+ ])
+ return float(np.sqrt(np.mean(np.square(
+ predictions - neutral_measurements))))
+
+
+def run_trial(job: dict) -> dict:
+ circuit = GridCircuit()
+ task = job["task"]
+ device_seed = job["device_seed"]
+ task_index = job["task_index"]
+ 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,
+ )
+ initial_gates = np.asarray(task["initial_gates_v"], dtype=float)
+ imperfection = GridSquareLawImperfection.sample_appendix_c(
+ circuit.edge_count, device_seed)
+ calibration_states, calibration_measurements = (
+ collect_neutral_observations(
+ circuit,
+ initial_gates,
+ imperfection,
+ count=job["calibration_observations"],
+ seed=device_seed + 1_000_003 + task_index,
+ )
+ )
+ heldout_states, heldout_measurements = collect_neutral_observations(
+ circuit,
+ initial_gates,
+ imperfection,
+ count=job["heldout_observations"],
+ seed=device_seed + 2_000_003 + task_index,
+ )
+ constant_predictor = make_predictor(
+ calibration_states, calibration_measurements, degree=0)
+ sdil_predictor = make_predictor(
+ calibration_states, calibration_measurements, degree=2)
+ predictors = {
+ "constant": constant_predictor,
+ "sdil": sdil_predictor,
+ "overclamp_sdil": sdil_predictor,
+ }
+ results = {}
+ for method in job["methods"]:
+ start = time.perf_counter()
+ if method.startswith("overclamp"):
+ result = train_grid_classifier(
+ circuit,
+ initial_gates,
+ dataset,
+ imperfection,
+ method=method,
+ predictor=predictors.get(method),
+ epochs=job["overclamp_epochs"],
+ overclamp_time_seconds_per_v=0.0025,
+ record_every=10,
+ early_stop_perfect_checkpoints=3,
+ )
+ else:
+ result = train_grid_classifier(
+ circuit,
+ initial_gates,
+ dataset,
+ imperfection,
+ method=method,
+ predictor=predictors.get(method),
+ epochs=job["standard_epochs"],
+ standard_learning_time_seconds=1e-3,
+ record_every=50,
+ )
+ result["wall_seconds"] = float(time.perf_counter() - start)
+ results[method] = result
+ return {
+ "task_index": task_index,
+ "source_file": task["source_file"],
+ "input_diameter_v": task["input_diameter_v"],
+ "classes": task["classes"],
+ "device_seed": device_seed,
+ "calibration": {
+ "observations": job["calibration_observations"],
+ "heldout_observations": job["heldout_observations"],
+ "raw_heldout_rmse_v_per_s": float(np.sqrt(np.mean(
+ np.square(heldout_measurements)))),
+ "constant_heldout_rmse_v_per_s": calibration_error(
+ constant_predictor, heldout_states, heldout_measurements),
+ "sdil_heldout_rmse_v_per_s": calibration_error(
+ sdil_predictor, heldout_states, heldout_measurements),
+ },
+ "methods": results,
+ }
+
+
+def paired_summary(
+ records: list[dict], first: str, second: str
+) -> dict:
+ first_errors = np.asarray([
+ record["methods"][first]["classification_error"]
+ for record in records
+ ])
+ second_errors = np.asarray([
+ record["methods"][second]["classification_error"]
+ for record in records
+ ])
+ return {
+ "comparison": f"{first}_versus_{second}",
+ "first_lower_error_fraction": float(np.mean(
+ first_errors < second_errors)),
+ "equal_error_fraction": float(np.mean(first_errors == second_errors)),
+ "first_higher_error_fraction": float(np.mean(
+ first_errors > second_errors)),
+ }
+
+
+def summarize(records: list[dict], methods: tuple[str, ...]) -> dict:
+ by_method = {}
+ for method in methods:
+ errors = np.asarray([
+ record["methods"][method]["classification_error"]
+ for record in records
+ ])
+ displacements = np.asarray([
+ record["methods"][method]["max_abs_clamp_displacement_v"]
+ for record in records
+ ])
+ by_method[method] = {
+ "mean_classification_error": float(np.mean(errors)),
+ "median_classification_error": float(np.median(errors)),
+ "zero_error_fraction": float(np.mean(errors == 0.0)),
+ "mean_max_abs_clamp_displacement_v": float(np.mean(displacements)),
+ "median_wall_seconds": float(np.median([
+ record["methods"][method]["wall_seconds"]
+ for record in records
+ ])),
+ }
+ calibration_keys = (
+ "raw_heldout_rmse_v_per_s",
+ "constant_heldout_rmse_v_per_s",
+ "sdil_heldout_rmse_v_per_s",
+ )
+ calibration = {
+ key: float(np.median([
+ record["calibration"][key] for record in records
+ ]))
+ for key in calibration_keys
+ }
+ pairwise = []
+ for first, second in (
+ ("sdil", "raw"),
+ ("sdil", "constant"),
+ ("sdil", "overclamp"),
+ ("overclamp_sdil", "overclamp"),
+ ):
+ if first in methods and second in methods:
+ pairwise.append(paired_summary(records, first, second))
+ return {
+ "trials": len(records),
+ "by_method": by_method,
+ "median_calibration_error": calibration,
+ "paired_classification": pairwise,
+ }
+
+
+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/p5_grid_bias_crossover.json"))
+ parser.add_argument("--rotations", type=int, default=2)
+ parser.add_argument(
+ "--device-seeds", default="20260829,20260830,20260831,20260832")
+ parser.add_argument("--calibration-observations", type=int, default=16)
+ parser.add_argument("--heldout-observations", type=int, default=64)
+ 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=8)
+ parser.add_argument("--methods", default=",".join(METHODS))
+ 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(","))
+ methods = tuple(args.methods.split(","))
+ unknown = set(methods) - set(METHODS)
+ if unknown:
+ raise ValueError(f"unknown methods: {sorted(unknown)}")
+ jobs = []
+ for task_index, task in enumerate(tasks):
+ for device_seed in device_seeds:
+ jobs.append({
+ "task_index": task_index,
+ "task": task,
+ "device_seed": device_seed,
+ "calibration_observations": args.calibration_observations,
+ "heldout_observations": args.heldout_observations,
+ "standard_epochs": args.standard_epochs,
+ "overclamp_epochs": args.overclamp_epochs,
+ "methods": methods,
+ })
+ records = []
+ with ProcessPoolExecutor(max_workers=args.workers) as executor:
+ futures = [executor.submit(run_trial, job) for job in jobs]
+ for completed, future in enumerate(as_completed(futures), start=1):
+ record = future.result()
+ records.append(record)
+ compact = ", ".join(
+ f"{method}={record['methods'][method]['classification_error']:.3f}"
+ for method in methods)
+ print(
+ f"trial {completed}/{len(jobs)} "
+ f"task={record['task_index']} seed={record['device_seed']}: "
+ f"{compact}",
+ flush=True,
+ )
+ records.sort(key=lambda record: (
+ record["task_index"], record["device_seed"]))
+ report = {
+ "analysis": "reconstructed_figure5_grid_hardware_bias_p5",
+ "confirmatory": False,
+ "autodiff_used": False,
+ "source_protocol": str(args.protocol),
+ "protocol": {
+ "task_count": len(tasks),
+ "rotations_per_input_diameter": args.rotations,
+ "input_diameters_v": sorted({
+ task["input_diameter_v"] for task in tasks}),
+ "device_seeds": 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,
+ },
+ "calibration_observations_per_trial": args.calibration_observations,
+ "calibration_scope": (
+ "per-edge degree-2 local predictor fitted once from neutral "
+ "free=clamped observations and then frozen"),
+ "constant_baseline": (
+ "per-edge degree-0 predictor fitted from the same observations"),
+ "standard_epochs": args.standard_epochs,
+ "standard_learning_time_seconds": 1e-3,
+ "overclamp_epochs_maximum": args.overclamp_epochs,
+ "overclamp_nudging": 32.0 / 129.0,
+ "overclamp_time_seconds_per_v": 0.0025,
+ "overclamp_early_stop_perfect_checkpoints": 3,
+ "methods": methods,
+ },
+ "records": records,
+ "summary": summarize(records, methods),
+ }
+ 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()