#!/usr/bin/env python3 """P1: test SDIL on the measured-state-dependent two-edge surrogate.""" 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 ( # noqa: E402 Circuit, LocalAffineBias, LocalPredictor, Task, calibrate_predictor, simulate_alternating_tasks, ) METHODS = ( "raw", "same_rms_noise", "frozen_constant", "frozen_sdil", "online_constant", "overclamp", "oracle", ) LABELS = { "raw": "structured bias", "same_rms_noise": "same-RMS noise", "frozen_constant": "constant calibration", "frozen_sdil": "SDIL", "online_constant": "online recalibration", "overclamp": "overclamp analogue", "oracle": "oracle subtraction", } COLORS = { "raw": "#CC3311", "same_rms_noise": "#BBBBBB", "frozen_constant": "#EE7733", "frozen_sdil": "#0077BB", "online_constant": "#AA4499", "overclamp": "#228833", "oracle": "#000000", } def load_pair(source: dict, name: str, strength: float) -> dict: record = source["pairs"][name] reference = np.asarray(record["reference_gate"], dtype=float) affine = record["local_affine_model"] field = LocalAffineBias( reference_gate=reference, bias_at_reference=np.asarray( affine["bias_at_reference_v_per_s"], dtype=float), local_slopes=np.asarray(affine["local_slopes_per_s"], dtype=float), ) states = [] for trace in record["traces"]: states.append(np.column_stack(( trace["retained_gate_minus"], trace["retained_gate_plus"]))) states = np.vstack(states) feature_scale = np.maximum(np.ptp(states, axis=0), 0.25) circuit = Circuit() beta_label = 0.14 if name == "experiment_1" else 0.18 tasks = ( Task("alpha", circuit.high, 0.31), Task("beta", circuit.low, beta_label), ) measurement = lambda gates: field(gates, strength) # noqa: E731 constant = LocalPredictor.zeros(reference, feature_scale, affine=False) sdil = LocalPredictor.zeros(reference, feature_scale, affine=True) calibration = { "epochs": 30, "learning_rate": 0.2, "states": int(len(states)), } calibration["neutral_observations_each"] = calibrate_predictor( constant, states, measurement, epochs=calibration["epochs"], learning_rate=calibration["learning_rate"], ) sdil_count = calibrate_predictor( sdil, states, measurement, epochs=calibration["epochs"], learning_rate=calibration["learning_rate"], ) if sdil_count != calibration["neutral_observations_each"]: raise AssertionError("calibration observation budgets disagree") bias_samples = np.asarray([measurement(state) for state in states]) calibration["constant_rmse"] = float(np.sqrt(np.mean([ np.mean((measurement(state) - constant.predict(state)) ** 2) for state in states ]))) calibration["sdil_rmse"] = float(np.sqrt(np.mean([ np.mean((measurement(state) - sdil.predict(state)) ** 2) for state in states ]))) calibration["constant_coefficients"] = constant.coefficients.tolist() calibration["sdil_coefficients"] = sdil.coefficients.tolist() return { "field": field, "states": states, "circuit": circuit, "tasks": tasks, "constant": constant, "sdil": sdil, "noise_std": np.sqrt(np.mean(bias_samples * bias_samples, axis=0)), "calibration": calibration, } def run_method( pair: dict, method: str, period: float, cycles: int, initial_gates: np.ndarray, strength: float, seed: int, ) -> dict: predictor = None if method in {"frozen_constant", "online_constant"}: predictor = pair["constant"] elif method == "frozen_sdil": predictor = pair["sdil"] return simulate_alternating_tasks( pair["circuit"], pair["tasks"], pair["field"], method=method, period_seconds=period, cycles=cycles, initial_gates=initial_gates, bias_strength=strength, predictor=predictor, online_predictor_rate=0.05, noise_standard_deviation=pair["noise_std"], seed=seed, ) def plot_report(report: dict, output: Path) -> None: fig, axes = plt.subplots(2, 2, figsize=(9.2, 7.0), sharex="col") for column, name in enumerate(("experiment_1", "experiment_2")): records = report["pairs"][name]["period_sweep"] for method in METHODS: selected = [record for record in records if record["method"] == method] period = np.asarray([record["period_seconds"] for record in selected]) error = np.asarray([record["mean_combined_error"] for record in selected]) span = np.asarray([record["mean_cycle_span"] for record in selected]) axes[0, column].loglog( period, error, "o-", color=COLORS[method], linewidth=1.3, markersize=3.5, label=LABELS[method]) axes[1, column].loglog( period, np.maximum(span, 1e-12), "o-", color=COLORS[method], linewidth=1.3, markersize=3.5, label=LABELS[method]) axes[0, column].set_title( f"{chr(ord('A') + column)} {name.replace('_', ' ')}: error floor") axes[0, column].set_ylabel("combined task error") axes[1, column].set_title( f"{chr(ord('C') + column)} {name.replace('_', ' ')}: cycle span") axes[1, column].set_xlabel("task-switching period (s)") axes[1, column].set_ylabel("gate-space cycle span (V)") for row in range(2): axes[row, column].grid(alpha=0.18) axes[0, 0].legend(frameon=False, fontsize=7, ncol=2) fig.suptitle( "Measured-state-dependent two-edge surrogate: frozen local SDIL", 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( "--state-dependence-json", type=Path, default=Path("results/physical_bias/p0_state_dependence.json")) parser.add_argument( "--output", type=Path, default=Path("results/physical_bias/p1_surrogate.json")) parser.add_argument( "--figure", type=Path, default=Path("results/figs/physical_bias_p1_surrogate.png")) parser.add_argument("--minimum-cycles", type=int, default=120) parser.add_argument("--total-nominal-time", type=float, default=6.0) parser.add_argument( "--periods", type=float, nargs="+", default=(0.002, 0.005, 0.01, 0.02, 0.05, 0.1, 0.2)) parser.add_argument("--bias-strength", type=float, default=1.0) parser.add_argument("--seed", type=int, default=20260806) return parser.parse_args() def main() -> None: args = parse_args() source = json.loads(args.state_dependence_json.read_text()) report = { "analysis": "physical_measured_state_dependent_surrogate_p1", "confirmatory": False, "physical_hardware_demonstration": False, "autodiff_used": False, "source_analysis": str(args.state_dependence_json), "protocol": { "minimum_cycles": args.minimum_cycles, "total_nominal_time_seconds": args.total_nominal_time, "periods_seconds": args.periods, "bias_strength": args.bias_strength, "initial_gates": [4.0, 4.0], "methods": METHODS, "overclamp_scope": ( "leading-order Appendix-F analogue; not the published classification endpoint" ), }, "pairs": {}, } initial_gates = np.asarray(report["protocol"]["initial_gates"], dtype=float) for pair_index, name in enumerate(("experiment_1", "experiment_2")): pair = load_pair(source, name, args.bias_strength) period_records = [] for period in args.periods: cycles = max( args.minimum_cycles, int(np.ceil(args.total_nominal_time / period)), ) for method_index, method in enumerate(METHODS): period_records.append(run_method( pair, method, period, cycles, initial_gates, args.bias_strength, args.seed + 1000 * pair_index + 10 * method_index, )) by_method = { method: [record for record in period_records if record["method"] == method] for method in METHODS } sdil_error = np.asarray([ record["mean_combined_error"] for record in by_method["frozen_sdil"]]) constant_error = np.asarray([ record["mean_combined_error"] for record in by_method["frozen_constant"]]) raw_error = np.asarray([ record["mean_combined_error"] for record in by_method["raw"]]) oracle_error = np.asarray([ record["mean_combined_error"] for record in by_method["oracle"]]) valid_gap = raw_error > oracle_error + 1e-16 gap_closed = ( (raw_error[valid_gap] - sdil_error[valid_gap]) / (raw_error[valid_gap] - oracle_error[valid_gap]) ) overclamp_error = np.asarray([ record["mean_combined_error"] for record in by_method["overclamp"]]) report["pairs"][name] = { "calibration": pair["calibration"], "noise_standard_deviation_v_per_s": pair["noise_std"].tolist(), "period_sweep": period_records, "summary": { "sdil_beats_frozen_constant_all_periods": bool(np.all( sdil_error < constant_error)), "median_raw_to_oracle_gap_closed_by_sdil": ( None if len(gap_closed) == 0 else float(np.median(gap_closed))), "overclamp_beats_sdil_all_periods": bool(np.all( overclamp_error < sdil_error)), "online_constant_neutral_observations_by_period": [ int(record["neutral_observations_during_learning"]) for record in by_method["online_constant"] ], "frozen_sdil_neutral_observations_by_period": [ int(record["neutral_observations_during_learning"]) for record in by_method["frozen_sdil"] ], }, } report["summary"] = { "sdil_beats_frozen_constant_both_pairs": bool(all( record["summary"]["sdil_beats_frozen_constant_all_periods"] for record in report["pairs"].values() )), "median_gap_closed_by_pair": { name: record["summary"]["median_raw_to_oracle_gap_closed_by_sdil"] for name, record 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()