import itertools import json import math import random from pathlib import Path import numpy as np import torch from torch import nn SEED = 1612 random.seed(SEED) np.random.seed(SEED) torch.manual_seed(SEED) torch.set_default_dtype(torch.float64) try: DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") except Exception: DEVICE = torch.device("cpu") # This experiment is intentionally small; on CUDA, an allocation/runtime error # falls back to CPU as required by the experiment protocol. def safe_device_call(fn): global DEVICE try: return fn(DEVICE) except Exception: DEVICE = torch.device("cpu") torch.cuda.empty_cache() if torch.cuda.is_available() else None return fn(DEVICE) def perms_and_signs(n, device): ps, ss = [], [] for p in itertools.permutations(range(n)): inv = sum(p[i] > p[j] for i in range(n) for j in range(i + 1, n)) ps.append(p) ss.append(-1.0 if inv % 2 else 1.0) return ps, torch.tensor(ss, device=device) class EquivariantVectorField(nn.Module): """v_i = MLP([x_i, mean(x), mean_j phi(x_i-x_j)]).""" def __init__(self, d, hidden=32): super().__init__() self.pair = nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, hidden), nn.Tanh()) self.out = nn.Sequential(nn.Linear(2 * d + hidden, hidden), nn.Tanh(), nn.Linear(hidden, d)) def forward(self, x, t=1.0): # x: [batch, N, d]; sum/mean over the unordered index set dif = x[:, :, None, :] - x[:, None, :, :] messages = self.pair(dif).mean(dim=2) pooled = x.mean(dim=1, keepdim=True).expand_as(x) return float(t) * self.out(torch.cat([x, pooled, messages], dim=-1)) class FlatVectorField(nn.Module): """Label-sensitive baseline with the same general hidden width.""" def __init__(self, n, d, hidden=32): super().__init__() self.net = nn.Sequential(nn.Linear(n * d + 1, hidden), nn.Tanh(), nn.Linear(hidden, hidden), nn.Tanh(), nn.Linear(hidden, n * d)) def forward(self, x, t=1.0): b = x.shape[0] return float(t) * self.net(torch.cat([x.reshape(b, -1), torch.full((b, 1), float(t), device=x.device)], dim=1)).reshape_as(x) def euler_flow(field, x, horizon=1.0, steps=16): z = x.clone() dt = horizon / steps for k in range(steps): z = z + dt * field(z, (k + 0.5) * dt) return z def relative(a, b, eps=1e-12): return ((a - b).pow(2).sum(dim=tuple(range(1, a.ndim))) / (b.pow(2).sum(dim=tuple(range(1, b.ndim))) + eps)).mean().sqrt().item() def equivariance_error(field, x, perm): xp = x[:, perm] return relative(field(xp, .7), field(x, .7)[:, perm]) def antisymmetrizer(base, x): # base accepts [B,N,D] and returns [B] scalar values n = x.shape[1] ps, signs = perms_and_signs(n, x.device) vals = torch.stack([base(x[:, p]) for p in ps], dim=1) return (vals * signs[None, :]).mean(dim=1) def antisym_error(base, x, perm): a = antisymmetrizer(base, x) ap = antisymmetrizer(base, x[:, perm]) inv = sum(perm[i] > perm[j] for i in range(len(perm)) for j in range(i + 1, len(perm))) sign = -1.0 if inv % 2 else 1.0 return (torch.abs(ap - sign * a) / (torch.abs(a) + 1e-12)).mean().item() class BaseScalar(nn.Module): def __init__(self, n, d, hidden=32): super().__init__() self.net = nn.Sequential(nn.Linear(n * d, hidden), nn.Tanh(), nn.Linear(hidden, hidden), nn.Tanh(), nn.Linear(hidden, 1)) def forward(self, x): return self.net(x.reshape(x.shape[0], -1)).squeeze(-1) class SymmetricJ(nn.Module): def __init__(self, d, hidden=24): super().__init__() self.net = nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, 1)) def forward(self, x): return self.net(x).squeeze(-1).mean(dim=1) def task_train(model, n, d, steps=350): opt = torch.optim.Adam(model.parameters(), lr=3e-3) for _ in range(steps): x = torch.randn(64, n, d, device=next(model.parameters()).device) # invariant target: mean squared radius plus a smooth interaction term target = x.pow(2).sum(-1).mean(-1) + 0.15 * torch.tanh((x[:, :, None] * x[:, None, :]).sum(-1).mean((1, 2))) pred = model(x).squeeze(-1) loss = (pred - target).pow(2).mean() opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): x = torch.randn(256, n, d, device=next(model.parameters()).device) target = x.pow(2).sum(-1).mean(-1) + 0.15 * torch.tanh((x[:, :, None] * x[:, None, :]).sum(-1).mean((1, 2))) p = torch.randperm(n, device=x.device) # Generalization to a fresh unseen permutation err = ((model(x[:, p]).squeeze(-1) - target) ** 2).mean().sqrt().item() return err def run(device): torch.manual_seed(SEED) n, d = 4, 2 x = torch.randn(48, n, d, device=device) perm = (1, 3, 0, 2) eq = EquivariantVectorField(d).to(device) flat = FlatVectorField(n, d).to(device) base = BaseScalar(n, d).to(device) sj = SymmetricJ(d).to(device) # Prediction 1: exact equivariance independent of flow scale and integration depth. scales = [0.0, 0.25, 0.5, 1.0, 2.0] eq_rows = [] for scale in scales: e1 = equivariance_error(eq, x * scale + 0.1, perm) e2 = equivariance_error(flat, x * scale + 0.1, perm) eq_rows.append({"scale": scale, "structured": e1, "vanilla": e2}) # Prediction 2: explicit antisymmetrization and symmetric factor are exact. as_rows = [] for k in [1, 2, 3, 4, 5, 6]: xx = torch.randn(24, k, d, device=device) b = BaseScalar(k, d).to(device) p = tuple(reversed(range(k))) aa = antisymmetrizer(b, xx) aap = antisymmetrizer(b, xx[:, p]) inv = sum(p[i] > p[j] for i in range(k) for j in range(i + 1, k)) signed_resid = aap - (-1.0 if inv % 2 else 1.0) * aa as_rows.append({"N": k, "antisym_error": antisym_error(b, xx, p), "antisym_abs_rmse": float(signed_resid.pow(2).mean().sqrt().item())}) # Symmetric Jastrow-like factor is invariant under all tested relabelings. with torch.no_grad(): jerr = relative(sj(x[:, perm]), sj(x)) # Prediction 3: paper's t*v flow gives displacement linear in t, through zero. ts = [0.0, 0.1, 0.25, 0.5, 1.0, 1.5, 2.0] disp = [] with torch.no_grad(): for t in ts: z = euler_flow(eq, x, horizon=1.0, steps=32) if t == 1.0 else None # Scaling the vector field by t is exactly equivalent to this small-step test. z = x.clone() for k in range(32): z = z + (1 / 32) * eq(z, t) disp.append(float((z - x).pow(2).sum().sqrt().item() / math.sqrt(x.numel()))) coef = np.polyfit(ts, disp, 1) pred_linear_r2 = float(1 - np.sum((np.asarray(disp) - np.polyval(coef, ts)) ** 2) / np.sum((np.asarray(disp) - np.mean(disp)) ** 2)) # Integration-depth signature: structured drift stays at numerical noise; baseline does not. depth_rows = [] for steps in [1, 2, 4, 8, 16, 32, 64]: with torch.no_grad(): zs = euler_flow(eq, x, steps=steps) zf = euler_flow(flat, x, steps=steps) xp = x[:, perm] es = relative(euler_flow(eq, xp, steps=steps), zs[:, perm]) ef = relative(euler_flow(flat, xp, steps=steps), zf[:, perm]) depth_rows.append({"steps": steps, "structured_flow_eq_error": es, "vanilla_flow_eq_error": ef}) # Secondary mini task: same target, equal small setup, unseen permutation. task_eq = EquivariantInvariantHead(d).to(device) task_flat = FlatInvariantHead(n, d).to(device) task = {"structured_rmse": task_train(task_eq, n, d), "vanilla_rmse": task_train(task_flat, n, d)} return {"device": str(device), "equivariance_sweep": eq_rows, "antisymmetry_sweep": as_rows, "symmetric_factor_error": jerr, "time_scale": {"t": ts, "displacement": disp, "slope": float(coef[0]), "intercept": float(coef[1]), "R2": pred_linear_r2}, "depth_sweep": depth_rows, "task": task} class EquivariantInvariantHead(nn.Module): def __init__(self, d, hidden=32): super().__init__(); self.item = nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, 1)) def forward(self, x): return self.item(x).mean(1) class FlatInvariantHead(nn.Module): def __init__(self, n, d, hidden=32): super().__init__(); self.net = nn.Sequential(nn.Linear(n*d, hidden), nn.Tanh(), nn.Linear(hidden, 1)) def forward(self, x): return self.net(x.reshape(x.shape[0], -1)) if __name__ == "__main__": result = safe_device_call(run) Path("results.json").write_text(json.dumps(result, indent=2)) print(json.dumps(result, indent=2))