From 7d592f7e6c3c3144e6c347996b6c1f048fa7847b Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Sat, 29 Aug 2026 17:26:42 -0500 Subject: feat: add sparse digital CLLN scaling ladder --- experiments/coupled_ladder_scaling.py | 341 ++++++++++++++++++++++++++++++++++ 1 file changed, 341 insertions(+) create mode 100644 experiments/coupled_ladder_scaling.py (limited to 'experiments') 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() + -- cgit v1.2.3