import json, random import numpy as np import torch from torch import nn SEED=23 def seed(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) def data(n,s): r=np.random.default_rng(s); x=r.integers(0,2,(n,8)); y=(x[:,0]!=x[:,-1]).astype('int64') return torch.tensor(x),torch.tensor(y) class Net(nn.Module): def __init__(self): super().__init__(); self.r=nn.GRU(4,24,batch_first=True); self.c=nn.Linear(24,2); self.s=nn.Linear(24,3) def forward(self,x,p): b,t=x.shape; e=torch.nn.functional.one_hot(x,2).float(); q=p[:,None,:].expand(b,t,2) h=self.r(torch.cat((e,q),-1))[1][0]; return self.c(h),torch.tanh(self.s(h)) def phases(K,d): 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) def probe(model,x,p): 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) _,s=model(xx,pp); m=s.reshape(x.shape[0],K,3).mean(0); return torch.linalg.vector_norm(m,dim=1).min(),m def run(reg, force_cpu=False): seed(SEED); dev='cpu' if force_cpu else ('cuda' if torch.cuda.is_available() else 'cpu') try: tr,yt=data(512,4); va,yv=data(256,5); tr,yt,va,yv=[z.to(dev) for z in (tr,yt,va,yv)] model=Net().to(dev); opt=torch.optim.Adam(model.parameters(),lr=3e-3); loss=nn.CrossEntropyLoss(); pg=phases(7,dev); probe_x=tr[:64] hist=[] for step in range(180): ix=torch.randint(0,len(tr),(64,),device=dev); logits,_=model(tr[ix],torch.zeros(64,2,device=dev)); L=loss(logits,yt[ix]) if reg: g,_=probe(model,probe_x,pg); L=L+0.3*torch.relu(torch.tensor(.12,device=dev)-g)**2 opt.zero_grad(); L.backward(); opt.step() if step%30==0: with torch.no_grad(): 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() hist.append((step,float(loss(model(va,torch.zeros(len(va),2,device=dev))[0],yv)),float(acc),float(g))) return hist except Exception: if dev=='cuda': torch.cuda.empty_cache(); return run(reg, True) raise if __name__=='__main__': b=run(False); r=run(True) print(json.dumps({'baseline':b,'gap_regularized':r,'final_baseline':b[-1],'final_gap_regularized':r[-1]},indent=2))