diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 13:00:18 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 13:00:18 -0500 |
| commit | 144c6a4333151bdfba976d3bf7f2b8be9645f895 (patch) | |
| tree | da21be00003ea84df2c2948873bde6ca970012c8 /experiments | |
| parent | fb5223f6ba409298af19c30f50b78e07b7a49d9e (diff) | |
exp: add device-level imperfection crossover
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/physical_imperfection_p3.py | 392 |
1 files changed, 392 insertions, 0 deletions
diff --git a/experiments/physical_imperfection_p3.py b/experiments/physical_imperfection_p3.py new file mode 100644 index 0000000..f573edd --- /dev/null +++ b/experiments/physical_imperfection_p3.py @@ -0,0 +1,392 @@ +#!/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() |
