Rank-One Delta Associative Memory / local_sequence_bench.py
Mechanism confirmed, baseline not beaten
1import json, math, random, time
2import numpy as np
3import torch
4from torch import nn
5
6DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
7
8def seed_all(s):
9 random.seed(s); np.random.seed(s); torch.manual_seed(s)
10 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
11
12def make_data(seed, n=320, T=24):
13 r=np.random.default_rng(seed)
14 z=np.zeros((n,T+1),np.float32); z[:,0]=r.normal(size=n)
15 for t in range(T): z[:,t+1]=.82*z[:,t]+.35*r.normal(size=n)
16 x=np.stack([z[:,:T], r.normal(size=(n,T)), r.normal(size=(n,T))],-1).astype('float32')
17 y=z[:,T,None].astype('float32')
18 return torch.tensor(x),torch.tensor(y)
19
20class DeltaMemory(nn.Module):
21 def __init__(self,d,h=16,beta=.5):
22 super().__init__(); self.k=nn.Linear(d,h); self.v=nn.Linear(d,h); self.out=nn.Linear(h,1); self.beta=beta
23 def forward(self,x):
24 B,T,D=x.shape; W=x.new_zeros(B,16,16); reads=[]; norms=[]; changes=[]
25 for t in range(T):
26 k=torch.tanh(self.k(x[:,t])); k=k/(k.norm(dim=-1,keepdim=True).clamp_min(1e-5))
27 v=self.v(x[:,t]); m=torch.einsum('bij,bj->bi',W,k); r=v-m
28 W=W+self.beta*r.unsqueeze(-1)*k.unsqueeze(-2)
29 reads.append(m); norms.append(W.square().mean().sqrt().detach()); changes.append((self.beta*r.norm(dim=-1)*k.norm(dim=-1)).mean().detach())
30 return self.out(reads[-1]), torch.stack(norms).mean(), torch.stack(changes).mean()
31
32class StandardGRU(nn.Module):
33 def __init__(self):
34 super().__init__(); self.rnn=nn.GRU(3,16,batch_first=True); self.head=nn.Linear(16,1)
35 def forward(self,x): return self.head(self.rnn(x)[0][:,-1])
36
37class DeltaGRU(nn.Module):
38 def __init__(self,beta=.5):
39 super().__init__(); self.rnn=nn.GRU(3,16,batch_first=True); self.mem=DeltaMemory(3,16,beta)
40 self.head=nn.Linear(16,1); self.mix=nn.Parameter(torch.tensor(0.0))
41 def forward(self,x):
42 h=self.rnn(x)[0][:,-1]; m,n,c=self.mem(x); g=torch.sigmoid(self.mix)
43 return self.head(h)+g*m,n,c
44
45def train(seed, idea, lr=3e-3, beta=.5, epochs=18):
46 seed_all(seed); xtr,ytr=make_data(seed); xte,yte=make_data(seed+10000)
47 model=DeltaGRU(beta).to(DEVICE) if idea else StandardGRU().to(DEVICE)
48 opt=torch.optim.Adam(model.parameters(),lr=lr); bs=64
49 xtr,ytr,xte,yte=[q.to(DEVICE) for q in (xtr,ytr,xte,yte)]
50 t0=time.time(); model.train()
51 for _ in range(epochs):
52 p=torch.randperm(len(xtr),device=DEVICE)
53 for j in range(0,len(xtr),bs):
54 out=model(xtr[p[j:j+bs]])
55 pred=out[0] if isinstance(out,tuple) else out
56 loss=((pred-ytr[p[j:j+bs]])**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
57 model.eval()
58 with torch.no_grad():
59 out=model(xte); pred=out[0] if isinstance(out,tuple) else out
60 mse=float(((pred-yte)**2).mean().cpu())
61 sig={}
62 if idea: sig={'state_norm':float(out[1].cpu()),'mean_rank_one_update_norm':float(out[2].cpu())}
63 return mse,sig,time.time()-t0
64
65def sign_perm(deltas):
66 obs=float(np.mean(deltas)); rng=np.random.default_rng(1198); count=0; n=20000
67 for _ in range(n):
68 if abs(np.mean(deltas*rng.choice([-1,1],len(deltas))))>=abs(obs): count+=1
69 return obs,float((count+1)/(n+1))
70
71def main():
72 # Shared union of hyperparameters: baseline evaluates every idea lr and nearby lr.
73 lrs=[1.5e-3,3e-3,6e-3]; betas=[.25,.5,.75]; seeds=list(range(8))
74 baseline_runs=[]
75 for lr in lrs:
76 vals=[train(s,False,lr=lr)[0] for s in seeds]
77 baseline_runs.append({'lr':lr,'mean':float(np.mean(vals)),'per_seed':vals})
78 best_base=min(baseline_runs,key=lambda q:q['mean'])
79 idea_runs=[]
80 for beta in betas:
81 vals=[]; ss=[]; tt=[]
82 for s in seeds:
83 v,sg,t=train(s,True,lr=best_base['lr'],beta=beta); vals.append(v); ss.append(sg); tt.append(t)
84 idea_runs.append({'beta':beta,'lr':best_base['lr'],'mean':float(np.mean(vals)),'per_seed':vals,'signature':ss,'seconds':float(np.mean(tt))})
85 best_idea=min(idea_runs,key=lambda q:q['mean']); delta=np.array(best_idea['per_seed'])-np.array(best_base['per_seed']); dm,p=sign_perm(delta)
86 sig=best_idea['signature']; report={'track':'local_sequence_forecast','model':'matched_gru_plus_delta_memory','baseline_sweep':baseline_runs,'idea_sweep':idea_runs,'best_baseline':best_base,'best_idea':best_idea,'paired_delta_mean':dm,'permutation_p_value':p,'mechanism_signature':{'predicted':'rank-one update has nonzero state change and finite state norm','observed_mean_state_norm':float(np.mean([q['state_norm'] for q in sig])),'observed_mean_update_norm':float(np.mean([q['mean_rank_one_update_norm'] for q in sig])),'confirmed':bool(np.mean([q['mean_rank_one_update_norm'] for q in sig])>0 and np.isfinite(np.mean([q['state_norm'] for q in sig])))},'device':DEVICE,'protocol_note':'Official bench package and README were absent; this is a local substitute, not an official bench_report.'}
87 open('bench_report.json','w').write(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
88
89if __name__=='__main__':
90 try: main()
91 except Exception as e:
92 if DEVICE=='cuda':
93 print('CUDA failed, rerun on CPU:',repr(e)); DEVICE='cpu'; main()
94 else: raise