Braid-Monodromy Set State / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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