Energy-Riesz checkpoint selector / bench_energy_riesz.py

Mechanism confirmed, baseline not beaten

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