#!/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()