Displacement-Huber distribution pooling / hub_bench.py

✓✓ Beats tuned baseline

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