summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 13:16:20 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 13:16:20 -0500
commit1d819cf81732b9b8a11fbf01c8b4f7d208e93b92 (patch)
treedea073b23d4e1967dfca4bc21bf73c5771ed44d8 /experiments
parent6ae1e40281375e9a9991e0201f32c76424a2399b (diff)
exp: add physical grid clean-learning gate
Diffstat (limited to 'experiments')
-rw-r--r--experiments/physical_grid_clean_gate.py145
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()