Symmetry-Preserving Flow Layer / symmetry_flow_experiment.py
Beats tuned baseline
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))