Behavior-Gap Clustered Neural Controllers / experiment.py
Mechanism failed
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()