import json import numpy as np import torch from torch import nn SEED=965 def pk_loss(z,x,eps=1e-6): u=z/(torch.linalg.vector_norm(z,dim=-1,keepdim=True)+eps) r2=(x*x).sum().clamp(max=1-eps) den=1+r2-2*(u*x).sum(-1) return (-torch.log((1-r2)/den.clamp_min(eps))).mean() class RNN(nn.Module): def __init__(self): super().__init__(); self.W=nn.Linear(3,2); self.U=nn.Linear(2,2,bias=False); self.out=nn.Linear(2,1) def forward(self,s): h=torch.zeros(s.shape[0],2,device=s.device); hs=[]; ys=[] for t in range(s.shape[1]): h=torch.tanh(self.W(torch.cat([s[:,t],h],-1))); hs.append(h); ys.append(self.out(h)) return torch.stack(ys,1),torch.stack(hs,1) def run(reg): torch.manual_seed(SEED); np.random.seed(SEED) dev='cuda' if torch.cuda.is_available() else 'cpu' try: device=torch.device(dev); model=RNN().to(device) g=torch.Generator(device=device); g.manual_seed(SEED) ntr,nva,T=192,64,25 t=torch.arange(ntr+nva+T+1,device=device).float() base=torch.sin(.18*t)+.15*torch.sin(.73*t) noise=.08*torch.randn(ntr+nva+T+1,generator=g,device=device) seq=(base+noise).unfold(0,T+1,1)[:ntr+nva] x=seq[:,:T].unsqueeze(-1); y=seq[:,1:T+1].unsqueeze(-1) opt=torch.optim.Adam(model.parameters(),lr=.015) losses=[] # fixed affine random-map composition estimate, as the implementation plan specifies q=.72; target=torch.tensor([.48,.20],device=device); probes=torch.tensor([[-.5,.1],[.1,-.4],[.4,.2]],device=device) xhat=target+(q**8)*(probes.mean(0)-target) for step in range(180): idx=torch.randperm(ntr,generator=g,device=device)[:64] pred,h=model(x[idx]); task=((pred-y[idx])**2).mean(); loss=task if reg: loss=loss+.01*pk_loss(h[:,-1],xhat.detach()) opt.zero_grad(); loss.backward(); opt.step(); losses.append(float(task.detach())) with torch.no_grad(): pred,h=model(x[ntr:]); mse=float(((pred-y[ntr:])**2).mean()) z=h[:,-1]; pair=torch.pdist(z).mean().item(); pl=pk_loss(z,xhat).item() return {'val_mse':mse,'pairwise_distance':pair,'pk_nll':pl,'last_train_task':losses[-1]} except Exception as e: if dev=='cuda': torch.cuda.empty_cache(); return run_cpu(reg) raise def run_cpu(reg): old=torch.cuda.is_available torch.cuda.is_available=lambda:False try:return run(reg) finally:torch.cuda.is_available=lambda:old if __name__=='__main__': out={'baseline':run(False),'poisson_kernel':run(True)} with open('mini_results.json','w') as f: json.dump(out,f,indent=2) print(json.dumps(out,indent=2))