import json, math, random import numpy as np import torch from torch import nn SEED=1077 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) def contraction_sweep(): rows=[]; steps=20 for rho in [0.5,0.8,0.9,0.95,1.0,1.05,1.2]: vals=np.array([rho**j for j in range(steps+1)],float) slope=float(np.polyfit(np.arange(1,steps+1),np.log(vals[1:]),1)[0]) rows.append(dict(rho=rho,predicted_slope=math.log(rho),observed_slope=slope, predicted_ratio=rho**steps,observed_ratio=float(vals[-1]), decays=bool(rho<1))) return rows def make_data(T=2400): u=np.random.uniform(-1,1,(T,1)).astype('float32'); x=np.zeros((T+1,2),dtype='float32') for k in range(T): x[k+1,0]=.91*x[k,0]+.16*np.tanh(x[k,1])+.20*u[k,0] x[k+1,1]=.86*x[k,1]-.13*np.tanh(x[k,0])+.12*u[k,0] y=x[:-1]+.015*np.random.randn(T,2).astype('float32') return torch.tensor(u),torch.tensor(y) class WorldModel(nn.Module): def __init__(self,h=16): super().__init__(); self.enc=nn.Sequential(nn.Linear(3,h),nn.Tanh(),nn.Linear(h,h),nn.Tanh()) self.trans=nn.Sequential(nn.Linear(h+1,h),nn.Tanh(),nn.Linear(h,h),nn.Tanh()) self.head=nn.Linear(h,2) def rollout_reset(self,ctx_u,ctx_y,u): z=self.enc(torch.cat([ctx_u,ctx_y],-1).mean(0)); out=[] for j in range(u.shape[0]): out.append(self.head(z)); z=self.trans(torch.cat([z,u[j]],-1)) return torch.stack(out) def rollout_carry(self,z,u): out=[] for j in range(u.shape[0]): out.append(self.head(z)); z=self.trans(torch.cat([z,u[j]],-1)) return torch.stack(out),z def train(mode,N=32,epochs=7): u,y=make_data(); model=WorldModel(); opt=torch.optim.Adam(model.parameters(),lr=3e-3) losses=[]; z=None; L=4 # Non-overlapping online batches; reset uses preceding context and detaches all batch boundaries. for ep in range(epochs): for start in range(L,start+1 if False else len(y)-N+1,N): if start+N>len(y): break if mode=='reset': cu=u[start-L:start]; cy=y[start-L:start] pred=model.rollout_reset(cu,cy,u[start:start+N]) else: if z is None: z=model.enc(torch.cat([u[start-L:start],y[start-L:start]],-1).mean(0)) pred,z=model.rollout_carry(z,u[start:start+N]); z=z.detach() loss=((pred-y[start:start+N])**2).mean(); opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5); opt.step() losses.append(float(loss)) return float(np.mean(losses[-20:])),float(np.mean(losses[:20])),losses def main(): contraction=contraction_sweep(); results={} for N in [16,64]: results[N]={m:train(m,N)[:2] for m in ['reset','carry']} print(json.dumps({'contraction':contraction,'training':results},indent=2)) with open('results.json','w') as f: json.dump({'contraction':contraction,'training':results},f,indent=2) if __name__=='__main__': main()