diff --git a/examples/advanced/image_classification.py b/examples/advanced/image_classification.py new file mode 100644 index 0000000..13c627d --- /dev/null +++ b/examples/advanced/image_classification.py @@ -0,0 +1,286 @@ +""" +Fashion-MNIST controlled-correlation p-bit demo. + +The image-classification task is expressed largely as a probabilistic circuit, +which stores the model through base h/J terms and named correlation components. +The solver controls how these components are activated over time. + +Image evidence is applied through h, while each class is represented by a learned +J/h correlation component that is gradually activated by CorrelationAnnealingSolver. + +The current model uses 720 p-bits. The implementation is intentionally simple. +Calculation speed has been optimized, but it can be further improved. + +This is POC demo and classification accuracy can also be improved. +Inference is slower/heavier than training in general. + +Results: + 04/09/2026 On 1000 test images accuracy is 55%, takes 370s on a laptop +""" +from pathlib import Path +from urllib.request import urlretrieve +import gzip, struct +import numpy as np +from sklearn.cluster import MiniBatchKMeans + +from p_kit.psl import PCircuit +from p_kit.solver.corr_ann_solver import CorrelationAnnealingSolver, staged_ramp + +from joblib import Parallel, delayed +import os +import time + +N_WORKERS = min(4, os.cpu_count() or 1) +# ---------------------------------------------------------------------- +# Configuration +# ---------------------------------------------------------------------- + +SIZE = 28 +N_FEATURES = 10 +PATCH_H, PATCH_W = 6, 7 +STRIDE = 3 +FEATURE_PATCHES = 40000 +FEATURE_BETA = 2.0 + +TRAIN_PER_CLASS = 2000 +OFFSETS = [(0,1),(0,2),(0,3),(1,0),(2,0),(3,0), + (1,1),(2,2),(1,-1),(2,-2),(1,2),(2,1),(1,-2),(2,-1)] + +I0 = 2.0 +SAMPLES = 10 +NT = 60 +TEST_LIMIT = 200 # max test images +SEED = 1234 + +rng = np.random.default_rng(SEED) + +# ---------------------------------------------------------------------- +# Fashion-MNIST +# ---------------------------------------------------------------------- + +DATA = Path("data/fashion") +BASE = "https://raw.githubusercontent.com/zalandoresearch/fashion-mnist/master/data/fashion/" +FILES = ["train-images-idx3-ubyte.gz", "train-labels-idx1-ubyte.gz", + "t10k-images-idx3-ubyte.gz", "t10k-labels-idx1-ubyte.gz"] + +DATA.mkdir(parents=True, exist_ok=True) +for name in FILES: + p = DATA / name + if not p.exists(): + print("Downloading", name) + urlretrieve(BASE + name, p) + +def load_images(path): + with gzip.open(path, "rb") as f: + _, n, r, c = struct.unpack(">IIII", f.read(16)) + return np.frombuffer(f.read(), np.uint8).reshape(n, r, c) / 255.0 + +def load_labels(path): + with gzip.open(path, "rb") as f: + _, n = struct.unpack(">II", f.read(8)) + return np.frombuffer(f.read(), np.uint8) + +x_train = load_images(DATA / FILES[0]) +y_train = load_labels(DATA / FILES[1]) +x_test = load_images(DATA / FILES[2]) +y_test = load_labels(DATA / FILES[3]) + +# ---------------------------------------------------------------------- +# Shared local categorical features +# ---------------------------------------------------------------------- + +def positions(size, patch): + p = list(range(0, size - patch + 1, STRIDE)) + if p[-1] != size - patch: + p.append(size - patch) + return p + +YS, XS = positions(SIZE, PATCH_H), positions(SIZE, PATCH_W) +GH, GW = len(YS), len(XS) +GRID_N = GH * GW +N_PBITS = GRID_N * N_FEATURES + +patches = [] +while len(patches) < FEATURE_PATCHES: + im = x_train[rng.integers(len(x_train))] + y, x = rng.choice(YS), rng.choice(XS) + p = im[y:y+PATCH_H, x:x+PATCH_W] + if p.mean() > .04: + patches.append(p.ravel()) + +km = MiniBatchKMeans(N_FEATURES - 1, batch_size=1024, n_init=5, + random_state=SEED).fit(patches) +templates = np.vstack([np.zeros(PATCH_H * PATCH_W), km.cluster_centers_]) + +def feature_prob(images): + images = np.asarray(images) + if images.ndim == 2: + images = images[None] + P = np.stack([images[:, y:y+PATCH_H, x:x+PATCH_W].reshape(len(images), -1) + for y in YS for x in XS], axis=1) + d = ((P[:,:,None] - templates[None,None]) ** 2).mean(-1) + z = (d - d.min(2, keepdims=True)) / (d.std(2, keepdims=True) + 1e-6) + z = -FEATURE_BETA * z + z -= z.max(2, keepdims=True) + e = np.exp(z) + return e / e.sum(2, keepdims=True) + +def image_state(image): + p = feature_prob(image)[0] + u = np.log(np.clip(p, 1e-6, 1)) + u -= u.mean(1, keepdims=True) + q = np.zeros_like(p) + q[np.arange(GRID_N), p.argmax(1)] = 1 + return (u / 2).ravel(), (2 * q - 1).ravel() + +# ---------------------------------------------------------------------- +# Learn class-dependent feature correlations +# ---------------------------------------------------------------------- + +def shifted(F, dy, dx): + ya, yb = (slice(0,-dy), slice(dy,None)) if dy else (slice(None), slice(None)) + if dx > 0: + xa, xb = slice(0,-dx), slice(dx,None) + elif dx < 0: + xa, xb = slice(-dx,None), slice(0,dx) + else: + xa = xb = slice(None) + return F[:,ya,xa], F[:,yb,xb] + +priors = np.zeros((10, N_FEATURES)) +joints = np.zeros((10, len(OFFSETS), N_FEATURES, N_FEATURES)) +acount = np.zeros((10, len(OFFSETS), N_FEATURES)) + +print("Learning class correlations...") + +for c in range(10): + ids = rng.choice(np.flatnonzero(y_train == c), TRAIN_PER_CLASS, replace=False) + F = feature_prob(x_train[ids]).reshape(-1, GH, GW, N_FEATURES) + priors[c] = F.sum((0,1,2)) + + for o, (dy, dx) in enumerate(OFFSETS): + a, b = shifted(F, dy, dx) + a, b = a.reshape(-1, N_FEATURES), b.reshape(-1, N_FEATURES) + acount[c,o] = a.sum(0) + joints[c,o] = a.T @ b + +alpha = 1.0 +prior_p = (priors + alpha/N_FEATURES) / (priors.sum(1, keepdims=True) + alpha) +global_prior = priors.sum(0) / priors.sum() + +cond = (joints + alpha/N_FEATURES) / (acount[:,:,:,None] + alpha) +global_cond = (joints.sum(0) + alpha/N_FEATURES) / (acount.sum(0)[:,:,None] + alpha) + +prior_w = np.clip(np.log((prior_p + 1e-12) / (global_prior + 1e-12)), -2.5, 2.5) +rel_w = np.clip(np.log((cond + 1e-12) / (global_cond[None] + 1e-12)), -2.5, 2.5) + +# ---------------------------------------------------------------------- +# Convert categorical correlations to Ising J/h +# ---------------------------------------------------------------------- + +def block(y, x): + s = (y * GW + x) * N_FEATURES + return np.arange(s, s + N_FEATURES) + +def pairs(dy, dx): + for y in range(GH): + for x in range(GW): + yy, xx = y + dy, x + dx + if 0 <= yy < GH and 0 <= xx < GW: + yield y, x, yy, xx + +def build_model(c): + J = np.zeros((N_PBITS, N_PBITS)) + h = np.zeros(N_PBITS) + offset = 0.0 + U = np.tile(prior_w[c], (GRID_N, 1)) + + for o, (dy, dx) in enumerate(OFFSETS): + W = rel_w[c,o] + row, col, mean = W.mean(1), W.mean(0), W.mean() + W0 = W - row[:,None] - col[None,:] + mean + + for y, x, yy, xx in pairs(dy, dx): + a, b = block(y,x), block(yy,xx) + J[np.ix_(a,b)] += W0 / 4 + J[np.ix_(b,a)] += W0.T / 4 + U[y*GW+x] += row + U[yy*GW+xx] += col + offset += mean + + for s in range(GRID_N): + u = U[s] + mean = u.mean() + h[s*N_FEATURES:(s+1)*N_FEATURES] = (u - mean) / 2 + offset -= mean + + return J, h, offset + +models = [build_model(c) for c in range(10)] +radius = max(np.max(np.abs(np.linalg.eigvalsh(J))) for J, _, _ in models) +scale = 1.0 / radius + +circuits, solvers, offsets = [], [], [] + +for c, (J, h, offset) in enumerate(models): + circuit = PCircuit(N_PBITS) + circuit.set_correlation_component("class", J=scale*J, h=scale*h) + + solver = CorrelationAnnealingSolver( + Nt=NT, dt=.1667, i0=I0, seed=SEED+c, + block_size=N_FEATURES, + component_schedules={"class": staged_ramp(.2, .8)} + ) + + circuits.append(circuit) + solvers.append(solver) + offsets.append(scale * offset) + +# ---------------------------------------------------------------------- +# Controlled-correlation classification +# ---------------------------------------------------------------------- + +def classify(image): + h_image, initial = image_state(image) + E = np.zeros(10) + + for c in range(10): + circuits[c].h = h_image + _, best_E = solvers[c].solve( + circuits[c], n_shots=SAMPLES, initial_state=initial, + return_best=True, target_scales={"class": 1.0} + ) + E[c] = best_E.mean() - I0 * offsets[c] + + return E.argmax() + +def classify_chunk(indices): + return [(i, classify(x_test[i])) for i in indices] + + +# ---------------------------------------------------------------------- +# Main code +# ---------------------------------------------------------------------- + +if __name__ == "__main__": + t0 = time.perf_counter() + + print(f"Running controlled-correlation demo ({N_PBITS} p-bits, {N_WORKERS} workers)...") + + n = min(TEST_LIMIT, len(x_test)) + chunks = np.array_split(np.arange(n), N_WORKERS) + + results = Parallel(n_jobs=N_WORKERS, verbose=10)( + delayed(classify_chunk)(chunk) for chunk in chunks + ) + + pred = np.empty(n, dtype=int) + for chunk in results: + for i, p in chunk: + pred[i] = p + + elapsed = time.perf_counter() - t0 + correct = np.sum(pred == y_test[:n]) + + print(f"Final accuracy: {100*correct/n:.2f}%") + print(f"Elapsed: {elapsed:.1f} s ({elapsed/60:.1f} min)") \ No newline at end of file diff --git a/p_kit/psl/p_circuit.py b/p_kit/psl/p_circuit.py index d191da8..9547d30 100644 --- a/p_kit/psl/p_circuit.py +++ b/p_kit/psl/p_circuit.py @@ -18,15 +18,15 @@ class PCircuit: biases J : np.array((n_pbits, n_pbits)) weights - ports: dict[str, object] - Circuit ports + ports : Dict[str, Any] | None + circuit ports """ def __init__(self, n_pbits: int, ports: Dict[str, Any] = None): self.n_pbits = n_pbits - self.ports = ports #Kept for copy behavior - + self.ports = ports + self.h = np.zeros((n_pbits,)) self.J = np.zeros((n_pbits, n_pbits)) self._connections = {} @@ -34,7 +34,8 @@ def __init__(self, n_pbits: int, ports: Dict[str, Any] = None): # identity (e.g. CaSuDaSolver's cache_J) can detect staleness # even though id(self.J) doesn't change across set_weight() calls. self._j_version = 0 - + self._correlation_components = {} + if ports: self._initialize_ports(ports) @@ -46,16 +47,14 @@ def _initialize_ports(self, port_attrs: Dict[str, Any]) -> None: for name, port in port_attrs.items(): port_indices[name] = idx idx += port.width - + # Set up each port for name, port in port_attrs.items(): new_port = Port(name=port.name, width=port.width) new_port.circuit = self new_port.index = port_indices[name] # Check ports name doesn't conflict with reserved attributes. - assert name != "ports" - assert name != "h" - assert name != "J" + assert name not in ("ports", "h", "J") setattr(self, name, new_port) def set_weight(self, from_pbit, to_pbit, weight, sym=True): @@ -64,8 +63,86 @@ def set_weight(self, from_pbit, to_pbit, weight, sym=True): self.J[to_pbit, from_pbit] = weight self._j_version += 1 + def add_correlation(self, i, j, weight=1.0, sym=True): + """Add positive coupling. i,j: p-bit indices; weight: strength; sym: symmetric coupling.""" + w = abs(weight) + self.J[i, j] += w + if sym: + self.J[j, i] += w + self._j_version += 1 + + def add_anticorrelation(self, i, j, weight=1.0, sym=True): + """Add negative coupling. i,j: p-bit indices; weight: strength; sym: symmetric coupling.""" + w = -abs(weight) + self.J[i, j] += w + if sym: + self.J[j, i] += w + self._j_version += 1 + + def add_group_coupling(self, group_a, group_b, weights=1.0, sym=True): + """Couple two groups. group_a/group_b: indices; weights: scalar or matrix; sym: symmetric coupling.""" + a = np.asarray(group_a, dtype=int) + b = np.asarray(group_b, dtype=int) + W = np.full((len(a), len(b)), float(weights)) if np.isscalar(weights) else np.asarray(weights, dtype=float) + if W.shape != (len(a), len(b)): + raise ValueError(f"weights shape {W.shape}, expected {(len(a), len(b))}") + self.J[np.ix_(a, b)] += W + if sym: + self.J[np.ix_(b, a)] += W.T + self._j_version += 1 + + def add_group_competition(self, group, strength=1.0): + """Add pairwise anticorrelation within a group. group: indices; strength: inhibition strength.""" + g = np.asarray(group, dtype=int) + w = abs(float(strength)) + for i in range(len(g)): + for j in range(i + 1, len(g)): + self.J[g[i], g[j]] -= w + self.J[g[j], g[i]] -= w + self._j_version += 1 + + def set_correlation_component(self, name, J=None, h=None): + """Set named dynamic component. name: identifier; J: coupling matrix; h: bias vector.""" + J = np.zeros_like(self.J) if J is None else np.asarray(J, dtype=float) + h = np.zeros_like(self.h) if h is None else np.asarray(h, dtype=float).reshape(-1) + if J.shape != self.J.shape: + raise ValueError(f"J shape {J.shape}, expected {self.J.shape}") + if h.shape != self.h.shape: + raise ValueError(f"h shape {h.shape}, expected {self.h.shape}") + self._correlation_components[name] = {"J": J.copy(), "h": h.copy()} + + def get_correlation_component(self, name): + """Return named component. name: component identifier.""" + return self._correlation_components[name] + + def remove_correlation_component(self, name): + """Remove named component. name: component identifier.""" + del self._correlation_components[name] + + def clear_correlation_components(self): + self._correlation_components.clear() + + def correlation_components(self): + return tuple(self._correlation_components) + + def effective_parameters(self, component_scales=None): + """Return effective J,h. component_scales: {name: scale} for active components.""" + J = self.J.copy() + h = self.h.copy() + if component_scales: + for name, scale in component_scales.items(): + c = self._correlation_components[name] + J += scale * c["J"] + h += scale * c["h"] + return J, h + def copy(self): new_circuit = PCircuit(self.n_pbits, self.ports) new_circuit.J = self.J new_circuit.h = self.h - return new_circuit + new_circuit._j_version = self._j_version + new_circuit._correlation_components = { + name: {"J": c["J"].copy(), "h": c["h"].copy()} + for name, c in self._correlation_components.items() + } + return new_circuit \ No newline at end of file diff --git a/p_kit/solver/corr_ann_solver.py b/p_kit/solver/corr_ann_solver.py new file mode 100644 index 0000000..712809a --- /dev/null +++ b/p_kit/solver/corr_ann_solver.py @@ -0,0 +1,299 @@ +import numpy as np + +from p_kit.solver.annealing import constant +from p_kit.solver.base_solver import Solver + + +def hold(value=1.0): + """Return a component schedule that holds one constant scale.""" + value = float(value) + def schedule(_solver, _run): + return value + return schedule + + +def linear_ramp(start_value=0.0, end_value=1.0, + start_fraction=0.0, end_fraction=1.0): + """Ramp a correlation component over a fraction of the solver run.""" + start_value, end_value = float(start_value), float(end_value) + start_fraction, end_fraction = float(start_fraction), float(end_fraction) + if not 0.0 <= start_fraction <= 1.0: + raise ValueError("start_fraction must be in [0, 1]") + if not 0.0 <= end_fraction <= 1.0: + raise ValueError("end_fraction must be in [0, 1]") + if end_fraction <= start_fraction: + raise ValueError("end_fraction must be greater than start_fraction") + + def schedule(solver, run): + x = 1.0 if solver.Nt <= 1 else run / (solver.Nt - 1) + if x <= start_fraction: + return start_value + if x >= end_fraction: + return end_value + t = (x - start_fraction) / (end_fraction - start_fraction) + return start_value + t * (end_value - start_value) + return schedule + + +def staged_ramp(start_fraction=0.20, end_fraction=0.80): + """0 -> 1 schedule: evidence-only, ramp correlations, full relaxation.""" + return linear_ramp(0.0, 1.0, start_fraction, end_fraction) + + +class CorrelationAnnealingSolver(Solver): + """Gibbs solver with dynamically controlled named J/h components. + + J(t) = J_base + sum_r alpha_r(t) J_r + h(t) = h_base + sum_r alpha_r(t) h_r + + block_size=None uses binary Gibbs updates; block_size=K uses exact + categorical K-p-bit updates. Validation and block-sparse J layouts are + cached across solves; component fields are updated incrementally. + """ + + def __init__(self, Nt, dt, i0, expected_mean=0, seed=None, backend=None, + tau=0.1, component_schedules=None, block_size=None): + super().__init__(Nt, dt, i0, expected_mean, seed, backend, tau) + if self.backend.xp is not np: + raise NotImplementedError("CorrelationAnnealingSolver currently supports NumpyBackend only") + if block_size is not None and int(block_size) < 2: + raise ValueError("block_size must be >= 2") + self.block_size = None if block_size is None else int(block_size) + self.component_schedules = dict(component_schedules or {}) + self._structure_key = None + self._layouts = None + + def _components(self, c): + return {name: c.get_correlation_component(name) + for name in c.correlation_components()} + + def _schedule_value(self, name, run): + schedule = self.component_schedules.get(name, 1.0) + return float(schedule(self, run) if callable(schedule) else schedule) + + def _cache_key(self, c, components): + return (id(c), id(c.J), getattr(c, "_j_version", 0), self.block_size, + tuple((name, id(comp["J"])) for name, comp in components.items())) + + def _validate_J_components(self, c, components, tol=1e-12): + matrices = [("base", np.asarray(c.J, dtype=float))] + matrices += [(name, np.asarray(comp["J"], dtype=float)) + for name, comp in components.items()] + for name, J in matrices: + if not np.allclose(J, J.T, atol=tol, rtol=0): + raise ValueError(f"component '{name}' must be symmetric for Gibbs energy sampling") + if self.block_size is None: + if np.max(np.abs(np.diag(J))) > tol: + raise ValueError(f"component '{name}' must have zero diagonal for binary Gibbs updates") + else: + K = self.block_size + for start in range(0, c.n_pbits, K): + if np.max(np.abs(J[start:start+K, start:start+K])) > tol: + raise ValueError( + f"component '{name}' has intra-block J couplings; " + "exact categorical block updates require zero coupling within each block" + ) + + def _block_layout(self, J): + K, n = self.block_size, len(J) + nb = n // K + rows, count = [], 0 + for a in range(nb): + row = [] + for b in range(nb): + W = J[a*K:(a+1)*K, b*K:(b+1)*K] + if np.any(W): + row.append((b, W)) + count += 1 + rows.append(row) + return rows if count < nb * nb / 2 else None + + def _prepare(self, c, components): + key = self._cache_key(c, components) + if key == self._structure_key: + return self._layouts + self._validate_J_components(c, components) + if self.block_size is None: + layouts = None + else: + layouts = {"base": self._block_layout(np.asarray(c.J, dtype=float))} + layouts.update({name: self._block_layout(np.asarray(comp["J"], dtype=float)) + for name, comp in components.items()}) + self._structure_key, self._layouts = key, layouts + return layouts + + def _validate_binary_initial(self, initial_state, n_shots, n_pbits): + if initial_state is None: + return np.where(self.random((n_shots, n_pbits)) < 0.5, -1.0, 1.0) + m = np.asarray(initial_state, dtype=float) + if m.ndim == 1: + if m.size != n_pbits: + raise ValueError(f"initial_state must contain {n_pbits} values") + m = np.tile(m, (n_shots, 1)) + elif m.shape != (n_shots, n_pbits): + raise ValueError(f"initial_state shape {m.shape}, expected {(n_shots, n_pbits)}") + if not np.all((m == -1) | (m == 1)): + raise ValueError("initial_state must contain only -1 or +1") + return m.copy() + + def _validate_block_initial(self, initial_state, n_shots, n_pbits): + K = self.block_size + if n_pbits % K: + raise ValueError(f"n_pbits ({n_pbits}) must be divisible by block_size ({K})") + nb = n_pbits // K + if initial_state is None: + m = -np.ones((n_shots, nb, K)) + winners = self._generator.integers(0, K, size=(n_shots, nb)) + m[np.arange(n_shots)[:, None], np.arange(nb), winners] = 1.0 + return m.reshape(n_shots, n_pbits) + m = self._validate_binary_initial(initial_state, n_shots, n_pbits) + if not np.all((m.reshape(n_shots, nb, K) > 0).sum(axis=2) == 1): + raise ValueError("block initial_state must contain exactly one +1 per block") + return m + + def _target_energy(self, m, fields, base_h, comp_h, scales): + field = fields["base"].copy() + h = base_h.copy() + for name, scale in scales.items(): + field += scale * fields[name] + h += scale * comp_h[name] + return self.i0 * (0.5 * np.sum(m * field, axis=1) + m @ h) + + def _apply_block_delta(self, field, live, delta, blk, J, layout, scale=1.0): + K = self.block_size + if layout is None: + d = delta @ J[blk*K:(blk+1)*K, :] + field += d + live += scale * d + return + for dst, W in layout[blk]: + idx = slice(dst*K, (dst+1)*K) + d = delta @ W + field[:, idx] += d + live[:, idx] += scale * d + + def _binary_sweep(self, m, base_J, comp_J, base_h, comp_h, fields, scales, beta): + live = fields["base"] + base_h + for name, scale in scales.items(): + live += scale * (fields[name] + comp_h[name]) + + for i in self._generator.permutation(m.shape[1]): + logits = np.clip(2.0 * beta * live[:, i], -60.0, 60.0) + p = 1.0 / (1.0 + np.exp(-logits)) + new = np.where(self.random((len(m),)) < p, 1.0, -1.0) + delta = new - m[:, i] + m[:, i] = new + + d = delta[:, None] * base_J[i][None, :] + fields["base"] += d + live += d + for name, scale in scales.items(): + d = delta[:, None] * comp_J[name][i][None, :] + fields[name] += d + live += scale * d + return m + + def _block_sweep(self, m, base_J, comp_J, base_h, comp_h, + fields, scales, layouts, beta): + K, nb = self.block_size, m.shape[1] // self.block_size + live = fields["base"] + base_h + for name, scale in scales.items(): + live += scale * (fields[name] + comp_h[name]) + + for blk in self._generator.permutation(nb): + idx = slice(blk*K, (blk+1)*K) + logits = 2.0 * beta * live[:, idx] + logits -= logits.max(axis=1, keepdims=True) + p = np.exp(logits) + p /= p.sum(axis=1, keepdims=True) + + cdf = np.cumsum(p, axis=1) + cdf[:, -1] = 1.0 + winner = (cdf < self.random((len(m), 1))).sum(axis=1) + + new = -np.ones((len(m), K)) + new[np.arange(len(m)), winner] = 1.0 + delta = new - m[:, idx] + m[:, idx] = new + + self._apply_block_delta(fields["base"], live, delta, blk, + base_J, layouts["base"]) + for name, scale in scales.items(): + self._apply_block_delta(fields[name], live, delta, blk, + comp_J[name], layouts[name], scale) + return m + + def solve(self, c, annealing_func=constant, n_shots=1, + initial_state=None, return_final=False, return_best=False, + target_scales=None): + """Run controlled-correlation Gibbs annealing.""" + if return_final and return_best: + raise ValueError("return_final and return_best are mutually exclusive") + + components = self._components(c) + layouts = self._prepare(c, components) + base_J = np.asarray(c.J, dtype=float) + base_h = np.asarray(c.h, dtype=float).reshape(-1) + comp_J = {name: np.asarray(comp["J"], dtype=float) for name, comp in components.items()} + comp_h = {name: np.asarray(comp["h"], dtype=float).reshape(-1) + for name, comp in components.items()} + + if self.block_size is None: + m = self._validate_binary_initial(initial_state, n_shots, c.n_pbits) + else: + m = self._validate_block_initial(initial_state, n_shots, c.n_pbits) + + fields = {"base": m @ base_J} + fields.update({name: m @ J for name, J in comp_J.items()}) + target = {name: float((target_scales or {}).get(name, 1.0)) + for name in components} + + best_m = m.copy() + best_E = self._target_energy(m, fields, base_h, comp_h, target) + + if not return_final and not return_best: + all_m = np.zeros((self.Nt, n_shots, c.n_pbits)) + all_I = np.zeros((self.Nt, c.n_pbits)) + E = np.zeros(self.Nt) + + for run in range(self.Nt): + scales = {name: self._schedule_value(name, run) for name in components} + beta = float(annealing_func(self, run)) + + if not return_final and not return_best: + live = fields["base"] + base_h + for name, scale in scales.items(): + live += scale * (fields[name] + comp_h[name]) + all_I[run] = (beta * live)[0] + + if self.block_size is None: + m = self._binary_sweep(m, base_J, comp_J, base_h, comp_h, + fields, scales, beta) + else: + m = self._block_sweep(m, base_J, comp_J, base_h, comp_h, + fields, scales, layouts, beta) + + target_E = self._target_energy(m, fields, base_h, comp_h, target) + improved = target_E > best_E + best_E[improved] = target_E[improved] + best_m[improved] = m[improved] + + if not return_final and not return_best: + all_m[run], E[run] = m, target_E[0] + + if return_best: + return best_m, best_E + if return_final: + return m[0] if n_shots == 1 else m + if n_shots == 1: + return all_I, all_m[:, 0, :], E + return all_m + + def copy(self): + return CorrelationAnnealingSolver( + Nt=self.Nt, dt=self.dt, i0=self.i0, + expected_mean=self.expected_mean, seed=self.seed, + backend=self.backend, tau=self.tau, + component_schedules=self.component_schedules, + block_size=self.block_size, + ) \ No newline at end of file