#!/usr/bin/env python3 """P3: device-seed crossover under Appendix-C component imperfections.""" from __future__ import annotations import argparse import json from pathlib import Path import sys import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from sdil.physical_coupled import Circuit, Task # noqa: E402 from sdil.physical_imperfection import ( # noqa: E402 DifferentialSquareLawImperfection, LocalPolynomialPredictor, edge_voltage_drops, fit_polynomial_predictor, simulate_imperfect_alternating_tasks, ) STANDARD_METHODS = ("clean", "raw", "constant", "sdil", "oracle_neutral") OVERCLAMP_METHODS = ("overclamp", "overclamp_sdil") LABELS = { "overclamp": "overclamping", "overclamp_sdil": "SDIL + overclamping", } COLORS = { "overclamp": "#CC3311", "overclamp_sdil": "#0077BB", } def make_tasks(circuit: Circuit, pair_name: str) -> tuple[Task, Task]: beta_label = 0.14 if pair_name == "experiment_1" else 0.18 return ( Task("alpha", circuit.high, 0.31), Task("beta", circuit.low, beta_label), ) def calibrate( circuit: Circuit, hardware: DifferentialSquareLawImperfection, observation_count: int, ) -> tuple[LocalPolynomialPredictor, LocalPolynomialPredictor, dict]: if observation_count < 3: raise ValueError("a quadratic predictor requires at least three states") outputs = np.linspace( circuit.low + 0.01, circuit.high - 0.01, observation_count, ) states = np.asarray([ edge_voltage_drops(circuit, output) for output in outputs]) neutral = np.asarray([ hardware.neutral_bias(circuit, output) for output in outputs]) center = np.mean(states, axis=0) scale = np.ptp(states, axis=0) constant = LocalPolynomialPredictor.zeros(center, scale, degree=0) sdil = LocalPolynomialPredictor.zeros(center, scale, degree=2) constant_count = fit_polynomial_predictor(constant, states, neutral) sdil_count = fit_polynomial_predictor(sdil, states, neutral) if constant_count != sdil_count: raise AssertionError("neutral observation budgets disagree") evaluation_outputs = np.linspace( circuit.low + 0.005, circuit.high - 0.005, 257) evaluation_states = np.asarray([ edge_voltage_drops(circuit, output) for output in evaluation_outputs]) evaluation_neutral = np.asarray([ hardware.neutral_bias(circuit, output) for output in evaluation_outputs]) constant_rmse = float(np.sqrt(np.mean([ np.square(measurement - constant.predict(state)) for state, measurement in zip(evaluation_states, evaluation_neutral) ]))) sdil_rmse = float(np.sqrt(np.mean([ np.square(measurement - sdil.predict(state)) for state, measurement in zip(evaluation_states, evaluation_neutral) ]))) return constant, sdil, { "neutral_observations_each": observation_count, "observation_output_range_v": [float(outputs[0]), float(outputs[-1])], "constant_heldout_rmse_v_per_s": constant_rmse, "sdil_heldout_rmse_v_per_s": sdil_rmse, "constant_coefficients": constant.coefficients.tolist(), "sdil_coefficients": sdil.coefficients.tolist(), } def run_one( circuit: Circuit, tasks: tuple[Task, Task], hardware: DifferentialSquareLawImperfection, *, method: str, period: float, cycles: int, predictor: LocalPolynomialPredictor | None = None, eta: float | None = None, ) -> dict: kwargs = {} if method.startswith("overclamp"): if eta is None: raise ValueError("overclamping requires eta") kwargs = { "overclamp_nudging": eta, "overclamp_magnitude": circuit.high, } result = simulate_imperfect_alternating_tasks( circuit, tasks, hardware, method=method, period_seconds=period, cycles=cycles, initial_gates=np.asarray((4.0, 4.0)), predictor=predictor, **kwargs, ) if eta is not None: result["overclamp_eta"] = eta return result def minimum_passing_record( records: list[dict], method: str, threshold: float ) -> dict | None: candidates = sorted( ( record for record in records if record["method"] == method and record["mean_combined_error"] <= threshold ), key=lambda record: record["overclamp_eta"], ) return candidates[0] if candidates else None def summarize_pair(devices: list[dict], threshold: float) -> dict: eta_ratios = [] exposure_ratios = [] paired_successes = 0 combination_no_larger = 0 standard_by_method = {method: [] for method in STANDARD_METHODS} for device in devices: for record in device["standard_clamping"]: standard_by_method[record["method"]].append( record["mean_combined_error"]) baseline = minimum_passing_record( device["overclamp_sweep"], "overclamp", threshold) combined = minimum_passing_record( device["overclamp_sweep"], "overclamp_sdil", threshold) if combined is not None and ( baseline is None or combined["overclamp_eta"] <= baseline["overclamp_eta"] ): combination_no_larger += 1 if baseline is not None and combined is not None: paired_successes += 1 eta_ratios.append( baseline["overclamp_eta"] / combined["overclamp_eta"]) exposure_ratios.append( baseline["clamp_displacement_l2_time_v2_s"] / combined["clamp_displacement_l2_time_v2_s"]) return { "device_count": len(devices), "standard_clamping_error": { method: { "median": float(np.median(values)), "minimum": float(np.min(values)), "maximum": float(np.max(values)), } for method, values in standard_by_method.items() }, "paired_devices_reaching_threshold": paired_successes, "combination_no_larger_eta_fraction": ( combination_no_larger / len(devices)), "median_eta_reduction_at_passing_endpoint": ( None if not eta_ratios else float(np.median(eta_ratios))), "median_clamp_l2_exposure_reduction_at_passing_endpoint": ( None if not exposure_ratios else float(np.median(exposure_ratios))), "eta_reduction_by_device": eta_ratios, "clamp_l2_exposure_reduction_by_device": exposure_ratios, } def aggregate( devices: list[dict], method: str, etas: list[float], key: str ) -> np.ndarray: rows = [] for eta in etas: values = np.asarray([ record[key] for device in devices for record in device["overclamp_sweep"] if record["method"] == method and record["overclamp_eta"] == eta ]) rows.append(( float(np.median(values)), float(np.quantile(values, 0.25)), float(np.quantile(values, 0.75)), )) return np.asarray(rows) def plot_report(report: dict, output: Path) -> None: etas = report["protocol"]["overclamp_eta_grid"] fig, axes = plt.subplots(2, 2, figsize=(9.2, 7.0), sharex="col") for column, name in enumerate(("experiment_1", "experiment_2")): devices = report["pairs"][name]["devices"] for method in OVERCLAMP_METHODS: error = aggregate( devices, method, etas, "mean_combined_error") exposure = aggregate( devices, method, etas, "clamp_displacement_l2_time_v2_s") axes[0, column].loglog( etas, error[:, 0], "o-", color=COLORS[method], label=LABELS[method]) axes[0, column].fill_between( etas, np.maximum(error[:, 1], 1e-32), error[:, 2], color=COLORS[method], alpha=0.16) axes[1, column].loglog( etas, exposure[:, 0], "o-", color=COLORS[method]) axes[1, column].fill_between( etas, exposure[:, 1], exposure[:, 2], color=COLORS[method], alpha=0.16) axes[0, column].axhline( report["protocol"]["precision_threshold_v2"], color="#666666", linestyle="--", linewidth=1.0) axes[0, column].set_title( f"{chr(ord('A') + column)} {name.replace('_', ' ')}: error") axes[0, column].set_ylabel("combined task error") axes[1, column].set_title( f"{chr(ord('C') + column)} clamp exposure") axes[1, column].set_xlabel("overclamping nudging strength η") axes[1, column].set_ylabel("Σ duration × displacement² (V²s)") for row in range(2): axes[row, column].grid(alpha=0.18) axes[0, 0].legend(frameon=False, fontsize=8) fig.suptitle( "Appendix-C device imperfections: SDIL complements overclamping", fontsize=11) fig.tight_layout() output.parent.mkdir(parents=True, exist_ok=True) fig.savefig(output, dpi=180) plt.close(fig) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument( "--output", type=Path, default=Path("results/physical_bias/p3_device_crossover.json")) parser.add_argument( "--figure", type=Path, default=Path("results/figs/physical_bias_p3_device_crossover.png")) parser.add_argument("--device-seeds", type=int, default=8) parser.add_argument("--seed", type=int, default=20260829) parser.add_argument("--neutral-observations", type=int, default=16) parser.add_argument("--period", type=float, default=0.02) parser.add_argument("--cycles", type=int, default=300) parser.add_argument( "--etas", type=float, nargs="+", default=(0.005, 0.01, 0.025, 0.05, 0.1, 0.25)) parser.add_argument("--precision-threshold-v2", type=float, default=1e-8) return parser.parse_args() def main() -> None: args = parse_args() circuit = Circuit() etas = list(args.etas) report = { "analysis": "appendix_c_device_imperfection_crossover_p3", "confirmatory": False, "physical_hardware_demonstration": False, "autodiff_used": False, "protocol": { "paper_source": "Dillavou et al. arXiv:2505.22887v2 Appendix C", "gain_standard_deviation": 0.01, "twin_mismatch_standard_deviation_v": 0.001, "multiplier_offset_standard_deviation_v_per_s": 2.3, "device_seeds": args.device_seeds, "seed_start": args.seed, "neutral_observations_each": args.neutral_observations, "period_seconds": args.period, "cycles": args.cycles, "overclamp_eta_grid": etas, "overclamp_target_magnitude_v": circuit.high, "precision_threshold_v2": args.precision_threshold_v2, }, "pairs": {}, } for pair_index, pair_name in enumerate(("experiment_1", "experiment_2")): tasks = make_tasks(circuit, pair_name) devices = [] for device_index in range(args.device_seeds): seed = args.seed + 1000 * pair_index + device_index hardware = DifferentialSquareLawImperfection.sample_appendix_c(seed) constant, sdil, calibration = calibrate( circuit, hardware, args.neutral_observations) standard_records = [] for method in STANDARD_METHODS: predictor = None if method == "constant": predictor = constant elif method == "sdil": predictor = sdil standard_records.append(run_one( circuit, tasks, hardware, method=method, period=args.period, cycles=args.cycles, predictor=predictor, )) overclamp_records = [] for eta in etas: for method in OVERCLAMP_METHODS: predictor = sdil if method == "overclamp_sdil" else None overclamp_records.append(run_one( circuit, tasks, hardware, method=method, period=args.period, cycles=args.cycles, predictor=predictor, eta=eta, )) devices.append({ "seed": seed, "hardware": hardware.as_dict(), "calibration": calibration, "standard_clamping": standard_records, "overclamp_sweep": overclamp_records, }) print( f"{pair_name}: completed device " f"{device_index + 1}/{args.device_seeds}", flush=True, ) report["pairs"][pair_name] = { "tasks": [ { "name": task.name, "input_voltage": task.input_voltage, "label_voltage": task.label_voltage, } for task in tasks ], "devices": devices, "summary": summarize_pair( devices, args.precision_threshold_v2), } report["summary"] = { "combination_no_larger_eta_fraction_by_pair": { name: pair["summary"]["combination_no_larger_eta_fraction"] for name, pair in report["pairs"].items() }, "median_eta_reduction_by_pair": { name: pair["summary"][ "median_eta_reduction_at_passing_endpoint"] for name, pair in report["pairs"].items() }, "median_clamp_l2_exposure_reduction_by_pair": { name: pair["summary"][ "median_clamp_l2_exposure_reduction_at_passing_endpoint"] for name, pair in report["pairs"].items() }, } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(report, indent=2) + "\n") plot_report(report, args.figure) print(json.dumps(report["summary"], indent=2)) print(f"wrote {args.output}") print(f"wrote {args.figure}") if __name__ == "__main__": main()