diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 17:26:42 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 17:26:42 -0500 |
| commit | 7d592f7e6c3c3144e6c347996b6c1f048fa7847b (patch) | |
| tree | 30bd2bb412d65f29483a76b4584dd90df35fc913 | |
| parent | ac5f34f5c839642d6ab0baf9e3f1b784fa905786 (diff) | |
feat: add sparse digital CLLN scaling ladder
| -rw-r--r-- | experiments/coupled_ladder_scaling.py | 341 | ||||
| -rw-r--r-- | sdil/coupled_ladder.py | 336 |
2 files changed, 677 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() + diff --git a/sdil/coupled_ladder.py b/sdil/coupled_ladder.py new file mode 100644 index 0000000..a9b8014 --- /dev/null +++ b/sdil/coupled_ladder.py @@ -0,0 +1,336 @@ +"""Sparse digital coupled-learning grids for controlled scaling experiments. + +The network is a linear resistor lattice. Its local learning signal is the +same free-minus-clamped voltage-square difference used by coupled learning. +Component imperfections reuse the per-edge measurement model used by the +hardware-realistic simulator in :mod:`sdil.physical_grid`. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np +from scipy.sparse import coo_matrix +from scipy.sparse.linalg import spsolve + +from sdil.physical_grid import ( + GridCircuit, + GridSquareLawImperfection, + RingClassificationDataset, + edge_voltage_drops, + output_difference, +) + + +Array = np.ndarray + + +def make_scaled_grid(side: int) -> GridCircuit: + """Return a square grid whose side-4 boundary layout matches Figure 5.""" + if side < 4 or side % 4: + raise ValueError("grid side must be a positive multiple of four") + + def node(row: int, column: int) -> int: + return row * side + column + + return GridCircuit( + rows=side, + columns=side, + source_nodes=( + node(3 * side // 4, 3 * side // 4), + node(3 * side // 4, side // 4), + node(side // 4, 3 * side // 4), + node(side // 4, side // 4), + ), + target_nodes=( + node(side // 2, side // 2), + node(side // 2, 0), + ), + ) + + +def tile_figure5_gates(base_gates: Array, side: int) -> Array: + """Tile a released 4-by-4 horizontal/vertical gate pattern.""" + gates = np.asarray(base_gates, dtype=float) + if gates.shape != (32,): + raise ValueError("the released base gate vector must have 32 entries") + if side < 4 or side % 4: + raise ValueError("grid side must be a positive multiple of four") + horizontal = gates[:16].reshape(4, 4) + vertical = gates[16:].reshape(4, 4) + tiled_horizontal = np.asarray([ + horizontal[row % 4, column % 4] + for row in range(side) + for column in range(side) + ]) + tiled_vertical = np.asarray([ + vertical[row % 4, column % 4] + for row in range(side) + for column in range(side) + ]) + return np.concatenate((tiled_horizontal, tiled_vertical)) + + +def solve_linear_grid_state( + circuit: GridCircuit, + gates: Array, + source_values: Array, + *, + target_values: Array | None = None, +) -> Array: + """Solve the sparse linear Kirchhoff system with fixed boundary nodes.""" + gates = np.asarray(gates, dtype=float) + sources = np.asarray(source_values, dtype=float) + if gates.shape != (circuit.edge_count,): + raise ValueError("gate vector has the wrong shape") + if sources.shape != (len(circuit.source_nodes),): + raise ValueError("source voltage vector has the wrong shape") + conductances = circuit.conductance_scale * ( + gates - circuit.threshold_voltage) + if np.any(conductances <= 0.0): + raise ValueError("all digital conductances must be positive") + + fixed = dict(zip(circuit.source_nodes, sources)) + if target_values is not None: + targets = np.asarray(target_values, dtype=float) + if targets.shape != (2,): + raise ValueError("target voltage vector must have shape (2,)") + fixed.update(zip(circuit.target_nodes, targets)) + + unknown_nodes = [ + node for node in range(circuit.node_count) if node not in fixed + ] + unknown_index = {node: index for index, node in enumerate(unknown_nodes)} + matrix_rows: list[int] = [] + matrix_columns: list[int] = [] + matrix_values: list[float] = [] + rhs = np.zeros(len(unknown_nodes), dtype=float) + + for conductance, (first, second) in zip( + conductances, circuit.edge_pairs + ): + for node, neighbour in ((first, second), (second, first)): + if node in fixed: + continue + row = unknown_index[node] + matrix_rows.append(row) + matrix_columns.append(row) + matrix_values.append(float(conductance)) + if neighbour in fixed: + rhs[row] += conductance * fixed[neighbour] + else: + matrix_rows.append(row) + matrix_columns.append(unknown_index[neighbour]) + matrix_values.append(float(-conductance)) + + matrix = coo_matrix( + (matrix_values, (matrix_rows, matrix_columns)), + shape=(len(unknown_nodes), len(unknown_nodes)), + ).tocsr() + unknown_voltages = np.asarray(spsolve(matrix, rhs), dtype=float) + if not np.all(np.isfinite(unknown_voltages)): + raise RuntimeError("linear grid solve returned a nonfinite state") + + voltages = np.empty(circuit.node_count, dtype=float) + for node, voltage in fixed.items(): + voltages[node] = voltage + voltages[unknown_nodes] = unknown_voltages + return voltages + + +def evaluate_digital_grid( + circuit: GridCircuit, + gates: Array, + dataset: RingClassificationDataset, +) -> dict: + outputs = [] + for inputs in dataset.inputs_v: + state = solve_linear_grid_state( + circuit, gates, circuit.source_values(*inputs)) + outputs.append(output_difference(circuit, state)) + outputs_array = np.asarray(outputs) + labels = np.asarray(dataset.labels_v) + errors = labels - outputs_array + active = labels * errors > 0.0 + return { + "classification_error": float(np.mean( + np.sign(outputs_array) != np.sign(labels))), + "hinge_loss_v2": float(np.mean(np.where( + active, np.square(errors), 0.0))), + "margin_success_fraction": float(np.mean(~active)), + "outputs_v": outputs_array.tolist(), + } + + +@dataclass(frozen=True) +class DigitalTrainingConfig: + epochs: int = 600 + record_every: int = 10 + standard_nudging: float = 128.0 / 129.0 + learning_time_seconds: float = 1.0e-3 + overclamp_nudging: float = 32.0 / 129.0 + overclamp_time_seconds_per_v: float = 2.5e-3 + overclamp_target_magnitude_v: float | None = None + + def __post_init__(self) -> None: + if self.epochs < 1 or self.record_every < 1: + raise ValueError("epochs and record interval must be positive") + if self.learning_time_seconds <= 0.0: + raise ValueError("learning time must be positive") + + +def train_digital_grid( + circuit: GridCircuit, + initial_gates: Array, + dataset: RingClassificationDataset, + imperfection: GridSquareLawImperfection, + *, + method: str, + config: DigitalTrainingConfig, + constant_bias_v_per_s: Array | None = None, + noise_seed: int = 0, +) -> dict: + """Train with a hand-written local coupled-learning update.""" + allowed = { + "clean", + "matched_noise", + "raw", + "constant", + "sdil", + "overclamp", + "overclamp_sdil", + } + if method not in allowed: + raise ValueError(f"unrecognized method {method}") + if method == "constant" and constant_bias_v_per_s is None: + raise ValueError("constant calibration requires a per-edge baseline") + + gates = np.asarray(initial_gates, dtype=float).copy() + if gates.shape != (circuit.edge_count,): + raise ValueError("initial gate vector has the wrong shape") + if constant_bias_v_per_s is not None: + constant_bias = np.asarray(constant_bias_v_per_s, dtype=float) + if constant_bias.shape != gates.shape: + raise ValueError("constant baseline has the wrong shape") + else: + constant_bias = None + + rng = np.random.default_rng(noise_seed) + local_updates = 0 + neutral_observations = 0 + clipped_updates = 0 + cumulative_learning_time = 0.0 + trace = [{ + "epoch": 0, + "local_updates": 0, + **evaluate_digital_grid(circuit, gates, dataset), + }] + + for epoch in range(1, config.epochs + 1): + for inputs, label in zip(dataset.inputs_v, dataset.labels_v): + sources = circuit.source_values(*inputs) + free_state = solve_linear_grid_state(circuit, gates, sources) + output_free = output_difference(circuit, free_state) + error = label - output_free + if label * error <= 0.0: + continue + + free_drops = edge_voltage_drops(circuit, free_state) + is_overclamp = method.startswith("overclamp") + if is_overclamp: + target_magnitude = ( + circuit.high_voltage + if config.overclamp_target_magnitude_v is None + else config.overclamp_target_magnitude_v + ) + output_clamped = output_free + config.overclamp_nudging * ( + target_magnitude * np.sign(error) - output_free) + duration = ( + config.overclamp_time_seconds_per_v * abs(error)) + else: + output_clamped = output_free + ( + config.standard_nudging * error) + duration = config.learning_time_seconds + + target_mean = float(np.mean( + free_state[list(circuit.target_nodes)])) + target_values = np.asarray(( + target_mean + 0.5 * output_clamped, + target_mean - 0.5 * output_clamped, + )) + clamped_state = solve_linear_grid_state( + circuit, gates, sources, target_values=target_values) + clamped_drops = edge_voltage_drops(circuit, clamped_state) + ideal_rate = imperfection.ideal_rate( + circuit.measured_learning_rate, free_drops, clamped_drops) + observed_rate = imperfection.observed_rate( + circuit.measured_learning_rate, free_drops, clamped_drops) + + if method == "clean": + applied_rate = ideal_rate + elif method == "matched_noise": + measurement_error = observed_rate - ideal_rate + random_sign = rng.choice((-1.0, 1.0), size=len(gates)) + applied_rate = ideal_rate + random_sign * np.abs( + measurement_error) + elif method == "constant": + applied_rate = observed_rate - constant_bias + elif method in {"sdil", "overclamp_sdil"}: + neutral_rate = imperfection.neutral_bias( + circuit.measured_learning_rate, free_drops) + applied_rate = observed_rate - neutral_rate + neutral_observations += 1 + else: + applied_rate = observed_rate + + proposed = gates + duration * applied_rate + clipped = np.clip( + proposed, circuit.gate_minimum, circuit.gate_maximum) + clipped_updates += int(np.any(clipped != proposed)) + gates = clipped + local_updates += 1 + cumulative_learning_time += duration + + if epoch % config.record_every == 0 or epoch == config.epochs: + trace.append({ + "epoch": epoch, + "local_updates": local_updates, + **evaluate_digital_grid(circuit, gates, dataset), + }) + + zero_records = [ + record for record in trace + if record["classification_error"] == 0.0 + ] + if zero_records: + epochs_to_zero = int(zero_records[0]["epoch"]) + updates_to_zero = int(zero_records[0]["local_updates"]) + reached_zero = True + else: + epochs_to_zero = config.epochs + updates_to_zero = local_updates + reached_zero = False + epoch_axis = np.asarray([record["epoch"] for record in trace]) + error_axis = np.asarray([ + record["classification_error"] for record in trace]) + error_auc = float(np.trapezoid(error_axis, epoch_axis) / config.epochs) + final = trace[-1] + return { + "method": method, + "classification_error": final["classification_error"], + "hinge_loss_v2": final["hinge_loss_v2"], + "margin_success_fraction": final["margin_success_fraction"], + "outputs_v": final["outputs_v"], + "reached_zero_error": reached_zero, + "restricted_epochs_to_zero_error": epochs_to_zero, + "restricted_updates_to_zero_error": updates_to_zero, + "classification_error_auc": error_auc, + "local_updates": local_updates, + "neutral_observations": neutral_observations, + "cumulative_learning_time_seconds": float(cumulative_learning_time), + "clipped_updates": clipped_updates, + "final_gates_v": gates.tolist(), + "trace": trace, + } + |
