summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/physical_bias_p1_smoke.py54
-rw-r--r--sdil/physical_coupled.py53
2 files changed, 99 insertions, 8 deletions
diff --git a/experiments/physical_bias_p1_smoke.py b/experiments/physical_bias_p1_smoke.py
index dc9b6f1..c23c9eb 100644
--- a/experiments/physical_bias_p1_smoke.py
+++ b/experiments/physical_bias_p1_smoke.py
@@ -20,6 +20,8 @@ from sdil.physical_coupled import (
free_output,
joint_solution,
local_replay_update,
+ overclamped_clean_rate,
+ simulate_alternating_tasks,
standard_clean_rate,
)
@@ -73,12 +75,64 @@ def main() -> None:
downstream[:] = np.random.default_rng(29).normal(size=downstream.shape)
update_c = local_replay_update(affine, gates, teaching, eligibility, 0.03)
assert np.array_equal(update_a, update_c)
+
+ _, output_free, output_clamped = overclamped_clean_rate(
+ circuit,
+ joint,
+ tasks[0],
+ nudging=0.25,
+ clamp_magnitude=circuit.high,
+ exact_target=True,
+ )
+ expected_clamped = output_free + 0.25 * (
+ circuit.high * np.sign(tasks[0].label_voltage - output_free)
+ - output_free
+ )
+ assert output_clamped == expected_clamped
+
+ exact_predictor = LocalPredictor.zeros(
+ field.reference_gate, scale, affine=True)
+ exact_predictor.coefficients[:, 0] = field.bias_at_reference
+ exact_predictor.coefficients[:, 1] = field.local_slopes * scale
+ shared = {
+ "period_seconds": 0.01,
+ "cycles": 8,
+ "initial_gates": np.asarray((4.0, 4.0)),
+ "overclamp_nudging": 0.25,
+ "overclamp_magnitude": circuit.high,
+ "overclamp_exact_target": True,
+ }
+ clean_overclamp = simulate_alternating_tasks(
+ circuit,
+ tasks,
+ field,
+ method="overclamp",
+ bias_strength=0.0,
+ **shared,
+ )
+ debiased_overclamp = simulate_alternating_tasks(
+ circuit,
+ tasks,
+ field,
+ method="overclamp_sdil",
+ bias_strength=1.0,
+ predictor=exact_predictor,
+ **shared,
+ )
+ assert np.allclose(
+ clean_overclamp["final_gates"],
+ debiased_overclamp["final_gates"],
+ atol=1e-12,
+ rtol=0.0,
+ )
+ assert debiased_overclamp["max_abs_clamp_displacement_v"] > 0.0
print({
"joint_solution": joint.tolist(),
"affine_calibration_mae": float(affine_error),
"constant_calibration_mae": float(constant_error),
"neutral_observations_each": observations_affine,
"local_replay_exact": True,
+ "sdil_overclamp_composition_exact": True,
"autodiff_used": False,
})
diff --git a/sdil/physical_coupled.py b/sdil/physical_coupled.py
index f451947..5e5b181 100644
--- a/sdil/physical_coupled.py
+++ b/sdil/physical_coupled.py
@@ -198,12 +198,18 @@ def overclamped_clean_rate(
*,
nudging: float = 0.25,
clamp_magnitude: Optional[float] = None,
+ exact_target: bool = False,
) -> tuple[Array, float, float]:
- """Leading-order overclamping signal from Appendix F, Eq. F6--F8."""
+ """Overclamping signal from Appendix F, Eq. F5 or its Eq. F6 limit."""
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)
+ if exact_target:
+ output_clamped = output_free + nudging * (
+ magnitude * np.sign(error) - output_free
+ )
+ else:
+ 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)
@@ -269,6 +275,9 @@ def simulate_alternating_tasks(
seed: int = 0,
summary_cycles: int = 20,
record_history: bool = False,
+ overclamp_nudging: float = 0.25,
+ overclamp_magnitude: Optional[float] = None,
+ overclamp_exact_target: bool = False,
) -> dict:
"""Alternate two tasks using explicit local circuit updates.
@@ -282,6 +291,7 @@ def simulate_alternating_tasks(
allowed = {
"raw", "frozen_constant", "frozen_sdil", "online_constant",
"online_sdil", "oracle", "same_rms_noise", "overclamp",
+ "overclamp_sdil",
}
if method not in allowed:
raise ValueError(f"unrecognized method {method}")
@@ -291,7 +301,8 @@ def simulate_alternating_tasks(
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"
+ "frozen_constant", "frozen_sdil", "online_constant", "online_sdil",
+ "overclamp_sdil",
} and predictor is None:
raise ValueError(f"{method} requires a predictor")
if method == "same_rms_noise" and noise_standard_deviation is None:
@@ -318,6 +329,9 @@ def simulate_alternating_tasks(
learning_on_time = 0.0
neutral_observations = 0
clipped_updates = 0
+ clamp_displacement_l1_time = 0.0
+ clamp_displacement_l2_time = 0.0
+ max_abs_clamp_displacement = 0.0
for _ in range(cycles):
half_endpoints = []
@@ -330,14 +344,26 @@ def simulate_alternating_tasks(
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)
+ if method in {"overclamp", "overclamp_sdil"}:
+ clean_rate, output_free, output_clamped = overclamped_clean_rate(
+ circuit,
+ gates,
+ task,
+ nudging=overclamp_nudging,
+ clamp_magnitude=overclamp_magnitude,
+ exact_target=overclamp_exact_target,
+ )
duration = nominal_step * abs(
task.label_voltage - output_free) / initial_error_scale
- residual_bias = physical_bias
+ if method == "overclamp":
+ residual_bias = physical_bias
+ else:
+ residual_bias = (
+ physical_bias - active_predictor.predict(gates)
+ )
else:
- clean_rate, _, _ = standard_clean_rate(circuit, gates, task)
+ clean_rate, output_free, output_clamped = standard_clean_rate(
+ circuit, gates, task)
duration = nominal_step
if method == "raw":
residual_bias = physical_bias
@@ -359,6 +385,11 @@ def simulate_alternating_tasks(
clipped_updates += int(np.any(clipped != proposed))
gates = clipped
learning_on_time += duration
+ displacement = abs(output_clamped - output_free)
+ clamp_displacement_l1_time += duration * displacement
+ clamp_displacement_l2_time += duration * displacement * displacement
+ max_abs_clamp_displacement = max(
+ max_abs_clamp_displacement, displacement)
half_endpoints.append(gates.copy())
half_task_errors.append(task_errors(circuit, gates, tasks))
half_task_errors_array = np.asarray(half_task_errors)
@@ -385,6 +416,12 @@ def simulate_alternating_tasks(
"std_cycle_span": float(np.std(spans)),
"neutral_observations_during_learning": neutral_observations,
"learning_on_time_seconds": float(learning_on_time),
+ "clamp_displacement_l1_time_v_s": float(
+ clamp_displacement_l1_time),
+ "clamp_displacement_l2_time_v2_s": float(
+ clamp_displacement_l2_time),
+ "max_abs_clamp_displacement_v": float(
+ max_abs_clamp_displacement),
"clipped_updates": clipped_updates,
"final_predictor_coefficients": (
None if active_predictor is None