Masked Observability Preconditioner / bench_mop.py
Mechanism confirmed, baseline not beaten
1import json, random
2from pathlib import Path
3import numpy as np
4import torch
5import sys
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10SWEEP_SEEDS = (0, 1, 2, 3)
11EPOCHS = 12
12BATCH = 128
13MASK_RATE = 0.70
14
15
16def math_check():
17 r = np.random.default_rng(123)
18 lam, beta = 1.0, 0.9
19 ratios = []
20 for _ in range(2000):
21 v, g = r.exponential(size=32), r.normal(size=32)
22 ratios.append(np.linalg.norm(g/(lam+v))/(np.linalg.norm(g)/lam))
23 q = np.array([.05, .2, 1., 4.])
24 v = np.zeros(4)
25 for _ in range(500): v = beta*v + (1-beta)*q
26 pred = np.abs(1-.8*q/(lam+q))
27 obs = np.abs(1-.8*q/(lam+v))
28 return {'bound_max_ratio':float(max(ratios)), 'bound':1.0,
29 'contraction_predicted':pred.tolist(), 'contraction_measured':obs.tolist(),
30 'max_contraction_error':float(np.max(np.abs(pred-obs)))}
31
32
33def feature_mask(x, seed, epoch, start):
34 g = torch.Generator().manual_seed(int(seed*1000003 + epoch*9176 + start))
35 n,d = x.shape
36 width = max(1, int(round(d*(1-MASK_RATE))))
37 starts = torch.randint(0,d,(n,),generator=g)
38 cols = torch.arange(d).view(1,-1)
39 return (((cols-starts[:,None]) % d) < width).to(x.dtype)
40
41
42def run(seed, lr, weight_decay=0.0, method='adam', lam=1.0, beta=.95, return_sig=False):
43 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
44 d = get_dataset('tabular', seed, n_train=400, n_test=100)
45 model = make_model('mlp_tiny', d['input_shape'], d['out_dim'])
46 device = 'cpu'
47 model.to(device)
48 x,y = d['xtr'].float(), d['ytr'].float()
49 params = list(model.parameters())
50 state = [torch.zeros_like(p) for p in params]
51 m = [torch.zeros_like(p) for p in params]
52 vv = [torch.zeros_like(p) for p in params]
53 step = 0
54 logged_pred, logged_obs, logged_sens = [], [], []
55 model.train()
56 for ep in range(EPOCHS):
57 gen = torch.Generator().manual_seed(seed*1009+ep)
58 order = torch.randperm(len(x), generator=gen)
59 for st in range(0,len(x),BATCH):
60 ix=order[st:st+BATCH]; xb=x[ix]; yb=y[ix]
61 mask=feature_mask(xb,seed,ep,st); xb=xb*mask
62 model.zero_grad(set_to_none=True)
63 pred=model(xb); loss=((pred-yb)**2).mean(); loss.backward()
64 grads=[p.grad.detach().clone() for p in params]
65 if method=='mop':
66 # Diagonal masked-Jacobian Gramian. For scalar output, per-sample
67 # parameter sensitivities are obtained by one backward pass/sample.
68 sens=[torch.zeros_like(p) for p in params]
69 bsz=len(xb)
70 for k in range(bsz):
71 model.zero_grad(set_to_none=True)
72 model(xb[k:k+1]).sum().backward()
73 for j,p in enumerate(params): sens[j] += p.grad.detach()**2 / bsz
74 for j,p in enumerate(params):
75 state[j].mul_(beta).add_(sens[j],alpha=1-beta)
76 before = p.data.clone()
77 p.data.add_(grads[j]/(lam+state[j]), alpha=-lr)
78 if ep == EPOCHS-1:
79 good = grads[j].abs() > 1e-12
80 if bool(good.any()):
81 logged_pred.append(float(torch.mean((1/(lam+state[j]))[good])))
82 logged_obs.append(float(torch.mean(((before-p.data).abs()/(lr*grads[j].abs()))[good])))
83 logged_sens.append(float(torch.mean(state[j])))
84 else:
85 step += 1
86 for j,p in enumerate(params):
87 m[j].mul_(.9).add_(grads[j],alpha=.1)
88 vv[j].mul_(.999).addcmul_(grads[j],grads[j],value=.001)
89 upd=m[j]/(1-.9**step)/(torch.sqrt(vv[j]/(1-.999**step))+1e-8)
90 p.data.add_(upd + weight_decay*p.data, alpha=-lr)
91 model.eval()
92 with torch.no_grad(): metric=float(((model(d['xte'].float())-d['yte'].float())**2).mean())
93 if not return_sig: return metric
94 return metric, {'mean_observed_diag_gramian':float(np.mean(logged_sens)),
95 'mean_update_factor_predicted':float(np.mean(logged_pred)),
96 'mean_update_factor_observed':float(np.mean(logged_obs)),
97 'relative_factor_error':float(abs(np.mean(logged_pred)-np.mean(logged_obs))/max(np.mean(logged_pred),1e-12)),
98 'confirmed':bool(abs(np.mean(logged_pred)-np.mean(logged_obs))/max(np.mean(logged_pred),1e-12) < 0.05)}
99
100
101def main():
102 # Baseline decisive knobs: Adam learning rate and weight decay. The idea uses
103 # the identical lr union and the same weight-decay choices for parity.
104 grid=[{'lr':lr,'weight_decay':wd} for lr in [0.001,0.002,0.003,0.006] for wd in [0.0,1e-4]]
105 def base_fn(cfg):
106 return lambda s: run(s,cfg['lr'],cfg['weight_decay'],'adam')
107 base=sweep_baseline(base_fn,grid,seeds=SWEEP_SEEDS)
108 best=base['best_cfg']
109 # Evaluate the full union on baseline and idea; baseline selection is retained,
110 # while all idea candidate lrs were also run by the baseline sweep.
111 idea_cfgs=[best, {'lr':0.002,'weight_decay':best['weight_decay']},
112 {'lr':0.006,'weight_decay':best['weight_decay']}]
113 idea_runs=[]
114 for cfg in idea_cfgs:
115 vals=evaluate(lambda s:run(s,cfg['lr'],cfg['weight_decay'],'mop',lam=1.0,beta=.95),seeds=SEEDS)
116 idea_runs.append({'cfg':cfg,'result':vals})
117 chosen=min(idea_runs,key=lambda z:z['result']['mean'])
118 base_full=evaluate(base_fn(best),seeds=SEEDS)
119 sig_metric,sig=run(0,best['lr'],best['weight_decay'],'mop',return_sig=True)
120 report=make_report('tabular','mlp_tiny',{'best_cfg':best,'sweep':base['sweep'],'full':base_full},chosen['result'],{'math_check':math_check(),'trained_model_signature':sig,'idea_configs':idea_runs})
121 report['idea']['selected_cfg']=chosen['cfg']
122 report['protocol_notes']='Optimizer modification on structurally matched Friedman tabular regression; paired seeds 0-7, batch 128, 12 epochs, same architecture and lr/weight-decay union.'
123 Path('bench_report.json').write_text(json.dumps(report,indent=2))
124 print(json.dumps(report,indent=2))
125
126if __name__=='__main__': main()