Utility-Weighted Left-Edge Quantization / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset, train_model, sweep_baseline, make_report
 8
 9SEEDS = tuple(range(8)); LRS = [1e-3, 3e-3, 1e-2]; EPOCHS = 18; CODES = 16
10
11class STEQuant(nn.Module):
12    def __init__(self, edges, values):
13        super().__init__(); self.register_buffer('edges',torch.tensor(edges,dtype=torch.float32)); self.register_buffer('values',torch.tensor(values,dtype=torch.float32))
14    def forward(self,x):
15        z=torch.relu(x); idx=torch.bucketize(z.detach(),self.edges[1:-1]); q=self.values[idx]; return z+(q-z).detach()
16
17def lloyd_edges_values(a,n):
18    x=np.asarray(a,dtype=np.float64).ravel(); xmax=max(float(np.quantile(x,.999)),1e-3); x=np.clip(x,0,xmax); c=np.quantile(x,(np.arange(n)+.5)/n)
19    for _ in range(30):
20        lab=np.searchsorted((c[:-1]+c[1:])/2,x); nc=c.copy()
21        for k in range(n):
22            v=x[lab==k]
23            if len(v): nc[k]=v.mean()
24        if np.max(abs(nc-c))<1e-5: break
25        c=nc
26    return np.r_[0,(c[:-1]+c[1:])/2,xmax],c
27
28def utility_edges_values(a,n):
29    x=np.asarray(a,dtype=np.float64).ravel(); xmax=max(float(np.quantile(x,.999)),1e-3); grid=np.linspace(0,xmax,1025)
30    hist,_=np.histogram(np.clip(x,0,xmax),bins=grid); mids=(grid[:-1]+grid[1:])/2
31    p=hist.astype(float)/max(hist.sum(),1); s=max(float(np.median(x)),1e-3); qp=1/((s+mids)*np.log1p(xmax/s)); g=np.sqrt(p*qp+1e-12)
32    # g is a per-bin mass proxy; assign it to the corresponding right edge.
33    c=np.r_[0,np.cumsum(g)]; c/=c[-1]
34    internal=np.interp(np.arange(1,n)/n,c,grid)
35    return np.r_[0,internal,xmax],np.r_[0,internal]
36
37def make_net(ds,kind,seed):
38    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed); probe=nn.Sequential(nn.Linear(10,64),nn.ReLU())
39    with torch.no_grad(): act=probe(ds['xtr']).numpy()
40    e,v=utility_edges_values(act,CODES) if kind=='idea' else lloyd_edges_values(act,CODES)
41    class Net(nn.Module):
42        def __init__(self):
43            super().__init__(); self.l1=nn.Linear(10,64); self.q=STEQuant(e,v); self.l2=nn.Linear(64,64); self.l3=nn.Linear(64,1)
44        def forward(self,x): return self.l3(torch.relu(self.l2(self.q(self.l1(x)))))
45    return Net(),e,v
46
47def run_one(seed,kind,lr):
48    ds=get_dataset('tabular',seed,n_train=400,n_test=400); net,e,v=make_net(ds,kind,seed); net,metric,_=train_model(net,ds,epochs=EPOCHS,lr=lr,batch=128)
49    if net is None:return float('nan'),{}
50    dev=next(net.parameters()).device
51    with torch.no_grad(): z=torch.relu(net.l1(ds['xte'].to(dev))); q=net.q(z); up=float((q>z+1e-6).float().mean().cpu()); gap=float((z-q).mean().cpu())
52    return metric,{'upward_fraction':up,'mean_left_gap':gap,'codes':CODES}
53
54def baseline_factory(cfg): return lambda seed:run_one(seed,'baseline',cfg['lr'])[0]
55
56def main():
57    grid=[{'lr':x} for x in LRS]; base=sweep_baseline(baseline_factory,grid,seeds=SEEDS[:4]); best=base.get('best_config',base.get('best',grid[1])); best_lr=float(best.get('lr',best.get('learning_rate',3e-3)))
58    idea_all=[]
59    for lr in [best_lr]+[x for x in LRS if x!=best_lr][:2]:
60        vals=[]; aux=[]
61        for s in SEEDS:
62            m,a=run_one(s,'idea',lr); vals.append(m); aux.append(a)
63        idea_all.append({'lr':lr,'per_seed':vals,'aux':aux,'mean':float(np.nanmean(vals))})
64    chosen=min(idea_all,key=lambda z:z['mean']); bvals=[]; baux=[]
65    for s in SEEDS:
66        m,a=run_one(s,'baseline',chosen['lr']); bvals.append(m); baux.append(a)
67    base_block={'sweep':base,'best_config':{'lr':chosen['lr']},'full':{'per_seed':bvals,'aux':baux,'mean':float(np.nanmean(bvals))}}
68    idea_block={'config':{'lr':chosen['lr']},'per_seed':chosen['per_seed'],'aux':chosen['aux'],'mean':chosen['mean'],'sweep':idea_all}
69    up=float(np.mean([a['upward_fraction'] for a in chosen['aux']])); gap=float(np.mean([a['mean_left_gap'] for a in chosen['aux']]))
70    sig={'predicted':{'upward_fraction':0.0,'left_gap_nonnegative':True},'observed':{'mean_upward_fraction':up,'mean_left_gap':gap},'confirmed':bool(up==0.0 and gap>=0)}
71    rep=make_report('tabular','mlp_tiny',base_block,idea_block,{'mechanism_signature':sig}); Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
72if __name__=='__main__':main()