Behavior-Gap Clustered Neural Controllers / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5
  6SEED=17
  7np.random.seed(SEED); torch.manual_seed(SEED)
  8
  9
 10def hankel_cols(seq, L):
 11    # seq: [T,d], each column is one flattened length-L window
 12    T,d=seq.shape
 13    return np.stack([seq[t:t+L].reshape(-1) for t in range(T-L+1)], axis=1)
 14
 15def behavior_matrix(A, B, C, inputs, L):
 16    # zero-state finite-horizon input/output graphs, concatenated over trials
 17    blocks=[]
 18    for u in inputs:
 19        x=np.zeros(A.shape[0]); ys=[]
 20        for t in range(len(u)):
 21            x=A@x+B@u[t]
 22            ys.append(C@x)
 23        blocks.append(np.vstack([hankel_cols(u,L), hankel_cols(np.asarray(ys),L)]))
 24    return np.concatenate(blocks, axis=1)
 25
 26def basis_and_projector(M, r):
 27    U,s,_=np.linalg.svd(M, full_matrices=False)
 28    rr=min(r, int(np.sum(s > max(M.shape)*s[0]*1e-10)), U.shape[1])
 29    rr=max(1,rr)
 30    Q=U[:,:rr]
 31    return Q, Q@Q.T, s
 32
 33def gap(P,Q):
 34    # spectral norm; symmetric difference is numerically stable here
 35    return float(np.linalg.norm(P-Q,2))
 36
 37def cluster(Qs, Ps, eps):
 38    groups=[]
 39    for i,P in enumerate(Ps):
 40        placed=False
 41        for g in groups:
 42            if gap(P, Ps[g[0]]) <= eps:
 43                g.append(i); placed=True; break
 44        if not placed: groups.append([i])
 45    return groups
 46
 47def adam_train(A0, targets, inputs, groups, steps=220, lr=.035, shared=True):
 48    # Learn each module's recurrent A from trajectory loss. Shared mode averages
 49    # first/second Adam statistics within each behavior cluster, while retaining A_i.
 50    A=[torch.tensor(a, dtype=torch.float32, requires_grad=True) for a in A0]
 51    n=len(A); m1=[torch.zeros_like(a) for a in A]; m2=[torch.zeros_like(a) for a in A]
 52    losses=[]
 53    for step in range(1,steps+1):
 54        ls=[]
 55        for i in range(n):
 56            x=torch.zeros(2); pred=[]
 57            for t in range(inputs.shape[1]):
 58                x=A[i]@x + torch.tensor(inputs[i,t],dtype=torch.float32)
 59                pred.append(x)
 60            pred=torch.stack(pred)
 61            ls.append(((pred-torch.tensor(targets[i],dtype=torch.float32))**2).mean())
 62        total=torch.stack(ls).sum(); total.backward()
 63        with torch.no_grad():
 64            for g in groups if shared else [[i] for i in range(n)]:
 65                grads=[A[i].grad for i in g]
 66                gm=torch.stack(grads).mean(0)
 67                # Shared controller state, but module-specific gradient direction/update.
 68                for i in g:
 69                    m1[i].mul_(0.9).add_(gm,alpha=.1)
 70                    m2[i].mul_(0.999).addcmul_(gm,gm,value=.001)
 71                    update=m1[i]/(1-.9**step)/(torch.sqrt(m2[i]/(1-.999**step))+1e-8)
 72                    A[i].add_(update,alpha=-lr)
 73                    A[i].grad.zero_()
 74        losses.append(float(torch.stack(ls).mean()))
 75    return losses, [a.detach().numpy() for a in A]
 76
 77def main():
 78    # 8 stable 2D systems: four near one behavior and four near another.
 79    n=8; d=2; T=28; Ls=[4,8,16]
 80    Atrue=[]
 81    for i in range(n):
 82        theta=(0.10 if i<4 else 0.72) + np.random.randn()*.025
 83        rad=(.82 if i<4 else .78)+np.random.randn()*.012
 84        Atrue.append(rad*np.array([[math.cos(theta),-math.sin(theta)],[math.sin(theta),math.cos(theta)]]))
 85    Atrue=np.array(Atrue); B=np.eye(2); C=np.eye(2)
 86    trials=[]
 87    for k in range(5): trials.append(np.random.randn(T,2)*.5)
 88    # math sanity: projector gap identity and angle behavior
 89    q,_=np.linalg.qr(np.random.randn(8,3)); r,_=np.linalg.qr(np.random.randn(8,3))
 90    P=q@q.T; R=r@r.T
 91    identity_err=abs(np.linalg.norm(P-R,2)-max(np.linalg.norm((np.eye(8)-R)@P,2),np.linalg.norm((np.eye(8)-P)@R,2)))
 92    angle=[]
 93    for th in [0,.1,.3,.7,1.2]:
 94        q1=np.array([[1.],[0.]]); q2=np.array([[math.cos(th)],[math.sin(th)]])
 95        angle.append((th,float(np.linalg.norm(q1@q1.T-q2@q2.T,2))))
 96    # behavior spaces and horizon effect
 97    gaps={}; clusters={}
 98    for L in Ls:
 99        Qs=[]; Ps=[]
100        for a in Atrue:
101            M=behavior_matrix(a,B,C,trials,L); q,p,s=basis_and_projector(M, min(12,M.shape[1]))
102            Qs.append(q); Ps.append(p)
103        gaps[str(L)] = [[round(gap(Ps[i],Ps[j]),5) for j in range(n)] for i in range(n)]
104        clusters[str(L)] = cluster(Qs,Ps,.35)
105    # Optimization uses L=8 clusters; targets are trajectories under true systems.
106    inp=np.random.randn(n,T,2)*.45
107    targets=[]
108    for i in range(n):
109        x=np.zeros(2); yy=[]
110        for t in range(T): x=Atrue[i]@x+inp[i,t]; yy.append(x.copy())
111        targets.append(np.asarray(yy))
112    targets=np.asarray(targets)
113    Ainit=Atrue+np.random.randn(*Atrue.shape)*.18
114    groups=clusters['8']
115    t0=time.perf_counter(); base,_=adam_train(Ainit,targets,inp,groups,shared=False); tb=time.perf_counter()-t0
116    t0=time.perf_counter(); idea,_=adam_train(Ainit,targets,inp,groups,shared=True); ti=time.perf_counter()-t0
117    result={'seed':SEED,'math':{'projector_identity_abs_error':identity_err,'angle_gap':angle},
118      'clusters_by_horizon':clusters,'within_cluster_max_gap':{str(L):max((gaps[str(L)][i][j] for g in range(len(clusters[str(L)]) ) for i in clusters[str(L)][g] for j in clusters[str(L)][g]),default=0) for L in Ls},
119      'optimization':{'groups_L8':groups,'baseline_independent_final_loss':base[-1],'idea_shared_moments_final_loss':idea[-1], 'baseline_start':base[0],'idea_start':idea[0], 'baseline_time_sec':tb,'idea_time_sec':ti,'loss_ratio_idea_over_baseline':idea[-1]/base[-1]},
120      'loss_curves':{'baseline':base,'idea':idea}}
121    Path('results.json').write_text(json.dumps(result,indent=2))
122    print(json.dumps({k:v for k,v in result.items() if k!='loss_curves'},indent=2))
123
124if __name__=='__main__': main()