| """Reproduce the paper's headline: one shared rule, constant params, any number of cells. |
| Run: python demo.py""" |
| import numpy as np, time |
| from gary_mesh import GaryMesh, ring_graph, harmonic_weights, Adam |
| def target_for(N): return 0.8*np.sin(2*(2*np.pi*np.arange(N)/N)) |
| Ntr,T=64,20 |
| pos,nbr=ring_graph(Ntr); Wh=harmonic_weights(pos,nbr); tgt=target_for(Ntr) |
| m=GaryMesh(C=12,H=48,seed=2); opt=Adam(m.params(),lr=3e-3) |
| for it in range(1,801): |
| rng=np.random.default_rng(it); S,c=m.rollout(Ntr,nbr,Wh,T,rng,train=True) |
| l,g=m.backward(S,tgt,c,nbr,Wh); gn=np.sqrt(sum((v**2).sum() for v in g.values())) |
| if gn>1.0: |
| for k in g: g[k]*=1.0/gn |
| opt.step(m.params(),g) |
| corr=abs(np.corrcoef(m.rollout(Ntr,nbr,Wh,T,np.random.default_rng(9))[0][:,0],tgt)[0,1]) |
| print(f"trained rule: {m.nparams()} params | train-ring correlation {corr:.3f}\n") |
| print("same rule, growing mesh, parameters FIXED:") |
| for N in [64,256,1024,4096,16384,65536]: |
| p,nb=ring_graph(N); wh=harmonic_weights(p,nb) |
| rng=np.random.default_rng(7); t0=time.time(); m.rollout(N,nb,wh,T,rng); dt=time.time()-t0 |
| print(f" {N:6d} cells | {m.nparams()} params | {1000*dt/N:.4f} ms/cell") |
|
|