diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 13:16:20 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 13:16:20 -0500 |
| commit | 1d819cf81732b9b8a11fbf01c8b4f7d208e93b92 (patch) | |
| tree | dea073b23d4e1967dfca4bc21bf73c5771ed44d8 | |
| parent | 6ae1e40281375e9a9991e0201f32c76424a2399b (diff) | |
exp: add physical grid clean-learning gate
| -rw-r--r-- | experiments/physical_grid_clean_gate.py | 145 |
1 files changed, 145 insertions, 0 deletions
diff --git a/experiments/physical_grid_clean_gate.py b/experiments/physical_grid_clean_gate.py new file mode 100644 index 0000000..2833f00 --- /dev/null +++ b/experiments/physical_grid_clean_gate.py @@ -0,0 +1,145 @@ +#!/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() |
