Gain-Rigid Sparse Attention / update_signature.py
Beats tuned baseline
1import json, random, sys
2import numpy as np
3import torch
4sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
5from bench import get_dataset, make_model, train_model
6from bench_gain_rigid import GainTransformer, EPOCHS, NTRAIN, NTEST
7
8seed=0; lr=0.003
9torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
10ds=get_dataset('sequence',seed,NTRAIN,NTEST)
11base=make_model('transformer_tiny',ds['input_shape'],ds['out_dim'])
12base, bm, _=train_model(base,ds,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
13torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
14idea=GainTransformer(ds['input_shape'][0],ds['out_dim'])
15idea, im, _=train_model(idea,ds,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
16dev_b=next(base.parameters()).device
17x=ds['xte'][:8].to(dev_b).clone().requires_grad_(True)
18gb=torch.autograd.grad(base(x).sum(),x)[0].abs().detach().cpu().numpy()
19dev_i=next(idea.parameters()).device
20x=ds['xte'][:8].to(dev_i).clone().requires_grad_(True)
21gi=torch.autograd.grad(idea(x).sum(),x)[0].abs().detach().cpu().numpy()
22threshold=1e-8
23sig={'prediction':'trained gain-rigid sparse attention should retain connected input influence, with nonzero gradient reach comparable to dense attention', 'observed':{'seed':seed,'baseline_test_mse':float(bm),'idea_test_mse':float(im),'baseline_input_gradient_mean':float(gb.mean()),'idea_input_gradient_mean':float(gi.mean()),'baseline_input_gradient_active_fraction':float((gb>threshold).mean()),'idea_input_gradient_active_fraction':float((gi>threshold).mean()),'gain_edges':int(idea.edge_count),'nodes':32},'confirmed':bool((gi>threshold).mean() >= .95*(gb>threshold).mean())}
24with open('bench_report.json') as f: rep=json.load(f)
25rep['mechanism_signature']=sig
26with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
27print(json.dumps(sig,indent=2))