import json, math, random from pathlib import Path import numpy as np import torch from torch import nn SEED = 7 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) device = "cuda" if torch.cuda.is_available() else "cpu" def math_check(): torch.manual_seed(SEED) vals = [] norm_err = [] for _ in range(20): B = torch.randn(5, 5) A = B - B.T dt = 0.2 I = torch.eye(5) R = torch.linalg.solve(I - .5*dt*A, I + .5*dt*A) vals.append(float(torch.linalg.matrix_norm(A + A.T))) v = torch.randn(5) norm_err.append(abs(float(torch.linalg.vector_norm(R@v) - torch.linalg.vector_norm(v)))) # Compare against an unconstrained Euler transport control. B = torch.randn(5, 5); A = B - B.T I = torch.eye(5); dt = .2; v = torch.randn(5) cay = torch.linalg.solve(I-.5*dt*A, I+.5*dt*A) euler = I + dt*A cay_err = abs(torch.linalg.vector_norm(cay@v)-torch.linalg.vector_norm(v)).item() euler_err = abs(torch.linalg.vector_norm(euler@v)-torch.linalg.vector_norm(v)).item() return {"max_skew_residual": max(vals), "max_cayley_norm_error": max(norm_err), "single_step_cayley_error": cay_err, "single_step_euler_control_error": euler_err} def make_data(num, T, seed): rng = np.random.default_rng(seed) X = np.zeros((num,T,5,4), dtype=np.float32) # x,y,vx,vy y = rng.integers(0,2,size=num) for b in range(num): sign = 1 if y[b] else -1 phase = rng.uniform(0, 2*np.pi) radius = rng.uniform(.65,.9) center = rng.normal(0,.08,2) # two objects perform a partial exchange; direction is the label for t in range(T): a = phase + sign * (0.95*np.pi) * t/(T-1) da = sign * (0.95*np.pi)/(T-1) p1 = center + radius*np.array([np.cos(a), np.sin(a)]) p2 = center + radius*np.array([np.cos(a+np.pi), np.sin(a+np.pi)]) v1 = radius*da*np.array([-np.sin(a),np.cos(a)]) v2 = radius*da*np.array([-np.sin(a+np.pi),np.cos(a+np.pi)]) pts=[np.r_[p1,v1],np.r_[p2,v2]] # stationary distractors make this a 5-object set 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]) pts=np.asarray(pts,dtype=np.float32) X[b,t]=pts[rng.permutation(5)] # independent hidden permutation per frame return torch.tensor(X), torch.tensor(y,dtype=torch.long) class DeepSetGRU(nn.Module): def __init__(self, d=4, hidden=32): super().__init__(); self.phi=nn.Sequential(nn.Linear(d,24),nn.Tanh(),nn.Linear(24,24),nn.Tanh()) self.rnn=nn.GRU(24,hidden,batch_first=True); self.out=nn.Linear(hidden,2) def forward(self,x): z=self.phi(x).mean(2); h,_=self.rnn(z); return self.out(h[:,-1]) class BraidState(nn.Module): def __init__(self, d=4, k=8, hidden=32): super().__init__(); self.k=k self.item=nn.Sequential(nn.Linear(d,24),nn.Tanh(),nn.Linear(24,24),nn.Tanh()) self.pair=nn.Sequential(nn.Linear(52,32),nn.Tanh(),nn.Linear(32,32),nn.Tanh()) self.amat=nn.Linear(32,k*k); self.weight=nn.Linear(32,1) self.zproj=nn.Linear(24,hidden); self.innov=nn.Sequential(nn.Linear(hidden+k,hidden),nn.Tanh(),nn.Linear(hidden, k)) self.out=nn.Sequential(nn.Linear(hidden+k,32),nn.Tanh(),nn.Linear(32,2)) def forward(self,x): B,T,N,D=x.shape; h=x.new_zeros(B,self.k) for t in range(T): xt=x[:,t]; item=self.item(xt); z=item.mean(1); A=x.new_zeros(B,self.k,self.k) for i in range(N): for j in range(i+1,N): # embeddings plus relative position and velocity 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)) raw=self.amat(q).view(B,self.k,self.k); skew=raw-raw.transpose(1,2) w=torch.sigmoid(self.weight(q)).view(B,1,1); A=A+w*skew I=torch.eye(self.k,device=x.device).expand(B,-1,-1) R=torch.linalg.solve(I-.5*A,I+.5*A) zp=torch.tanh(self.zproj(z)); h=(R@h.unsqueeze(-1)).squeeze(-1); h=h+self.innov(torch.cat([zp,h],-1)) return self.out(torch.cat([zp,h],-1)) def train(model, train, test, epochs=40): model.to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3); lossfn=nn.CrossEntropyLoss() tx,ty=train; vx,vy=test; tx,ty=tx.to(device),ty.to(device); vx,vy=vx.to(device),vy.to(device) best=0 for ep in range(epochs): model.train(); opt.zero_grad(); loss=lossfn(model(tx),ty); loss.backward(); opt.step() model.eval() with torch.no_grad(): acc=(model(vx).argmax(1)==vy).float().mean().item() best=max(best,acc) return best def permutation_check(): torch.manual_seed(SEED) x,_=make_data(4,6,99) m=BraidState().to(device).eval() x=x.to(device) perms=torch.stack([torch.randperm(5) for _ in range(4)]) xp=torch.stack([x[b,:,perms[b]] for b in range(4)]) with torch.no_grad(): a=m(x); b=m(xp) return float((a-b).abs().max().item()) def main(): check=math_check() train_data=make_data(640,12,11); test_data=make_data(320,12,12); long_data=make_data(320,24,13) results={"device":device,"math_check":check,"max_permutation_logit_difference":permutation_check()} # fixed seeds and identical data; long test checks extrapolation in time for name, cls in [("deepset_gru",DeepSetGRU),("braid_monodromy",BraidState)]: torch.manual_seed(SEED) m=cls(); acc=train(m,train_data,test_data) m.eval(); lx,ly=long_data; lx,ly=lx.to(device),ly.to(device) with torch.no_grad(): longacc=(m(lx).argmax(1)==ly).float().mean().item() results[name]={"params":sum(p.numel() for p in m.parameters()),"accuracy_T12":acc,"accuracy_T24":longacc} Path("results.json").write_text(json.dumps(results,indent=2)) print(json.dumps(results,indent=2)) if __name__=='__main__': try: main() except Exception as e: if device=='cuda': print('CUDA failed, rerun on CPU:',repr(e)); device='cpu'; torch.cuda.empty_cache(); main() else: raise