import json, math, time from pathlib import Path import numpy as np import torch SEED=17 np.random.seed(SEED); torch.manual_seed(SEED) def hankel_cols(seq, L): # seq: [T,d], each column is one flattened length-L window T,d=seq.shape return np.stack([seq[t:t+L].reshape(-1) for t in range(T-L+1)], axis=1) def behavior_matrix(A, B, C, inputs, L): # zero-state finite-horizon input/output graphs, concatenated over trials blocks=[] for u in inputs: x=np.zeros(A.shape[0]); ys=[] for t in range(len(u)): x=A@x+B@u[t] ys.append(C@x) blocks.append(np.vstack([hankel_cols(u,L), hankel_cols(np.asarray(ys),L)])) return np.concatenate(blocks, axis=1) def basis_and_projector(M, r): U,s,_=np.linalg.svd(M, full_matrices=False) rr=min(r, int(np.sum(s > max(M.shape)*s[0]*1e-10)), U.shape[1]) rr=max(1,rr) Q=U[:,:rr] return Q, Q@Q.T, s def gap(P,Q): # spectral norm; symmetric difference is numerically stable here return float(np.linalg.norm(P-Q,2)) def cluster(Qs, Ps, eps): groups=[] for i,P in enumerate(Ps): placed=False for g in groups: if gap(P, Ps[g[0]]) <= eps: g.append(i); placed=True; break if not placed: groups.append([i]) return groups def adam_train(A0, targets, inputs, groups, steps=220, lr=.035, shared=True): # Learn each module's recurrent A from trajectory loss. Shared mode averages # first/second Adam statistics within each behavior cluster, while retaining A_i. A=[torch.tensor(a, dtype=torch.float32, requires_grad=True) for a in A0] n=len(A); m1=[torch.zeros_like(a) for a in A]; m2=[torch.zeros_like(a) for a in A] losses=[] for step in range(1,steps+1): ls=[] for i in range(n): x=torch.zeros(2); pred=[] for t in range(inputs.shape[1]): x=A[i]@x + torch.tensor(inputs[i,t],dtype=torch.float32) pred.append(x) pred=torch.stack(pred) ls.append(((pred-torch.tensor(targets[i],dtype=torch.float32))**2).mean()) total=torch.stack(ls).sum(); total.backward() with torch.no_grad(): for g in groups if shared else [[i] for i in range(n)]: grads=[A[i].grad for i in g] gm=torch.stack(grads).mean(0) # Shared controller state, but module-specific gradient direction/update. for i in g: m1[i].mul_(0.9).add_(gm,alpha=.1) m2[i].mul_(0.999).addcmul_(gm,gm,value=.001) update=m1[i]/(1-.9**step)/(torch.sqrt(m2[i]/(1-.999**step))+1e-8) A[i].add_(update,alpha=-lr) A[i].grad.zero_() losses.append(float(torch.stack(ls).mean())) return losses, [a.detach().numpy() for a in A] def main(): # 8 stable 2D systems: four near one behavior and four near another. n=8; d=2; T=28; Ls=[4,8,16] Atrue=[] for i in range(n): theta=(0.10 if i<4 else 0.72) + np.random.randn()*.025 rad=(.82 if i<4 else .78)+np.random.randn()*.012 Atrue.append(rad*np.array([[math.cos(theta),-math.sin(theta)],[math.sin(theta),math.cos(theta)]])) Atrue=np.array(Atrue); B=np.eye(2); C=np.eye(2) trials=[] for k in range(5): trials.append(np.random.randn(T,2)*.5) # math sanity: projector gap identity and angle behavior q,_=np.linalg.qr(np.random.randn(8,3)); r,_=np.linalg.qr(np.random.randn(8,3)) P=q@q.T; R=r@r.T 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))) angle=[] for th in [0,.1,.3,.7,1.2]: q1=np.array([[1.],[0.]]); q2=np.array([[math.cos(th)],[math.sin(th)]]) angle.append((th,float(np.linalg.norm(q1@q1.T-q2@q2.T,2)))) # behavior spaces and horizon effect gaps={}; clusters={} for L in Ls: Qs=[]; Ps=[] for a in Atrue: M=behavior_matrix(a,B,C,trials,L); q,p,s=basis_and_projector(M, min(12,M.shape[1])) Qs.append(q); Ps.append(p) gaps[str(L)] = [[round(gap(Ps[i],Ps[j]),5) for j in range(n)] for i in range(n)] clusters[str(L)] = cluster(Qs,Ps,.35) # Optimization uses L=8 clusters; targets are trajectories under true systems. inp=np.random.randn(n,T,2)*.45 targets=[] for i in range(n): x=np.zeros(2); yy=[] for t in range(T): x=Atrue[i]@x+inp[i,t]; yy.append(x.copy()) targets.append(np.asarray(yy)) targets=np.asarray(targets) Ainit=Atrue+np.random.randn(*Atrue.shape)*.18 groups=clusters['8'] t0=time.perf_counter(); base,_=adam_train(Ainit,targets,inp,groups,shared=False); tb=time.perf_counter()-t0 t0=time.perf_counter(); idea,_=adam_train(Ainit,targets,inp,groups,shared=True); ti=time.perf_counter()-t0 result={'seed':SEED,'math':{'projector_identity_abs_error':identity_err,'angle_gap':angle}, '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}, '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]}, 'loss_curves':{'baseline':base,'idea':idea}} Path('results.json').write_text(json.dumps(result,indent=2)) print(json.dumps({k:v for k,v in result.items() if k!='loss_curves'},indent=2)) if __name__=='__main__': main()