Braid-Monodromy Set State / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED = 7
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10device = "cuda" if torch.cuda.is_available() else "cpu"
11
12
13def math_check():
14 torch.manual_seed(SEED)
15 vals = []
16 norm_err = []
17 for _ in range(20):
18 B = torch.randn(5, 5)
19 A = B - B.T
20 dt = 0.2
21 I = torch.eye(5)
22 R = torch.linalg.solve(I - .5*dt*A, I + .5*dt*A)
23 vals.append(float(torch.linalg.matrix_norm(A + A.T)))
24 v = torch.randn(5)
25 norm_err.append(abs(float(torch.linalg.vector_norm(R@v) - torch.linalg.vector_norm(v))))
26 # Compare against an unconstrained Euler transport control.
27 B = torch.randn(5, 5); A = B - B.T
28 I = torch.eye(5); dt = .2; v = torch.randn(5)
29 cay = torch.linalg.solve(I-.5*dt*A, I+.5*dt*A)
30 euler = I + dt*A
31 cay_err = abs(torch.linalg.vector_norm(cay@v)-torch.linalg.vector_norm(v)).item()
32 euler_err = abs(torch.linalg.vector_norm(euler@v)-torch.linalg.vector_norm(v)).item()
33 return {"max_skew_residual": max(vals), "max_cayley_norm_error": max(norm_err),
34 "single_step_cayley_error": cay_err, "single_step_euler_control_error": euler_err}
35
36
37def make_data(num, T, seed):
38 rng = np.random.default_rng(seed)
39 X = np.zeros((num,T,5,4), dtype=np.float32) # x,y,vx,vy
40 y = rng.integers(0,2,size=num)
41 for b in range(num):
42 sign = 1 if y[b] else -1
43 phase = rng.uniform(0, 2*np.pi)
44 radius = rng.uniform(.65,.9)
45 center = rng.normal(0,.08,2)
46 # two objects perform a partial exchange; direction is the label
47 for t in range(T):
48 a = phase + sign * (0.95*np.pi) * t/(T-1)
49 da = sign * (0.95*np.pi)/(T-1)
50 p1 = center + radius*np.array([np.cos(a), np.sin(a)])
51 p2 = center + radius*np.array([np.cos(a+np.pi), np.sin(a+np.pi)])
52 v1 = radius*da*np.array([-np.sin(a),np.cos(a)])
53 v2 = radius*da*np.array([-np.sin(a+np.pi),np.cos(a+np.pi)])
54 pts=[np.r_[p1,v1],np.r_[p2,v2]]
55 # stationary distractors make this a 5-object set
56 for c in [(-1.0,-.65),(.95,-.6),(.05,1.05)]: pts.append(np.r_[np.array(c)+rng.normal(0,.015,2),0,0])
57 pts=np.asarray(pts,dtype=np.float32)
58 X[b,t]=pts[rng.permutation(5)] # independent hidden permutation per frame
59 return torch.tensor(X), torch.tensor(y,dtype=torch.long)
60
61class DeepSetGRU(nn.Module):
62 def __init__(self, d=4, hidden=32):
63 super().__init__(); self.phi=nn.Sequential(nn.Linear(d,24),nn.Tanh(),nn.Linear(24,24),nn.Tanh())
64 self.rnn=nn.GRU(24,hidden,batch_first=True); self.out=nn.Linear(hidden,2)
65 def forward(self,x):
66 z=self.phi(x).mean(2); h,_=self.rnn(z); return self.out(h[:,-1])
67
68class BraidState(nn.Module):
69 def __init__(self, d=4, k=8, hidden=32):
70 super().__init__(); self.k=k
71 self.item=nn.Sequential(nn.Linear(d,24),nn.Tanh(),nn.Linear(24,24),nn.Tanh())
72 self.pair=nn.Sequential(nn.Linear(52,32),nn.Tanh(),nn.Linear(32,32),nn.Tanh())
73 self.amat=nn.Linear(32,k*k); self.weight=nn.Linear(32,1)
74 self.zproj=nn.Linear(24,hidden); self.innov=nn.Sequential(nn.Linear(hidden+k,hidden),nn.Tanh(),nn.Linear(hidden, k))
75 self.out=nn.Sequential(nn.Linear(hidden+k,32),nn.Tanh(),nn.Linear(32,2))
76 def forward(self,x):
77 B,T,N,D=x.shape; h=x.new_zeros(B,self.k)
78 for t in range(T):
79 xt=x[:,t]; item=self.item(xt); z=item.mean(1); A=x.new_zeros(B,self.k,self.k)
80 for i in range(N):
81 for j in range(i+1,N):
82 # embeddings plus relative position and velocity
83 q=self.pair(torch.cat([item[:,i]+item[:,j], torch.abs(item[:,i]-item[:,j]), torch.abs(xt[:,i,0:2]-xt[:,j,0:2]), torch.abs(xt[:,i,2:4]-xt[:,j,2:4])],-1))
84 raw=self.amat(q).view(B,self.k,self.k); skew=raw-raw.transpose(1,2)
85 w=torch.sigmoid(self.weight(q)).view(B,1,1); A=A+w*skew
86 I=torch.eye(self.k,device=x.device).expand(B,-1,-1)
87 R=torch.linalg.solve(I-.5*A,I+.5*A)
88 zp=torch.tanh(self.zproj(z)); h=(R@h.unsqueeze(-1)).squeeze(-1); h=h+self.innov(torch.cat([zp,h],-1))
89 return self.out(torch.cat([zp,h],-1))
90
91def train(model, train, test, epochs=40):
92 model.to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3); lossfn=nn.CrossEntropyLoss()
93 tx,ty=train; vx,vy=test; tx,ty=tx.to(device),ty.to(device); vx,vy=vx.to(device),vy.to(device)
94 best=0
95 for ep in range(epochs):
96 model.train(); opt.zero_grad(); loss=lossfn(model(tx),ty); loss.backward(); opt.step()
97 model.eval()
98 with torch.no_grad(): acc=(model(vx).argmax(1)==vy).float().mean().item()
99 best=max(best,acc)
100 return best
101
102def permutation_check():
103 torch.manual_seed(SEED)
104 x,_=make_data(4,6,99)
105 m=BraidState().to(device).eval()
106 x=x.to(device)
107 perms=torch.stack([torch.randperm(5) for _ in range(4)])
108 xp=torch.stack([x[b,:,perms[b]] for b in range(4)])
109 with torch.no_grad():
110 a=m(x); b=m(xp)
111 return float((a-b).abs().max().item())
112
113def main():
114 check=math_check()
115 train_data=make_data(640,12,11); test_data=make_data(320,12,12); long_data=make_data(320,24,13)
116 results={"device":device,"math_check":check,"max_permutation_logit_difference":permutation_check()}
117 # fixed seeds and identical data; long test checks extrapolation in time
118 for name, cls in [("deepset_gru",DeepSetGRU),("braid_monodromy",BraidState)]:
119 torch.manual_seed(SEED)
120 m=cls(); acc=train(m,train_data,test_data)
121 m.eval(); lx,ly=long_data; lx,ly=lx.to(device),ly.to(device)
122 with torch.no_grad(): longacc=(m(lx).argmax(1)==ly).float().mean().item()
123 results[name]={"params":sum(p.numel() for p in m.parameters()),"accuracy_T12":acc,"accuracy_T24":longacc}
124 Path("results.json").write_text(json.dumps(results,indent=2))
125 print(json.dumps(results,indent=2))
126
127if __name__=='__main__':
128 try: main()
129 except Exception as e:
130 if device=='cuda':
131 print('CUDA failed, rerun on CPU:',repr(e)); device='cpu'; torch.cuda.empty_cache(); main()
132 else: raise