Utility-Weighted Left-Edge Quantization / bench_experiment.py
Mechanism confirmed, baseline not beaten
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()