From 591b4c833f8b36be80f258a164b4737319a54526 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Sat, 29 Aug 2026 13:20:53 -0500 Subject: exp: add physical grid bias crossover --- experiments/physical_grid_bias_p5.py | 378 +++++++++++++++++++++++++++++++++++ 1 file changed, 378 insertions(+) create mode 100644 experiments/physical_grid_bias_p5.py 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() -- cgit v1.2.3