summaryrefslogtreecommitdiff
path: root/experiments/coupled_ladder_scaling.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/coupled_ladder_scaling.py')
-rw-r--r--experiments/coupled_ladder_scaling.py341
1 files changed, 341 insertions, 0 deletions
diff --git a/experiments/coupled_ladder_scaling.py b/experiments/coupled_ladder_scaling.py
new file mode 100644
index 0000000..ce3509a
--- /dev/null
+++ b/experiments/coupled_ladder_scaling.py
@@ -0,0 +1,341 @@
+#!/usr/bin/env python3
+"""Run the paired digital CLLN size ladder under component imperfection."""
+
+from __future__ import annotations
+
+import argparse
+from concurrent.futures import ProcessPoolExecutor, as_completed
+import json
+from pathlib import Path
+import sys
+import time
+
+import numpy as np
+
+ROOT = Path(__file__).resolve().parents[1]
+sys.path.insert(0, str(ROOT))
+
+from sdil.coupled_ladder import ( # noqa: E402
+ DigitalTrainingConfig,
+ make_scaled_grid,
+ solve_linear_grid_state,
+ tile_figure5_gates,
+ train_digital_grid,
+)
+from sdil.physical_grid import ( # noqa: E402
+ GridSquareLawImperfection,
+ RingClassificationDataset,
+ edge_voltage_drops,
+)
+
+
+METHODS = (
+ "clean",
+ "matched_noise",
+ "raw",
+ "constant",
+ "sdil",
+ "overclamp",
+ "overclamp_sdil",
+)
+
+
+def select_tasks(protocol: dict, rotations: int) -> list[dict]:
+ if rotations < 1 or rotations > 8:
+ raise ValueError("rotations must be between one and eight")
+ standard = [
+ record for record in protocol["experiments"]
+ if record["method"] == "standard"
+ ]
+ selected = []
+ for diameter in protocol["protocol_checks"]["input_diameters_v"]:
+ 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_mapping(specification: str, cast) -> dict:
+ mapping = {}
+ for entry in specification.split(","):
+ key, value = entry.split(":", maxsplit=1)
+ mapping[int(key)] = cast(value)
+ return mapping
+
+
+def constant_calibration(
+ circuit,
+ gates: np.ndarray,
+ imperfection: GridSquareLawImperfection,
+ *,
+ observations: int,
+ seed: int,
+) -> np.ndarray:
+ rng = np.random.default_rng(seed)
+ measurements = []
+ for _ in range(observations):
+ inputs = rng.uniform(
+ circuit.low_voltage, circuit.high_voltage, size=2)
+ state = solve_linear_grid_state(
+ circuit, gates, circuit.source_values(*inputs))
+ drops = edge_voltage_drops(circuit, state)
+ measurements.append(imperfection.neutral_bias(
+ circuit.measured_learning_rate, drops))
+ return np.mean(np.asarray(measurements), axis=0)
+
+
+def compact_result(result: dict) -> dict:
+ result = dict(result)
+ result.pop("final_gates_v", None)
+ result["trace"] = [{
+ key: value for key, value in record.items()
+ if key != "outputs_v"
+ } for record in result["trace"]]
+ return result
+
+
+def run_job(job: dict) -> dict:
+ side = job["side"]
+ task = job["task"]
+ task_index = job["task_index"]
+ device_seed = job["device_seed"]
+ circuit = make_scaled_grid(side)
+ gates = tile_figure5_gates(
+ np.asarray(task["initial_gates_v"], dtype=float), side)
+ 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,
+ )
+ imperfection = GridSquareLawImperfection.sample_appendix_c(
+ circuit.edge_count,
+ device_seed,
+ gain_standard_deviation=job["gain_standard_deviation"],
+ twin_mismatch_standard_deviation_v=(
+ job["twin_mismatch_standard_deviation_v"]),
+ multiplier_offset_standard_deviation_v_per_s=(
+ job["multiplier_offset_standard_deviation_v_per_s"]),
+ )
+ constant_bias = None
+ if "constant" in job["methods"]:
+ constant_bias = constant_calibration(
+ circuit,
+ gates,
+ imperfection,
+ observations=job["calibration_observations"],
+ seed=device_seed + 1_000_003 + task_index + 1009 * side,
+ )
+ config = DigitalTrainingConfig(
+ epochs=job["epochs"],
+ record_every=job["record_every"],
+ learning_time_seconds=job["learning_time_seconds"],
+ )
+ methods = {}
+ for method in job["methods"]:
+ start = time.perf_counter()
+ try:
+ result = train_digital_grid(
+ circuit,
+ gates,
+ dataset,
+ imperfection,
+ method=method,
+ config=config,
+ constant_bias_v_per_s=constant_bias,
+ noise_seed=(
+ device_seed + 2_000_003 + task_index + 1009 * side),
+ )
+ result = compact_result(result)
+ result["status"] = "completed"
+ except (RuntimeError, ValueError) as error:
+ result = {
+ "method": method,
+ "status": "failed",
+ "failure_message": str(error),
+ "classification_error": 1.0,
+ "hinge_loss_v2": None,
+ "reached_zero_error": False,
+ "restricted_epochs_to_zero_error": job["epochs"],
+ "restricted_updates_to_zero_error": (
+ job["epochs"] * len(dataset.labels_v)),
+ "classification_error_auc": 1.0,
+ "local_updates": None,
+ "neutral_observations": None,
+ "trace": [],
+ }
+ result["wall_seconds"] = float(time.perf_counter() - start)
+ methods[method] = result
+ return {
+ "side": side,
+ "nodes": circuit.node_count,
+ "learnable_edges": circuit.edge_count,
+ "task_index": task_index,
+ "input_diameter_v": task["input_diameter_v"],
+ "classes": task["classes"],
+ "device_seed": device_seed,
+ "methods": methods,
+ }
+
+
+def summarize(records: list[dict], methods: tuple[str, ...]) -> dict:
+ summary = {}
+ for side in sorted({record["side"] for record in records}):
+ side_records = [record for record in records if record["side"] == side]
+ method_summary = {}
+ for method in methods:
+ values = [record["methods"][method] for record in side_records]
+ errors = np.asarray([
+ value["classification_error"] for value in values])
+ method_summary[method] = {
+ "trials": len(values),
+ "failures": int(sum(
+ value["status"] != "completed" for value in values)),
+ "mean_classification_error": float(np.mean(errors)),
+ "zero_error_fraction": float(np.mean(errors == 0.0)),
+ "mean_classification_error_auc": float(np.mean([
+ value["classification_error_auc"] for value in values
+ ])),
+ "mean_restricted_epochs_to_zero_error": float(np.mean([
+ value["restricted_epochs_to_zero_error"]
+ for value in values
+ ])),
+ "median_wall_seconds": float(np.median([
+ value["wall_seconds"] for value in values
+ ])),
+ }
+ summary[str(side)] = {
+ "nodes": side_records[0]["nodes"],
+ "learnable_edges": side_records[0]["learnable_edges"],
+ "methods": method_summary,
+ }
+ return summary
+
+
+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/coupled_ladder/p0_pilot.json"))
+ parser.add_argument("--sizes", default="4,8,12,16,24,32")
+ parser.add_argument("--rotations", type=int, default=1)
+ parser.add_argument("--device-seeds", default="20260829")
+ parser.add_argument(
+ "--methods", default="clean,matched_noise,raw,sdil")
+ parser.add_argument("--epochs", type=int, default=600)
+ parser.add_argument("--record-every", type=int, default=10)
+ parser.add_argument(
+ "--learning-times",
+ default=(
+ "4:0.001,8:0.001,12:0.001,16:0.001,"
+ "24:0.001,32:0.001"),
+ )
+ parser.add_argument("--calibration-observations", type=int, default=16)
+ parser.add_argument("--gain-standard-deviation", type=float, default=0.01)
+ parser.add_argument(
+ "--twin-mismatch-standard-deviation-v", type=float, default=0.001)
+ parser.add_argument(
+ "--multiplier-offset-standard-deviation-v-per-s",
+ type=float,
+ default=2.3,
+ )
+ parser.add_argument("--workers", type=int, default=8)
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ protocol = json.loads(args.protocol.read_text())
+ sizes = tuple(int(value) for value in args.sizes.split(","))
+ learning_times = parse_mapping(args.learning_times, float)
+ missing_times = set(sizes) - set(learning_times)
+ if missing_times:
+ raise ValueError(
+ f"learning times are missing sizes {sorted(missing_times)}")
+ methods = tuple(args.methods.split(","))
+ unknown_methods = set(methods) - set(METHODS)
+ if unknown_methods:
+ raise ValueError(f"unknown methods: {sorted(unknown_methods)}")
+ tasks = select_tasks(protocol, args.rotations)
+ device_seeds = tuple(
+ int(value) for value in args.device_seeds.split(","))
+ jobs = [{
+ "side": side,
+ "task_index": task_index,
+ "task": task,
+ "device_seed": device_seed,
+ "methods": methods,
+ "epochs": args.epochs,
+ "record_every": args.record_every,
+ "learning_time_seconds": learning_times[side],
+ "calibration_observations": args.calibration_observations,
+ "gain_standard_deviation": args.gain_standard_deviation,
+ "twin_mismatch_standard_deviation_v": (
+ args.twin_mismatch_standard_deviation_v),
+ "multiplier_offset_standard_deviation_v_per_s": (
+ args.multiplier_offset_standard_deviation_v_per_s),
+ } for side in sizes
+ for task_index, task in enumerate(tasks)
+ for device_seed in device_seeds]
+
+ records = []
+ with ProcessPoolExecutor(max_workers=args.workers) as executor:
+ futures = [executor.submit(run_job, job) for job in jobs]
+ for completed, future in enumerate(as_completed(futures), start=1):
+ record = future.result()
+ records.append(record)
+ compact = ", ".join(
+ f"{method}={record['methods'][method]['classification_error']:.3f}"
+ for method in methods)
+ print(
+ f"{completed}/{len(jobs)} side={record['side']} "
+ f"task={record['task_index']} seed={record['device_seed']}: "
+ f"{compact}",
+ flush=True,
+ )
+ records.sort(key=lambda record: (
+ record["side"], record["task_index"], record["device_seed"]))
+ report = {
+ "analysis": "digital_coupled_learning_size_ladder",
+ "confirmatory": False,
+ "autodiff_used": False,
+ "source_protocol": str(args.protocol),
+ "protocol": {
+ "sizes": sizes,
+ "rotations_per_input_diameter": args.rotations,
+ "task_count": len(tasks),
+ "device_seeds": device_seeds,
+ "methods": methods,
+ "epochs": args.epochs,
+ "record_every": args.record_every,
+ "learning_time_seconds_by_side": learning_times,
+ "calibration_observations": args.calibration_observations,
+ "component_imperfection": {
+ "gain_standard_deviation": args.gain_standard_deviation,
+ "twin_mismatch_standard_deviation_v": (
+ args.twin_mismatch_standard_deviation_v),
+ "multiplier_offset_standard_deviation_v_per_s": (
+ args.multiplier_offset_standard_deviation_v_per_s),
+ },
+ "pairing": "task, initial gates, and component draw",
+ },
+ "records": records,
+ "summary": summarize(records, methods),
+ }
+ 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()
+