Energy-Riesz checkpoint selector / bench_energy_riesz.py
Mechanism confirmed, baseline not beaten
1import sys, json, time
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, make_model, evaluate, sweep_baseline, make_report
8
9ROOT=Path(__file__).parent
10TRACK='poisson_boundary'; MODEL='mlp_tiny'; EPOCHS=24; BATCH=64
11# Same union of learning rates is used by baseline and selector.
12LRS=[1e-3,3e-3,1e-2]
13LEVELS=[2,3,4]
14
15def device(): return 'cuda' if torch.cuda.is_available() else 'cpu'
16def seed_all(s):
17 np.random.seed(s); torch.manual_seed(s)
18
19def wrapped(net,x):
20 return (x[:,0:1]*(1-x[:,0:1])*x[:,1:2]*(1-x[:,1:2]))*net(x)
21
22def make_basis(level):
23 # homogeneous sine basis, nested in the sense of increasing mode sets
24 return [(p,q) for p in range(1,level+1) for q in range(1,level+1)]
25
26def riesz_score(net, level, dev, nq=18):
27 # Tensor quadrature estimates a(u,phi) and l(phi), then solves A z=b.
28 z=torch.linspace(0.0,1.0,nq,device=dev)[1:-1]
29 X,Y=torch.meshgrid(z,z,indexing='ij'); pts=torch.stack([X.reshape(-1),Y.reshape(-1)],1).requires_grad_(True)
30 with torch.enable_grad():
31 u=wrapped(net,pts).reshape(-1)
32 gu=torch.autograd.grad(u.sum(),pts,create_graph=False)[0]
33 modes=make_basis(level); A=np.zeros((len(modes),len(modes))); b=np.zeros(len(modes))
34 wt=1.0/((nq-1)**2)
35 for i,(p,q) in enumerate(modes):
36 phi=torch.sin(np.pi*p*pts[:,0])*torch.sin(np.pi*q*pts[:,1])
37 gp=torch.stack([np.pi*p*torch.cos(np.pi*p*pts[:,0])*torch.sin(np.pi*q*pts[:,1]), np.pi*q*torch.sin(np.pi*p*pts[:,0])*torch.cos(np.pi*q*pts[:,1])],1)
38 # -Delta exact manufactured solution f=2*pi^2*sin(pi x)sin(pi y)
39 load=(2*np.pi**2*torch.sin(np.pi*pts[:,0])*torch.sin(np.pi*pts[:,1])*phi).mean().item()
40 b[i]=load-(gu*gp).sum(1).mean().item()
41 for j,(r,s) in enumerate(modes):
42 A[i,j]=((np.pi**2*(p*r+q*s)/2)*0 + 0) # overwritten by quadrature below
43 phij=torch.sin(np.pi*r*pts[:,0])*torch.sin(np.pi*s*pts[:,1])
44 gij=torch.stack([np.pi*r*torch.cos(np.pi*r*pts[:,0])*torch.sin(np.pi*s*pts[:,1]), np.pi*s*torch.sin(np.pi*r*pts[:,0])*torch.cos(np.pi*s*pts[:,1])],1)
45 A[i,j]=(gp*gij).sum(1).mean().item()
46 zz=np.linalg.solve(A+1e-9*np.eye(len(modes)),b)
47 return float(np.sqrt(max(0,zz@A@zz)))
48
49def train_system(seed, lr, selector=False, level=3, return_sig=False):
50 seed_all(seed); ds=get_dataset(TRACK,400,400)
51 dev=device()
52 try:
53 net=make_model(MODEL,tuple(ds['xtr'].shape[1:]),1).to(dev)
54 x=ds['xtr'].to(dev); y=ds['ytr'].to(dev); xe=ds['xte'].to(dev); ye=ds['yte'].to(dev)
55 opt=torch.optim.Adam(net.parameters(),lr=lr); lossf=nn.MSELoss(); checkpoints=[]; losses=[]
56 g=torch.Generator(device=dev); g.manual_seed(seed)
57 for ep in range(EPOCHS):
58 net.train(); perm=torch.randperm(len(x),generator=g,device=dev)
59 for ii in perm.split(BATCH):
60 pred=wrapped(net,x[ii]); loss=lossf(pred,y[ii]); opt.zero_grad(); loss.backward(); opt.step()
61 net.eval()
62 with torch.no_grad(): losses.append(float(lossf(wrapped(net,x),y).item()))
63 if selector and (ep%4==3 or ep==EPOCHS-1):
64 checkpoints.append({k:v.detach().cpu().clone() for k,v in net.state_dict().items()})
65 if not selector:
66 with torch.no_grad(): metric=float(lossf(wrapped(net,xe),ye).item())
67 return (metric,{"final_train_loss":losses[-1]}) if return_sig else metric
68 scores=[]
69 for state in checkpoints:
70 net.load_state_dict(state); net.eval(); scores.append(riesz_score(net,level,dev))
71 chosen=int(np.argmin(scores)); net.load_state_dict(checkpoints[chosen]); net.eval()
72 with torch.no_grad(): metric=float(lossf(wrapped(net,xe),ye).item())
73 sig={"selected_index":chosen,"scores":scores,"nested_probe":{}}
74 for lv in LEVELS:
75 net.load_state_dict(checkpoints[chosen]); sig['nested_probe'][str(lv)]=riesz_score(net,lv,dev)
76 sig['predicted_monotone']=all(sig['nested_probe'][str(a)]<=sig['nested_probe'][str(b)]+1e-5 for a,b in zip(LEVELS,LEVELS[1:]))
77 sig['confirmed']=bool(sig['predicted_monotone'])
78 return (metric,sig) if return_sig else metric
79 except Exception:
80 if dev=='cuda':
81 torch.cuda.empty_cache(); return train_system_cpu(seed,lr,selector,level,return_sig)
82 raise
83
84def train_system_cpu(seed,lr,selector=False,level=3,return_sig=False):
85 old=torch.cuda.is_available
86 # Explicitly duplicate with CPU by temporarily using a local implementation flag.
87 global device
88 orig=device; device=lambda:'cpu'
89 try: return train_system(seed,lr,selector,level,return_sig)
90 finally: device=orig
91
92def baseline_fn(cfg): return lambda s: train_system(s,cfg['lr'],False)
93def idea_fn(cfg): return lambda s: train_system(s,cfg['lr'],True,cfg['level'])
94
95def main():
96 t=time.time(); grid=[{'lr':lr,'level':3} for lr in LRS]
97 base=sweep_baseline(baseline_fn, [{'lr':lr} for lr in LRS])
98 # idea sweep has same lr union; level is the selector knob, baseline uses equivalent shared grid.
99 idea_cfgs=[{'lr':base['best_cfg']['lr'],'level':lv} for lv in LEVELS]
100 idea_runs=[]
101 for cfg in idea_cfgs:
102 r=evaluate(idea_fn(cfg)); idea_runs.append((cfg,r))
103 best_cfg,idea=min(idea_runs,key=lambda z:z[1]['mean'])
104 sig_metric,sig=train_system(0,best_cfg['lr'],True,best_cfg['level'],True)
105 extra={'mechanism_signature':sig,'custom_track':{'name':TRACK,'file':'poisson_track.py','domain':'pde'}}
106 rep=make_report(TRACK,MODEL,base,idea,extra)
107 rep['idea_config_sweep']=[{'cfg':c,'mean':r['mean'],'per_seed':r['per_seed']} for c,r in idea_runs]
108 rep['runtime_seconds']=time.time()-t; rep['protocol_note']='Custom PDE track required; baseline and idea share mlp_tiny and Adam, differing only in archived checkpoint selector.'
109 (ROOT/'bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
110if __name__=='__main__': main()