"""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")