Displacement-Huber distribution pooling / hub_bench.py
Beats tuned baseline
1import sys, json, time
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import train_model, evaluate, sweep_baseline, make_report
9
10# Custom matched structure: each example is a sequence of token distributions,
11# and the target depends on the robust location of the clean token population.
12META = {'name':'distribution_token_sequence','domain':'sequence',
13 'description':'Regression from a sequence of quantile-vector token distributions; some training/test tokens are grossly corrupted.'}
14
15def get_dataset(seed, n_train=400, n_test=400):
16 rng = np.random.RandomState(seed)
17 nt, m = 12, 16
18 def make(n, train=False):
19 y = rng.uniform(-1, 1, n).astype(np.float32)
20 # Token distributions have a shared location signal plus token noise.
21 centers = y[:,None] + rng.normal(0, .08, (n,nt)).astype(np.float32)
22 levels = np.linspace(-.45,.45,m,dtype=np.float32)
23 q = centers[:,:,None] + levels[None,None,:] + rng.normal(0,.035,(n,nt,m)).astype(np.float32)
24 # contamination is present only in the observed input; target remains y.
25 rate = .30
26 bad = rng.rand(n,nt) < rate
27 q += bad[:,:,None] * rng.choice([-1.,1.], size=(n,nt,1)).astype(np.float32) * 3.0
28 return q.astype(np.float32), y[:,None]
29 xtr,ytr=make(n_train,True); xte,yte=make(n_test,False)
30 return {'xtr':torch.tensor(xtr), 'ytr':torch.tensor(ytr),
31 'xte':torch.tensor(xte), 'yte':torch.tensor(yte),
32 'task':'regression','metric':'mse','input_shape':(nt,m),'out_dim':1,
33 'track':'distribution_token_sequence'}
34
35class PoolNet(nn.Module):
36 def __init__(self, m=16, mode='mean', delta=0.8):
37 super().__init__(); self.mode=mode; self.delta=delta
38 self.token = nn.Sequential(nn.Linear(m,32), nn.ReLU(), nn.Linear(32,16))
39 self.head = nn.Sequential(nn.Linear(16,32),nn.ReLU(),nn.Linear(32,1))
40 def pool(self,q):
41 # q [B,T,M]. Pointwise Huber barycenter, safeguarded by fixed bracket.
42 if self.mode=='mean': return q.mean(1)
43 z=q.mean(1); lo=q.min(1).values; hi=q.max(1).values
44 for _ in range(8):
45 r=z[:,None,:]-q
46 psi=r.clamp(-self.delta,self.delta)
47 active=(r.abs()<self.delta).float()
48 score=psi.mean(1); cur=active.mean(1)
49 newton=z-score/(cur+1e-5)
50 ok=(newton>=lo)&(newton<=hi)&(cur>1e-5)
51 z=torch.where(ok,newton,(lo+hi)/2)
52 r=z[:,None,:]-q; score=r.clamp(-self.delta,self.delta).mean(1)
53 lo=torch.where(score<0,z,lo); hi=torch.where(score>0,z,hi)
54 return z.sort(dim=-1).values
55 def forward(self,x): return self.head(self.pool(self.token(x)))
56
57def seed_all(s):
58 np.random.seed(s); torch.manual_seed(s)
59 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
60
61def run(cfg, mode, seed):
62 seed_all(seed); d=get_dataset(seed,400,400)
63 net=PoolNet(mode=mode,delta=cfg.get('delta',.8))
64 _, metric, hist=train_model(net,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *a,**k:None)
65 return float(metric)
66
67def main():
68 # Union parity: every lr tried for Huber is also swept for arithmetic.
69 grid=[{'lr':x,'epochs':12,'delta':d} for x in (0.001,0.003,0.006) for d in (0.5,0.8,1.2)]
70 base_grid=[{'lr':x,'epochs':12,'delta':d} for x in (0.001,0.003,0.006) for d in (0.5,0.8,1.2)]
71 base=sweep_baseline(lambda c: lambda s: run(c,'mean',s),base_grid)
72 # Equal 3-config idea sweep at the baseline-selected lr plus nearby deltas.
73 bestlr=base['best_cfg']['lr']
74 idea_grid=[{'lr':bestlr,'epochs':12,'delta':d} for d in (0.5,0.8,1.2)]
75 ir=[]
76 for c in idea_grid: ir.append((c,evaluate(lambda s,c=c:run(c,'huber',s))))
77 bestc, idea=min(ir,key=lambda z:z[1]['mean'])
78 # Signature measured from independently trained benchmark systems.
79 seed_all(123); ds_sig=get_dataset(123,400,400)
80 bnet,_,_=train_model(PoolNet(mode='mean'),ds_sig,epochs=12,lr=base['best_cfg']['lr'],batch=128,log=lambda *a,**k:None)
81 hnet,_,_=train_model(PoolNet(mode='huber',delta=bestc['delta']),ds_sig,epochs=12,lr=bestc['lr'],batch=128,log=lambda *a,**k:None)
82 with torch.no_grad():
83 x=torch.zeros(1,12,16); x[:,:,:]=torch.linspace(-.45,.45,16)
84 def shift(net):
85 clean=net.pool(x).mean(); xx=x.clone(); xx[:,0,:]+=10.0
86 return float((net.pool(xx).mean()-clean).abs())
87 changed=float(shift(hnet)); changedb=float(shift(bnet))
88 sig={'prediction':'a gross token should shift Huber pooled quantiles less than arithmetic pooling',
89 'predicted_bounded_score':bestc['delta']/12.,'observed_huber_shift':changed,
90 'observed_mean_shift':changedb,'confirmed':bool(changed < changedb)}
91 rep=make_report('distribution_token_sequence','custom_poolnet',base,idea,
92 {'prediction':sig['prediction'],'predicted_bounded_score':sig['predicted_bounded_score'],
93 'observed_huber_shift':changed,'observed_mean_shift':changedb,'confirmed':sig['confirmed'],
94 'custom_track':{'name':'distribution_token_sequence','file':'hub_bench.py','domain':'sequence'},
95 'idea_grid':[{**c,'mean':r['mean']} for c,r in ir], 'selected_idea_cfg':bestc})
96 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
97 print(json.dumps(rep,indent=2))
98if __name__=='__main__': main()