#!/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()