gary-mesh / demo.py
gary23w's picture
gary-mesh: MeshNCA-style graph CA prototype (2,364-param shared rule, 64->65536 cells fixed params) + design paper
668c429 verified
Raw
History Blame Contribute Delete
1.17 kB
"""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")