Encoder-reset recursive world-model training / experiment.py
Beats tuned baseline
1import json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED=1077
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8torch.set_num_threads(4)
9
10def contraction_sweep():
11 rows=[]; steps=20
12 for rho in [0.5,0.8,0.9,0.95,1.0,1.05,1.2]:
13 vals=np.array([rho**j for j in range(steps+1)],float)
14 slope=float(np.polyfit(np.arange(1,steps+1),np.log(vals[1:]),1)[0])
15 rows.append(dict(rho=rho,predicted_slope=math.log(rho),observed_slope=slope,
16 predicted_ratio=rho**steps,observed_ratio=float(vals[-1]),
17 decays=bool(rho<1)))
18 return rows
19
20def make_data(T=2400):
21 u=np.random.uniform(-1,1,(T,1)).astype('float32'); x=np.zeros((T+1,2),dtype='float32')
22 for k in range(T):
23 x[k+1,0]=.91*x[k,0]+.16*np.tanh(x[k,1])+.20*u[k,0]
24 x[k+1,1]=.86*x[k,1]-.13*np.tanh(x[k,0])+.12*u[k,0]
25 y=x[:-1]+.015*np.random.randn(T,2).astype('float32')
26 return torch.tensor(u),torch.tensor(y)
27
28class WorldModel(nn.Module):
29 def __init__(self,h=16):
30 super().__init__(); self.enc=nn.Sequential(nn.Linear(3,h),nn.Tanh(),nn.Linear(h,h),nn.Tanh())
31 self.trans=nn.Sequential(nn.Linear(h+1,h),nn.Tanh(),nn.Linear(h,h),nn.Tanh())
32 self.head=nn.Linear(h,2)
33 def rollout_reset(self,ctx_u,ctx_y,u):
34 z=self.enc(torch.cat([ctx_u,ctx_y],-1).mean(0)); out=[]
35 for j in range(u.shape[0]): out.append(self.head(z)); z=self.trans(torch.cat([z,u[j]],-1))
36 return torch.stack(out)
37 def rollout_carry(self,z,u):
38 out=[]
39 for j in range(u.shape[0]): out.append(self.head(z)); z=self.trans(torch.cat([z,u[j]],-1))
40 return torch.stack(out),z
41
42def train(mode,N=32,epochs=7):
43 u,y=make_data(); model=WorldModel(); opt=torch.optim.Adam(model.parameters(),lr=3e-3)
44 losses=[]; z=None; L=4
45 # Non-overlapping online batches; reset uses preceding context and detaches all batch boundaries.
46 for ep in range(epochs):
47 for start in range(L,start+1 if False else len(y)-N+1,N):
48 if start+N>len(y): break
49 if mode=='reset':
50 cu=u[start-L:start]; cy=y[start-L:start]
51 pred=model.rollout_reset(cu,cy,u[start:start+N])
52 else:
53 if z is None: z=model.enc(torch.cat([u[start-L:start],y[start-L:start]],-1).mean(0))
54 pred,z=model.rollout_carry(z,u[start:start+N]); z=z.detach()
55 loss=((pred-y[start:start+N])**2).mean(); opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5); opt.step()
56 losses.append(float(loss))
57 return float(np.mean(losses[-20:])),float(np.mean(losses[:20])),losses
58
59def main():
60 contraction=contraction_sweep(); results={}
61 for N in [16,64]:
62 results[N]={m:train(m,N)[:2] for m in ['reset','carry']}
63 print(json.dumps({'contraction':contraction,'training':results},indent=2))
64 with open('results.json','w') as f: json.dump({'contraction':contraction,'training':results},f,indent=2)
65if __name__=='__main__': main()