Intrinsic-Rank Filter Memory for Actor-Critic / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import json, random, sys
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
7SEEDS=tuple(range(8)); SWEEP_SEEDS=(0,1,2,3); LRS=[1e-3,3e-3,1e-2]; TAUS=[1e-2,1e-3,1e-4]; EPOCHS=12; NTR,NTE=400,200
8
9def seed_all(seed):
10 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
11 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
12
13def causal_lift(x):
14 x=np.asarray(x,dtype=np.float32).reshape(-1,8,3); f=np.zeros_like(x); a=.72
15 for t in range(8): f[:,t]=a*(f[:,t-1] if t else 0)+(1-a)*x[:,t]
16 return np.concatenate([x,f,f],axis=2).astype(np.float32)
17
18def prepare(seed):
19 d=get_dataset('dynamics',seed=seed,n_train=NTR,n_test=NTE); ztr=causal_lift(d['xtr'].numpy()); zte=causal_lift(d['xte'].numpy())
20 rows=ztr.reshape(-1,9); mean=rows.mean(0); _,s,vt=np.linalg.svd(rows-mean,full_matrices=False)
21 d.update(_mean=mean.astype(np.float32),_s=s.astype(np.float32),_vt=vt.astype(np.float32),_ztr=ztr,_zte=zte); return d
22
23class SharedGRU(nn.Module):
24 def __init__(self,input_dim):
25 super().__init__(); self.rnn=nn.GRU(input_dim,64,batch_first=True); self.head=nn.Linear(64,1)
26 def forward(self,x):
27 _,h=self.rnn(x.view(x.shape[0],8,-1)); return self.head(h[-1])
28
29def transformed_ds(d,kind,tau=1e-3):
30 out=dict(d)
31 if kind=='baseline': tr,te=d['_ztr'],d['_zte']
32 else:
33 r=max(1,int(np.sum(d['_s']/d['_s'][0]>tau))); tr=(d['_ztr']-d['_mean'])@d['_vt'][:r].T; te=(d['_zte']-d['_mean'])@d['_vt'][:r].T
34 out['xtr']=torch.as_tensor(tr.reshape(NTR,-1)); out['xte']=torch.as_tensor(te.reshape(NTE,-1)); out['input_shape']=tuple(out['xtr'].shape[1:]); out['out_dim']=1; return out
35
36def run(kind,lr,seed,tau=1e-3,return_model=False):
37 seed_all(seed); d=prepare(seed); q=transformed_ds(d,kind,tau); width=9 if kind=='baseline' else max(1,int(np.sum(d['_s']/d['_s'][0]>tau)))
38 net,metric,hist=train_model(SharedGRU(width),q,epochs=EPOCHS,lr=float(lr),batch=128,log=lambda *a,**k:None)
39 if metric is None: raise RuntimeError('training failed')
40 return (float(metric),net,q,d) if return_model else float(metric)
41
42def base_factory(cfg): return lambda seed: run('baseline',cfg['lr'],seed)
43def idea_factory(cfg): return lambda seed: run('idea',cfg['lr'],seed,cfg['tau'])
44
45def mechanism_signature():
46 d=prepare(9001); ratios=d['_s']/d['_s'][0]; rank=int(np.sum(ratios>1e-3)); vals={}
47 for kind in ('baseline','idea'):
48 metric,net,q,_=run(kind,3e-3,0,1e-3,True)
49 # Avoid shared-GPU cuDNN allocation failures during probing.
50 torch.backends.cudnn.enabled=False; net=net.to('cpu'); net.eval(); x=q['xte'][:32]
51 with torch.no_grad(): p0=net(x); xx=x.clone(); xx[:,0]+=0.25; p1=net(xx)
52 vals[kind]={'metric':metric,'prediction_rms_under_input_perturbation':float(torch.sqrt(torch.mean((p1-p0)**2)))}
53 torch.backends.cudnn.enabled=True
54 return {'prediction':'causal lifted ambient width 9 has deterministic duplicate directions and SVD rank 6','predicted_intrinsic_rank':6,'observed_rank_tau_1e-3':rank,'ambient_coordinate_width':9,'trained_model_behavior':vals,'confirmed':bool(rank==6 and all(np.isfinite(v['prediction_rms_under_input_perturbation']) for v in vals.values()))}
55
56def main():
57 grid=[{'lr':lr,'tau':tau} for lr in LRS for tau in TAUS]; base=sweep_baseline(base_factory,grid,seeds=SWEEP_SEEDS)
58 trials=[{'cfg':c,'result':evaluate(idea_factory(c),SEEDS)} for c in grid]; best=min(trials,key=lambda z:z['result']['mean'])
59 rep=make_report('dynamics','rnn_small',base,best['result'],{'idea_config':best['cfg'],'idea_sweep':trials,'mechanism_signature':mechanism_signature()}); rep['mechanism_signature']=rep.pop('mechanism_signature')
60 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
61 print(json.dumps(rep,indent=2))
62if __name__=='__main__': main()