Encoder-reset recursive world-model training / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 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()