summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 13:13:52 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 13:13:52 -0500
commitb5c1b5be1628664977fb86bd7456b4176b204320 (patch)
tree1b8917552e69282de0d9beb255087301b64e3405
parentb9742072e2d9a80808fea1f314309f31476b80c8 (diff)
feat: train the reconstructed physical grid
-rw-r--r--experiments/physical_grid_smoke.py31
-rw-r--r--sdil/physical_grid.py200
2 files changed, 231 insertions, 0 deletions
diff --git a/experiments/physical_grid_smoke.py b/experiments/physical_grid_smoke.py
index a84ed84..d5ababd 100644
--- a/experiments/physical_grid_smoke.py
+++ b/experiments/physical_grid_smoke.py
@@ -15,11 +15,14 @@ from sdil.physical_grid import ( # noqa: E402
EdgePolynomialPredictor,
GridCircuit,
GridSquareLawImperfection,
+ RingClassificationDataset,
_residual_and_jacobian,
edge_voltage_drops,
fit_edge_predictor,
output_difference,
solve_grid_state,
+ evaluate_grid_classifier,
+ train_grid_classifier,
)
@@ -106,6 +109,32 @@ def main() -> None:
])))
assert quadratic_rmse < 1e-7
assert constant_rmse > 1e-3
+
+ angles = np.arange(8) * (2.0 * np.pi / 8.0)
+ midpoint = 0.5 * (circuit.low_voltage + circuit.high_voltage)
+ ring = RingClassificationDataset(
+ inputs_v=np.column_stack((
+ midpoint + 0.18 * np.cos(angles),
+ midpoint - 0.18 * np.sin(angles),
+ )),
+ labels_v=np.asarray([-0.018] * 4 + [0.018] * 4),
+ )
+ initial_gates = np.random.default_rng(7).normal(
+ 2.33, 0.02, circuit.edge_count)
+ initial_metrics, _ = evaluate_grid_classifier(
+ circuit, initial_gates, ring)
+ trained = train_grid_classifier(
+ circuit,
+ initial_gates,
+ ring,
+ ideal,
+ method="clean",
+ epochs=30,
+ standard_learning_time_seconds=1e-3,
+ record_every=10,
+ )
+ assert trained["hinge_loss_v2"] < initial_metrics["hinge_loss_v2"]
+ assert trained["local_updates"] > 0
print({
"nodes": circuit.node_count,
"edges": circuit.edge_count,
@@ -113,6 +142,8 @@ def main() -> None:
"target_nodes": circuit.target_nodes,
"constant_neutral_rmse_v_per_s": constant_rmse,
"quadratic_neutral_rmse_v_per_s": quadratic_rmse,
+ "clean_hinge_before": initial_metrics["hinge_loss_v2"],
+ "clean_hinge_after": trained["hinge_loss_v2"],
"autodiff_used": False,
})
diff --git a/sdil/physical_grid.py b/sdil/physical_grid.py
index 10d667b..bb6b4a4 100644
--- a/sdil/physical_grid.py
+++ b/sdil/physical_grid.py
@@ -312,3 +312,203 @@ def fit_edge_predictor(
predictor.coefficients[edge] = np.linalg.solve(
gram + ridge * np.eye(gram.shape[0]), rhs)
return int(len(states))
+
+
+@dataclass(frozen=True)
+class RingClassificationDataset:
+ inputs_v: Array
+ labels_v: Array
+
+ def __post_init__(self) -> None:
+ inputs = np.asarray(self.inputs_v)
+ labels = np.asarray(self.labels_v)
+ if inputs.ndim != 2 or inputs.shape[1] != 2:
+ raise ValueError("ring inputs must have shape (samples, 2)")
+ if labels.shape != (len(inputs),):
+ raise ValueError("ring labels must match the sample count")
+ if np.any(labels == 0.0):
+ raise ValueError("classification labels must be signed")
+
+
+def evaluate_grid_classifier(
+ circuit: GridCircuit,
+ gates: Array,
+ dataset: RingClassificationDataset,
+ *,
+ initial_states: list[Array | None] | None = None,
+) -> tuple[dict, list[Array]]:
+ if initial_states is None:
+ initial_states = [None] * len(dataset.labels_v)
+ states = []
+ outputs = []
+ for index, inputs in enumerate(dataset.inputs_v):
+ state = solve_grid_state(
+ circuit,
+ gates,
+ circuit.source_values(*inputs),
+ initial_state=initial_states[index],
+ )
+ states.append(state)
+ 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(),
+ }, states
+
+
+def train_grid_classifier(
+ circuit: GridCircuit,
+ initial_gates: Array,
+ dataset: RingClassificationDataset,
+ imperfection: GridSquareLawImperfection,
+ *,
+ method: str,
+ epochs: int,
+ predictor: EdgePolynomialPredictor | None = None,
+ standard_nudging: float = 128.0 / 129.0,
+ standard_learning_time_seconds: float = 1.0e-3,
+ overclamp_nudging: float = 32.0 / 129.0,
+ overclamp_target_magnitude_v: float | None = None,
+ overclamp_time_seconds_per_v: float = 0.05,
+ record_every: int = 50,
+) -> dict:
+ """Train the physical grid with explicit local voltage-square updates."""
+ allowed = {
+ "clean",
+ "raw",
+ "constant",
+ "sdil",
+ "oracle_neutral",
+ "overclamp_clean",
+ "overclamp",
+ "overclamp_constant",
+ "overclamp_sdil",
+ "overclamp_oracle_neutral",
+ }
+ if method not in allowed:
+ raise ValueError(f"unrecognized method {method}")
+ if epochs < 1:
+ raise ValueError("epochs must be positive")
+ if method in {
+ "constant", "sdil", "overclamp_constant", "overclamp_sdil"
+ } and predictor is None:
+ raise ValueError(f"{method} requires a predictor")
+ gates = np.asarray(initial_gates, dtype=float).copy()
+ if gates.shape != (circuit.edge_count,):
+ raise ValueError("initial gate vector has the wrong shape")
+ active_predictor = predictor.copy() if predictor is not None else None
+ target_magnitude = (
+ circuit.high_voltage
+ if overclamp_target_magnitude_v is None
+ else overclamp_target_magnitude_v)
+ is_overclamp = method.startswith("overclamp")
+ free_cache: list[Array | None] = [None] * len(dataset.labels_v)
+ clamped_cache: list[Array | None] = [None] * len(dataset.labels_v)
+ trace = []
+ cumulative_learning_time = 0.0
+ clamp_l2_time = 0.0
+ max_clamp_displacement = 0.0
+ local_updates = 0
+ clipped_updates = 0
+
+ for epoch in range(epochs):
+ for sample, (inputs, label) in enumerate(zip(
+ dataset.inputs_v, dataset.labels_v
+ )):
+ sources = circuit.source_values(*inputs)
+ free_state = solve_grid_state(
+ circuit,
+ gates,
+ sources,
+ initial_state=free_cache[sample],
+ )
+ free_cache[sample] = free_state
+ output_free = output_difference(circuit, free_state)
+ error = label - output_free
+ if label * error <= 0.0:
+ continue
+ if is_overclamp:
+ output_clamped = output_free + overclamp_nudging * (
+ target_magnitude * np.sign(error) - output_free)
+ duration = overclamp_time_seconds_per_v * abs(error)
+ else:
+ output_clamped = output_free + standard_nudging * error
+ duration = standard_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_grid_state(
+ circuit,
+ gates,
+ sources,
+ target_values=target_values,
+ initial_state=(
+ free_state if clamped_cache[sample] is None
+ else clamped_cache[sample]),
+ )
+ clamped_cache[sample] = clamped_state
+ free_drops = edge_voltage_drops(circuit, free_state)
+ 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)
+ neutral_bias = imperfection.neutral_bias(
+ circuit.measured_learning_rate, free_drops)
+ if method in {"clean", "overclamp_clean"}:
+ applied_rate = ideal_rate
+ elif method in {
+ "oracle_neutral", "overclamp_oracle_neutral"
+ }:
+ applied_rate = observed_rate - neutral_bias
+ elif method in {
+ "constant", "sdil", "overclamp_constant", "overclamp_sdil"
+ }:
+ applied_rate = observed_rate - active_predictor.predict(
+ free_drops)
+ 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
+ displacement = abs(output_clamped - output_free)
+ clamp_l2_time += duration * displacement * displacement
+ max_clamp_displacement = max(
+ max_clamp_displacement, displacement)
+ if epoch % record_every == 0 or epoch == epochs - 1:
+ metrics, free_cache = evaluate_grid_classifier(
+ circuit, gates, dataset, initial_states=free_cache)
+ trace.append({"epoch": epoch, **metrics})
+
+ final = trace[-1]
+ return {
+ "method": method,
+ "epochs": epochs,
+ "initial_gates_v": np.asarray(initial_gates).tolist(),
+ "final_gates_v": gates.tolist(),
+ "classification_error": final["classification_error"],
+ "hinge_loss_v2": final["hinge_loss_v2"],
+ "margin_success_fraction": final["margin_success_fraction"],
+ "outputs_v": final["outputs_v"],
+ "local_updates": local_updates,
+ "cumulative_learning_time_seconds": float(cumulative_learning_time),
+ "clamp_displacement_l2_time_v2_s": float(clamp_l2_time),
+ "max_abs_clamp_displacement_v": float(max_clamp_displacement),
+ "clipped_updates": clipped_updates,
+ "trace": trace,
+ }