From b5c1b5be1628664977fb86bd7456b4176b204320 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Sat, 29 Aug 2026 13:13:52 -0500 Subject: feat: train the reconstructed physical grid --- sdil/physical_grid.py | 200 ++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 200 insertions(+) (limited to 'sdil/physical_grid.py') 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, + } -- cgit v1.2.3