Chern-Gap Monitor for Finite-Horizon Collapse / train_experiment.py
Mechanism confirmed, baseline not beaten
1import json, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED=23
7def seed(s):
8 random.seed(s); np.random.seed(s); torch.manual_seed(s)
9
10def data(n,s):
11 r=np.random.default_rng(s); x=r.integers(0,2,(n,8)); y=(x[:,0]!=x[:,-1]).astype('int64')
12 return torch.tensor(x),torch.tensor(y)
13
14class Net(nn.Module):
15 def __init__(self):
16 super().__init__(); self.r=nn.GRU(4,24,batch_first=True); self.c=nn.Linear(24,2); self.s=nn.Linear(24,3)
17 def forward(self,x,p):
18 b,t=x.shape; e=torch.nn.functional.one_hot(x,2).float(); q=p[:,None,:].expand(b,t,2)
19 h=self.r(torch.cat((e,q),-1))[1][0]; return self.c(h),torch.tanh(self.s(h))
20
21def phases(K,d):
22 z=torch.linspace(0,2*np.pi,K+1,device=d)[:-1]; a,b=torch.meshgrid(z,z,indexing='ij'); return torch.stack((a.flatten(),b.flatten()),-1)
23
24def probe(model,x,p):
25 K=p.shape[0]; xx=x[:,None,:].expand(-1,K,-1).reshape(-1,x.shape[1]); pp=p[None,:,:].expand(x.shape[0],-1,-1).reshape(-1,2)
26 _,s=model(xx,pp); m=s.reshape(x.shape[0],K,3).mean(0); return torch.linalg.vector_norm(m,dim=1).min(),m
27
28def run(reg, force_cpu=False):
29 seed(SEED); dev='cpu' if force_cpu else ('cuda' if torch.cuda.is_available() else 'cpu')
30 try:
31 tr,yt=data(512,4); va,yv=data(256,5); tr,yt,va,yv=[z.to(dev) for z in (tr,yt,va,yv)]
32 model=Net().to(dev); opt=torch.optim.Adam(model.parameters(),lr=3e-3); loss=nn.CrossEntropyLoss(); pg=phases(7,dev); probe_x=tr[:64]
33 hist=[]
34 for step in range(180):
35 ix=torch.randint(0,len(tr),(64,),device=dev); logits,_=model(tr[ix],torch.zeros(64,2,device=dev)); L=loss(logits,yt[ix])
36 if reg:
37 g,_=probe(model,probe_x,pg); L=L+0.3*torch.relu(torch.tensor(.12,device=dev)-g)**2
38 opt.zero_grad(); L.backward(); opt.step()
39 if step%30==0:
40 with torch.no_grad():
41 g,m=probe(model,probe_x,pg); pred=model(va,torch.zeros(len(va),2,device=dev))[0].argmax(1); acc=(pred==yv).float().mean()
42 hist.append((step,float(loss(model(va,torch.zeros(len(va),2,device=dev))[0],yv)),float(acc),float(g)))
43 return hist
44 except Exception:
45 if dev=='cuda':
46 torch.cuda.empty_cache(); return run(reg, True)
47 raise
48
49if __name__=='__main__':
50 b=run(False); r=run(True)
51 print(json.dumps({'baseline':b,'gap_regularized':r,'final_baseline':b[-1],'final_gap_regularized':r[-1]},indent=2))