Symmetry-Preserving Flow Layer / symmetry_flow_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import itertools
  2import json
  3import math
  4import random
  5from pathlib import Path
  6
  7import numpy as np
  8import torch
  9from torch import nn
 10
 11SEED = 1612
 12random.seed(SEED)
 13np.random.seed(SEED)
 14torch.manual_seed(SEED)
 15torch.set_default_dtype(torch.float64)
 16
 17try:
 18    DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 19except Exception:
 20    DEVICE = torch.device("cpu")
 21
 22# This experiment is intentionally small; on CUDA, an allocation/runtime error
 23# falls back to CPU as required by the experiment protocol.
 24def safe_device_call(fn):
 25    global DEVICE
 26    try:
 27        return fn(DEVICE)
 28    except Exception:
 29        DEVICE = torch.device("cpu")
 30        torch.cuda.empty_cache() if torch.cuda.is_available() else None
 31        return fn(DEVICE)
 32
 33
 34def perms_and_signs(n, device):
 35    ps, ss = [], []
 36    for p in itertools.permutations(range(n)):
 37        inv = sum(p[i] > p[j] for i in range(n) for j in range(i + 1, n))
 38        ps.append(p)
 39        ss.append(-1.0 if inv % 2 else 1.0)
 40    return ps, torch.tensor(ss, device=device)
 41
 42
 43class EquivariantVectorField(nn.Module):
 44    """v_i = MLP([x_i, mean(x), mean_j phi(x_i-x_j)])."""
 45    def __init__(self, d, hidden=32):
 46        super().__init__()
 47        self.pair = nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, hidden), nn.Tanh())
 48        self.out = nn.Sequential(nn.Linear(2 * d + hidden, hidden), nn.Tanh(), nn.Linear(hidden, d))
 49
 50    def forward(self, x, t=1.0):
 51        # x: [batch, N, d]; sum/mean over the unordered index set
 52        dif = x[:, :, None, :] - x[:, None, :, :]
 53        messages = self.pair(dif).mean(dim=2)
 54        pooled = x.mean(dim=1, keepdim=True).expand_as(x)
 55        return float(t) * self.out(torch.cat([x, pooled, messages], dim=-1))
 56
 57
 58class FlatVectorField(nn.Module):
 59    """Label-sensitive baseline with the same general hidden width."""
 60    def __init__(self, n, d, hidden=32):
 61        super().__init__()
 62        self.net = nn.Sequential(nn.Linear(n * d + 1, hidden), nn.Tanh(), nn.Linear(hidden, hidden), nn.Tanh(), nn.Linear(hidden, n * d))
 63
 64    def forward(self, x, t=1.0):
 65        b = x.shape[0]
 66        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)
 67
 68
 69def euler_flow(field, x, horizon=1.0, steps=16):
 70    z = x.clone()
 71    dt = horizon / steps
 72    for k in range(steps):
 73        z = z + dt * field(z, (k + 0.5) * dt)
 74    return z
 75
 76
 77def relative(a, b, eps=1e-12):
 78    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()
 79
 80
 81def equivariance_error(field, x, perm):
 82    xp = x[:, perm]
 83    return relative(field(xp, .7), field(x, .7)[:, perm])
 84
 85
 86def antisymmetrizer(base, x):
 87    # base accepts [B,N,D] and returns [B] scalar values
 88    n = x.shape[1]
 89    ps, signs = perms_and_signs(n, x.device)
 90    vals = torch.stack([base(x[:, p]) for p in ps], dim=1)
 91    return (vals * signs[None, :]).mean(dim=1)
 92
 93
 94def antisym_error(base, x, perm):
 95    a = antisymmetrizer(base, x)
 96    ap = antisymmetrizer(base, x[:, perm])
 97    inv = sum(perm[i] > perm[j] for i in range(len(perm)) for j in range(i + 1, len(perm)))
 98    sign = -1.0 if inv % 2 else 1.0
 99    return (torch.abs(ap - sign * a) / (torch.abs(a) + 1e-12)).mean().item()
100
101
102class BaseScalar(nn.Module):
103    def __init__(self, n, d, hidden=32):
104        super().__init__()
105        self.net = nn.Sequential(nn.Linear(n * d, hidden), nn.Tanh(), nn.Linear(hidden, hidden), nn.Tanh(), nn.Linear(hidden, 1))
106
107    def forward(self, x):
108        return self.net(x.reshape(x.shape[0], -1)).squeeze(-1)
109
110
111class SymmetricJ(nn.Module):
112    def __init__(self, d, hidden=24):
113        super().__init__()
114        self.net = nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, 1))
115
116    def forward(self, x):
117        return self.net(x).squeeze(-1).mean(dim=1)
118
119
120def task_train(model, n, d, steps=350):
121    opt = torch.optim.Adam(model.parameters(), lr=3e-3)
122    for _ in range(steps):
123        x = torch.randn(64, n, d, device=next(model.parameters()).device)
124        # invariant target: mean squared radius plus a smooth interaction term
125        target = x.pow(2).sum(-1).mean(-1) + 0.15 * torch.tanh((x[:, :, None] * x[:, None, :]).sum(-1).mean((1, 2)))
126        pred = model(x).squeeze(-1)
127        loss = (pred - target).pow(2).mean()
128        opt.zero_grad(); loss.backward(); opt.step()
129    with torch.no_grad():
130        x = torch.randn(256, n, d, device=next(model.parameters()).device)
131        target = x.pow(2).sum(-1).mean(-1) + 0.15 * torch.tanh((x[:, :, None] * x[:, None, :]).sum(-1).mean((1, 2)))
132        p = torch.randperm(n, device=x.device)
133        # Generalization to a fresh unseen permutation
134        err = ((model(x[:, p]).squeeze(-1) - target) ** 2).mean().sqrt().item()
135    return err
136
137
138def run(device):
139    torch.manual_seed(SEED)
140    n, d = 4, 2
141    x = torch.randn(48, n, d, device=device)
142    perm = (1, 3, 0, 2)
143    eq = EquivariantVectorField(d).to(device)
144    flat = FlatVectorField(n, d).to(device)
145    base = BaseScalar(n, d).to(device)
146    sj = SymmetricJ(d).to(device)
147
148    # Prediction 1: exact equivariance independent of flow scale and integration depth.
149    scales = [0.0, 0.25, 0.5, 1.0, 2.0]
150    eq_rows = []
151    for scale in scales:
152        e1 = equivariance_error(eq, x * scale + 0.1, perm)
153        e2 = equivariance_error(flat, x * scale + 0.1, perm)
154        eq_rows.append({"scale": scale, "structured": e1, "vanilla": e2})
155
156    # Prediction 2: explicit antisymmetrization and symmetric factor are exact.
157    as_rows = []
158    for k in [1, 2, 3, 4, 5, 6]:
159        xx = torch.randn(24, k, d, device=device)
160        b = BaseScalar(k, d).to(device)
161        p = tuple(reversed(range(k)))
162        aa = antisymmetrizer(b, xx)
163        aap = antisymmetrizer(b, xx[:, p])
164        inv = sum(p[i] > p[j] for i in range(k) for j in range(i + 1, k))
165        signed_resid = aap - (-1.0 if inv % 2 else 1.0) * aa
166        as_rows.append({"N": k, "antisym_error": antisym_error(b, xx, p),
167                        "antisym_abs_rmse": float(signed_resid.pow(2).mean().sqrt().item())})
168
169    # Symmetric Jastrow-like factor is invariant under all tested relabelings.
170    with torch.no_grad():
171        jerr = relative(sj(x[:, perm]), sj(x))
172
173    # Prediction 3: paper's t*v flow gives displacement linear in t, through zero.
174    ts = [0.0, 0.1, 0.25, 0.5, 1.0, 1.5, 2.0]
175    disp = []
176    with torch.no_grad():
177        for t in ts:
178            z = euler_flow(eq, x, horizon=1.0, steps=32) if t == 1.0 else None
179            # Scaling the vector field by t is exactly equivalent to this small-step test.
180            z = x.clone()
181            for k in range(32):
182                z = z + (1 / 32) * eq(z, t)
183            disp.append(float((z - x).pow(2).sum().sqrt().item() / math.sqrt(x.numel())))
184    coef = np.polyfit(ts, disp, 1)
185    pred_linear_r2 = float(1 - np.sum((np.asarray(disp) - np.polyval(coef, ts)) ** 2) / np.sum((np.asarray(disp) - np.mean(disp)) ** 2))
186
187    # Integration-depth signature: structured drift stays at numerical noise; baseline does not.
188    depth_rows = []
189    for steps in [1, 2, 4, 8, 16, 32, 64]:
190        with torch.no_grad():
191            zs = euler_flow(eq, x, steps=steps)
192            zf = euler_flow(flat, x, steps=steps)
193            xp = x[:, perm]
194            es = relative(euler_flow(eq, xp, steps=steps), zs[:, perm])
195            ef = relative(euler_flow(flat, xp, steps=steps), zf[:, perm])
196        depth_rows.append({"steps": steps, "structured_flow_eq_error": es, "vanilla_flow_eq_error": ef})
197
198    # Secondary mini task: same target, equal small setup, unseen permutation.
199    task_eq = EquivariantInvariantHead(d).to(device)
200    task_flat = FlatInvariantHead(n, d).to(device)
201    task = {"structured_rmse": task_train(task_eq, n, d), "vanilla_rmse": task_train(task_flat, n, d)}
202    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}
203
204
205class EquivariantInvariantHead(nn.Module):
206    def __init__(self, d, hidden=32):
207        super().__init__(); self.item = nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, 1))
208    def forward(self, x): return self.item(x).mean(1)
209
210class FlatInvariantHead(nn.Module):
211    def __init__(self, n, d, hidden=32):
212        super().__init__(); self.net = nn.Sequential(nn.Linear(n*d, hidden), nn.Tanh(), nn.Linear(hidden, 1))
213    def forward(self, x): return self.net(x.reshape(x.shape[0], -1))
214
215
216if __name__ == "__main__":
217    result = safe_device_call(run)
218    Path("results.json").write_text(json.dumps(result, indent=2))
219    print(json.dumps(result, indent=2))