"""Backpropagation-free simulator for the Dillavou two-edge circuit. The circuit equations and standard/overclamping updates follow Appendix D/F of Dillavou et al. (arXiv:2505.22887v2). Every adaptive operation is an explicit NumPy local rule; this module intentionally has no autodiff path. """ from __future__ import annotations from dataclasses import dataclass from typing import Callable, Iterable, Optional import numpy as np Array = np.ndarray @dataclass(frozen=True) class Circuit: high: float = 0.4351 low: float = 0.0181 conductance_per_gate: float = 8.5e-4 threshold_voltage: float = 0.7 fixed_conductance: float = 1.0 / 500.0 measured_learning_rate: float = 2040.0 integration_step_seconds: float = 2.0e-4 gate_minimum: float = 1.0 gate_maximum: float = 5.2 @dataclass(frozen=True) class Task: name: str input_voltage: float label_voltage: float @dataclass(frozen=True) class LocalAffineBias: reference_gate: Array bias_at_reference: Array local_slopes: Array def __post_init__(self) -> None: for value in ( self.reference_gate, self.bias_at_reference, self.local_slopes ): if np.asarray(value).shape != (2,): raise ValueError("two-edge bias arrays must have shape (2,)") def __call__(self, gates: Array, strength: float = 1.0) -> Array: gates = np.asarray(gates, dtype=float) if gates.shape != (2,): raise ValueError("gates must have shape (2,)") return strength * ( self.bias_at_reference + self.local_slopes * (gates - self.reference_gate) ) @dataclass class LocalPredictor: """Independent per-edge affine filters trained by normalized LMS.""" reference_gate: Array feature_scale: Array coefficients: Array affine: bool @classmethod def zeros( cls, reference_gate: Array, feature_scale: Array, *, affine: bool, ) -> "LocalPredictor": width = 2 if affine else 1 return cls( reference_gate=np.asarray(reference_gate, dtype=float).copy(), feature_scale=np.asarray(feature_scale, dtype=float).copy(), coefficients=np.zeros((2, width), dtype=float), affine=affine, ) def features(self, gates: Array) -> Array: gates = np.asarray(gates, dtype=float) if gates.shape != (2,): raise ValueError("gates must have shape (2,)") if not self.affine: return np.ones((2, 1), dtype=float) normalized = (gates - self.reference_gate) / self.feature_scale return np.column_stack((np.ones(2, dtype=float), normalized)) def predict(self, gates: Array) -> Array: return np.sum(self.coefficients * self.features(gates), axis=1) def update(self, gates: Array, neutral_measurement: Array, rate: float) -> Array: """One local normalized-LMS update and its pre-update residual.""" neutral_measurement = np.asarray(neutral_measurement, dtype=float) if neutral_measurement.shape != (2,): raise ValueError("neutral measurement must have shape (2,)") features = self.features(gates) residual = neutral_measurement - np.sum( self.coefficients * features, axis=1) normalization = np.sum(features * features, axis=1, keepdims=True) self.coefficients += ( rate * residual[:, None] * features / np.maximum(normalization, 1e-12) ) return residual def copy(self) -> "LocalPredictor": return LocalPredictor( reference_gate=self.reference_gate.copy(), feature_scale=self.feature_scale.copy(), coefficients=self.coefficients.copy(), affine=self.affine, ) def free_output(circuit: Circuit, gates: Array, input_voltage: float) -> float: """Appendix D, Eq. D13, with gates ordered (minus, plus).""" gate_minus, gate_plus = np.asarray(gates, dtype=float) scale = circuit.conductance_per_gate numerator = ( input_voltage * circuit.fixed_conductance + scale * ( gate_plus * circuit.high + gate_minus * circuit.low - (circuit.low + circuit.high) * circuit.threshold_voltage ) ) denominator = ( circuit.fixed_conductance + scale * ( gate_plus + gate_minus - 2.0 * circuit.threshold_voltage ) ) if denominator <= 0.0: raise ValueError("nonpositive effective conductance") return float(numerator / denominator) def solution_line(circuit: Circuit, task: Task) -> tuple[float, float]: """Return slope/intercept of gate_plus versus gate_minus at zero error.""" label = task.label_voltage scale = circuit.conductance_per_gate denominator = scale * (circuit.high - label) if denominator == 0.0: raise ValueError("label coincides with high boundary") slope = -(circuit.low - label) / (circuit.high - label) intercept = -( circuit.fixed_conductance * (task.input_voltage - label) + scale * circuit.threshold_voltage * (2.0 * label - circuit.low - circuit.high) ) / denominator return float(slope), float(intercept) def joint_solution(circuit: Circuit, tasks: Iterable[Task]) -> Array: tasks = tuple(tasks) if len(tasks) != 2: raise ValueError("joint_solution expects exactly two tasks") slope_a, intercept_a = solution_line(circuit, tasks[0]) slope_b, intercept_b = solution_line(circuit, tasks[1]) if slope_a == slope_b: raise ValueError("parallel solution lines have no unique joint solution") gate_minus = (intercept_b - intercept_a) / (slope_a - slope_b) return np.asarray( (gate_minus, slope_a * gate_minus + intercept_a), dtype=float) def voltage_drop_squares(circuit: Circuit, output: float) -> Array: return np.asarray( ((output - circuit.low) ** 2, (circuit.high - output) ** 2), dtype=float, ) def standard_clean_rate( circuit: Circuit, gates: Array, task: Task, nudging: float = 1.0 ) -> tuple[Array, float, float]: output_free = free_output(circuit, gates, task.input_voltage) output_clamped = output_free + nudging * ( task.label_voltage - output_free) rate = circuit.measured_learning_rate * ( voltage_drop_squares(circuit, output_free) - voltage_drop_squares(circuit, output_clamped) ) return rate, output_free, output_clamped def overclamped_clean_rate( circuit: Circuit, gates: Array, task: Task, *, nudging: float = 0.25, clamp_magnitude: Optional[float] = None, ) -> tuple[Array, float, float]: """Leading-order overclamping signal from Appendix F, Eq. F6--F8.""" output_free = free_output(circuit, gates, task.input_voltage) error = task.label_voltage - output_free magnitude = circuit.high if clamp_magnitude is None else clamp_magnitude output_clamped = output_free + nudging * magnitude * np.sign(error) rate = circuit.measured_learning_rate * ( voltage_drop_squares(circuit, output_free) - voltage_drop_squares(circuit, output_clamped) ) return rate, output_free, output_clamped def calibrate_predictor( predictor: LocalPredictor, states: Array, measurement: Callable[[Array], Array], *, epochs: int, learning_rate: float, ) -> int: """Sequential local LMS calibration; returns neutral observation count.""" states = np.asarray(states, dtype=float) if states.ndim != 2 or states.shape[1] != 2: raise ValueError("calibration states must have shape (observations, 2)") if epochs < 1: raise ValueError("epochs must be positive") count = 0 for _ in range(epochs): for gates in states: predictor.update(gates, measurement(gates), learning_rate) count += 1 return count def local_replay_update( predictor: LocalPredictor, gates: Array, teaching_measurement: Array, eligibility: Array, learning_rate: float, ) -> Array: """The complete stored-tuple SDIL update, independent of any task/model.""" residual = np.asarray(teaching_measurement) - predictor.predict(gates) return learning_rate * residual * np.asarray(eligibility) def task_errors(circuit: Circuit, gates: Array, tasks: Iterable[Task]) -> Array: return np.asarray([ (task.label_voltage - free_output( circuit, gates, task.input_voltage)) ** 2 for task in tasks ], dtype=float) def simulate_alternating_tasks( circuit: Circuit, tasks: Iterable[Task], bias_field: LocalAffineBias, *, method: str, period_seconds: float, cycles: int, initial_gates: Array, bias_strength: float = 1.0, predictor: Optional[LocalPredictor] = None, online_predictor_rate: float = 0.05, noise_standard_deviation: Optional[Array] = None, seed: int = 0, summary_cycles: int = 20, record_history: bool = False, ) -> dict: """Alternate two tasks using explicit local circuit updates. `online_constant` and `online_sdil` take one neutral observation at the beginning of each half-cycle. Frozen predictors take none during task learning. The overclamping implementation uses the leading-order constant-displacement signal of Eq. F6 and the error-proportional update duration of Eq. F8; it is an analogue for these regression tasks, not a reproduction of the paper's classification experiment. """ allowed = { "raw", "frozen_constant", "frozen_sdil", "online_constant", "online_sdil", "oracle", "same_rms_noise", "overclamp", } if method not in allowed: raise ValueError(f"unrecognized method {method}") tasks = tuple(tasks) if len(tasks) != 2: raise ValueError("exactly two alternating tasks are required") if period_seconds <= 0.0 or cycles < 1: raise ValueError("period and cycles must be positive") if method in { "frozen_constant", "frozen_sdil", "online_constant", "online_sdil" } and predictor is None: raise ValueError(f"{method} requires a predictor") if method == "same_rms_noise" and noise_standard_deviation is None: raise ValueError("same_rms_noise requires a standard deviation") active_predictor = predictor.copy() if predictor is not None else None gates = np.asarray(initial_gates, dtype=float).copy() if gates.shape != (2,): raise ValueError("initial gates must have shape (2,)") nominal_step = circuit.integration_step_seconds half_steps = max(1, int(round(period_seconds / (2.0 * nominal_step)))) rng = np.random.default_rng(seed) initial_error_scale = float(np.mean([ abs(task.label_voltage - free_output( circuit, gates, task.input_voltage)) for task in tasks ])) initial_error_scale = max(initial_error_scale, 1e-6) combined_error_history = [] cycle_span_history = [] gate_history = [] task_error_history = [] learning_on_time = 0.0 neutral_observations = 0 clipped_updates = 0 for _ in range(cycles): half_endpoints = [] half_task_errors = [] for task in tasks: if method in {"online_constant", "online_sdil"}: neutral = bias_field(gates, bias_strength) active_predictor.update( gates, neutral, online_predictor_rate) neutral_observations += 1 for _ in range(half_steps): physical_bias = bias_field(gates, bias_strength) if method == "overclamp": clean_rate, output_free, _ = overclamped_clean_rate( circuit, gates, task) duration = nominal_step * abs( task.label_voltage - output_free) / initial_error_scale residual_bias = physical_bias else: clean_rate, _, _ = standard_clean_rate(circuit, gates, task) duration = nominal_step if method == "raw": residual_bias = physical_bias elif method == "oracle": residual_bias = np.zeros(2, dtype=float) elif method == "same_rms_noise": residual_bias = rng.normal( loc=0.0, scale=np.asarray(noise_standard_deviation, dtype=float), size=2, ) else: residual_bias = ( physical_bias - active_predictor.predict(gates) ) proposed = gates + duration * (clean_rate + residual_bias) clipped = np.clip( proposed, circuit.gate_minimum, circuit.gate_maximum) clipped_updates += int(np.any(clipped != proposed)) gates = clipped learning_on_time += duration half_endpoints.append(gates.copy()) half_task_errors.append(task_errors(circuit, gates, tasks)) half_task_errors_array = np.asarray(half_task_errors) combined_error_history.append(float(np.mean(half_task_errors_array))) cycle_span_history.append(float(np.linalg.norm( half_endpoints[1] - half_endpoints[0]))) gate_history.append(np.asarray(half_endpoints).tolist()) task_error_history.append(half_task_errors_array.tolist()) summary_count = min(summary_cycles, cycles) combined = np.asarray(combined_error_history[-summary_count:]) spans = np.asarray(cycle_span_history[-summary_count:]) result = { "method": method, "period_seconds": period_seconds, "cycles": cycles, "half_steps": half_steps, "initial_gates": np.asarray(initial_gates, dtype=float).tolist(), "final_gates": gates.tolist(), "bias_strength": bias_strength, "mean_combined_error": float(np.mean(combined)), "std_combined_error": float(np.std(combined)), "mean_cycle_span": float(np.mean(spans)), "std_cycle_span": float(np.std(spans)), "neutral_observations_during_learning": neutral_observations, "learning_on_time_seconds": float(learning_on_time), "clipped_updates": clipped_updates, "final_predictor_coefficients": ( None if active_predictor is None else active_predictor.coefficients.tolist() ), } if record_history: result.update({ "combined_error_history": combined_error_history, "cycle_span_history": cycle_span_history, "half_cycle_gate_history": gate_history, "half_cycle_task_error_history": task_error_history, }) return result