diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 13:13:52 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 13:13:52 -0500 |
| commit | b5c1b5be1628664977fb86bd7456b4176b204320 (patch) | |
| tree | 1b8917552e69282de0d9beb255087301b64e3405 | |
| parent | b9742072e2d9a80808fea1f314309f31476b80c8 (diff) | |
feat: train the reconstructed physical grid
| -rw-r--r-- | experiments/physical_grid_smoke.py | 31 | ||||
| -rw-r--r-- | sdil/physical_grid.py | 200 |
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, + } |
