#!/usr/bin/env python3 """Clean-learning gate for the reconstructed Figure-5 physical grid.""" from __future__ import annotations import argparse import json from pathlib import Path import sys import numpy as np ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from sdil.physical_grid import ( # noqa: E402 GridCircuit, GridSquareLawImperfection, RingClassificationDataset, train_grid_classifier, ) 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_diameters = (diameters[0], diameters[len(diameters) // 2], diameters[-1]) selected = [] for diameter in selected_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 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/p4_grid_clean_gate.json")) parser.add_argument("--rotations", type=int, default=2) parser.add_argument("--standard-epochs", type=int, default=600) parser.add_argument("--overclamp-epochs", type=int, default=1000) return parser.parse_args() def main() -> None: args = parse_args() protocol = json.loads(args.protocol.read_text()) circuit = GridCircuit() ideal = GridSquareLawImperfection.ideal(circuit.edge_count) tasks = select_tasks(protocol, args.rotations) records = [] for task_index, task in enumerate(tasks): 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) standard = train_grid_classifier( circuit, initial_gates, dataset, ideal, method="clean", epochs=args.standard_epochs, standard_learning_time_seconds=1e-3, record_every=50, ) overclamp = train_grid_classifier( circuit, initial_gates, dataset, ideal, method="overclamp_clean", epochs=args.overclamp_epochs, overclamp_time_seconds_per_v=0.0025, record_every=10, early_stop_perfect_checkpoints=3, ) records.append({ "task_index": task_index, "source_file": task["source_file"], "input_diameter_v": task["input_diameter_v"], "classes": task["classes"], "standard_clean": standard, "overclamp_clean": overclamp, }) print( f"task {task_index + 1}/{len(tasks)}: " f"standard={standard['classification_error']:.3f}, " f"overclamp={overclamp['classification_error']:.3f}", flush=True, ) report = { "analysis": "reconstructed_figure5_grid_clean_gate_p4", "confirmatory": False, "autodiff_used": False, "source_protocol": str(args.protocol), "protocol": { "selected_diameters_v": sorted({ record["input_diameter_v"] for record in records}), "rotations_per_diameter": args.rotations, "standard_epochs": args.standard_epochs, "standard_learning_time_seconds": 1e-3, "overclamp_epochs_maximum": args.overclamp_epochs, "overclamp_time_seconds_per_v": 0.0025, "overclamp_early_stop_perfect_checkpoints": 3, }, "records": records, "summary": { "tasks": len(records), "standard_clean_zero_error_tasks": sum( record["standard_clean"]["classification_error"] == 0.0 for record in records), "overclamp_clean_zero_error_tasks": sum( record["overclamp_clean"]["classification_error"] == 0.0 for record in records), }, } report["summary"]["gate_passed"] = bool( report["summary"]["standard_clean_zero_error_tasks"] == len(records) and report["summary"]["overclamp_clean_zero_error_tasks"] == len(records) ) 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()