"""Adversarial probe of the Onsager-corrected Sinkhorn (adaptive-TAP) proposal. Iteration under test: h = beta * V mu T - b_t * mu_prev ; mu = Sinkhorn(exp(h)) b_t = 0 recovers naive mean field, which is algebraically identical to the PGD form of entropic Gromov-Wasserstein already measured at 0.034 on omit-size. We sweep the Onsager coefficient over ten orders of magnitude so the verdict does not depend on getting the R-transform prescription exactly right: if no scalar b works, no principled derivation of b can rescue it. """ import json, sys, time import numpy as np import torch sys.path.insert(0, "/home/yurenh2/emm") from worldalign.synth_fast_gate import fast_pair_descent from worldalign.spectral_match import grampa, umeyama DEV = torch.device("cuda:0") DT = torch.float64 def standardise(M): m = ~np.eye(len(M), dtype=bool) v = M[m] out = (M - v.mean()) / v.std() np.fill_diagonal(out, 0.0) return out def sinkhorn_log(h, iters=60): """Doubly stochastic (row and col sums 1) projection of exp(h), batched.""" lg = h - h.amax(dim=(-2, -1), keepdim=True) for _ in range(iters): lg = lg - torch.logsumexp(lg, dim=-1, keepdim=True) lg = lg - torch.logsumexp(lg, dim=-2, keepdim=True) return torch.exp(lg) def hutch_chi(h, mu, eps=1e-3, iters=60): """(1/N^2) tr(d Sinkhorn / d h) by a random probe.""" Z = torch.randn_like(h) mu2 = sinkhorn_log(h + eps * Z, iters) return ((Z * (mu2 - mu)).sum(dim=(-2, -1)) / (eps * h.shape[-1] ** 2)) def r_transform(spec_prod, chi_grid): """Numerical R-transform of an empirical spectrum: invert Stieltjes.""" # G(z) = mean 1/(z - l); R(g) = G^{-1}(g) - 1/g lo, hi = spec_prod.max() * 1.0001, spec_prod.max() * 1e6 out = [] for g in chi_grid: a, b = lo, hi for _ in range(200): m = 0.5 * (a + b) if np.mean(1.0 / (m - spec_prod)) > g: a = m else: b = m z = 0.5 * (a + b) out.append(z - 1.0 / g) return np.array(out) def run_amp(V, T, betas, b_scale, mode, restarts, iters_per_beta, seed, spec_prod=None, damping=0.0): N = V.shape[-1] g = torch.Generator(device=DEV).manual_seed(seed) mu = torch.rand(restarts, N, N, generator=g, device=DEV, dtype=DT) * 0.01 mu = sinkhorn_log(mu) mu_prev = torch.full_like(mu, 1.0 / N) traj = [] for beta in betas: for _ in range(iters_per_beta): drive = beta * torch.matmul(torch.matmul(V, mu), T) if mode == "mf": b = torch.zeros(restarts, device=DEV, dtype=DT) else: chi = hutch_chi(drive, mu) if mode == "iid": coef = beta * beta * N # sum_j V_ij^2 ~ N elif mode == "ri": cg = np.clip(chi.abs().cpu().numpy(), 1e-12, None) coef = torch.tensor( r_transform(spec_prod * beta, cg), device=DEV, dtype=DT) b = b_scale * coef * chi h = drive - b.view(-1, 1, 1) * mu_prev new = sinkhorn_log(h) if damping: new = (1 - damping) * new + damping * mu mu_prev, mu = mu, new if not torch.isfinite(mu).all(): return None, traj traj.append(float(mu.amax(dim=-1).mean())) return mu, traj def score(mu, hidden, V_gpu, Tsh_gpu, N, refine=True): from scipy.optimize import linear_sum_assignment accs, refs = [], [] for b in range(mu.shape[0]): _, col = linear_sum_assignment(-mu[b].cpu().numpy()) a = float((hidden[col] == np.arange(N)).mean()) accs.append(a) if refine: fin = fast_pair_descent( Tsh_gpu, V_gpu, torch.from_numpy(np.ascontiguousarray(col)).to(DEV), 600) refs.append(float((hidden[fin.cpu().numpy()] == np.arange(N)).mean())) return accs, refs def main(): d = torch.load("/home/yurenh2/emm/artifacts/synth_v1/omit_size.pt", map_location="cpu", weights_only=False) V = standardise(d["visual_field"].double().numpy()) T = standardise(d["text_field"].double().numpy()) N = len(V) lv = np.linalg.eigvalsh(V) results = [] for trial in range(2): rng = np.random.default_rng(trial) hidden = rng.permutation(N) Tsh = standardise(T[np.ix_(hidden, hidden)]) lt = np.linalg.eigvalsh(Tsh) spec_prod = np.outer(lv, lt).ravel() / N # operator scale of V (x) T / N Vg = torch.from_numpy(V).to(DEV).to(DT) Tg = torch.from_numpy(Tsh).to(DEV).to(DT) # sanity: GRAMPA / Umeyama reference on this exact instance if trial == 0: gp = grampa(V, Tsh, 1.0) um = umeyama(V, Tsh) print("REF grampa %.4f umeyama %.4f chance %.4f" % (float((hidden[gp] == np.arange(N)).mean()), float((hidden[um] == np.arange(N)).mean()), 1 / N), flush=True) betas_ladder = list(np.geomspace(0.02, 2.0, 20)) for mode in ("mf", "iid", "ri"): scales = [0.0] if mode == "mf" else [1e-4, 1e-2, 0.1, 0.3, 1.0, 3.0, 10.0, 100.0] for sc in scales: t0 = time.time() mu, traj = run_amp(Vg, Tg, betas_ladder, sc, mode, 8, 25, 100 + trial, spec_prod=spec_prod) if mu is None: print("trial%d %-4s scale=%-8g DIVERGED after %d beta steps (%.1fs)" % (trial, mode, sc, len(traj), time.time() - t0), flush=True) results.append(dict(trial=trial, mode=mode, scale=sc, diverged=True, betas_done=len(traj))) continue accs, refs = score(mu, hidden, Vg, Tg, N) print("trial%d %-4s scale=%-8g raw mean %.4f max %.4f | refined mean %.4f max %.4f | sharp %.3f (%.1fs)" % (trial, mode, sc, np.mean(accs), np.max(accs), np.mean(refs), np.max(refs), traj[-1], time.time() - t0), flush=True) results.append(dict(trial=trial, mode=mode, scale=sc, diverged=False, raw=accs, refined=refs, sharpness=traj[-1])) with open("/home/yurenh2/emm/logs/amp_probe.json", "w") as f: json.dump(results, f, indent=1) print("DONE") if __name__ == "__main__": main()