Rank-One Delta Associative Memory / local_sequence_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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