summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/coupled_ladder_scaling.py341
-rw-r--r--sdil/coupled_ladder.py336
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,
+ }
+