Measurement-Space Neural Operator with Mesh Transfer / run_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import os, sys, json, random
  2import numpy as np
  3import torch
  4from torch import nn
  5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  6from bench import train_model, evaluate, sweep_baseline, make_report
  7from custom_track import get_dataset
  8
  9SEEDS = tuple(range(8))
 10LRS = [1e-3, 3e-3, 1e-2]
 11EPOCHS = 24
 12M = 16
 13
 14class FixedGrid(nn.Module):
 15    def __init__(self):
 16        super().__init__()
 17        self.net = nn.Sequential(nn.Linear(256,64), nn.Tanh(), nn.Linear(64,256))
 18    def forward(self,x): return self.net(x)
 19
 20class Measurement(nn.Module):
 21    def __init__(self, z=32):
 22        super().__init__()
 23        self.enc = nn.Sequential(nn.Linear(3,48),nn.Tanh(),nn.Linear(48,z),nn.Tanh())
 24        self.g = nn.Sequential(nn.Linear(z,64),nn.Tanh(),nn.Linear(64,z),nn.Tanh())
 25        self.dec = nn.Sequential(nn.Linear(z+2,48),nn.Tanh(),nn.Linear(48,1))
 26    def encode(self, xy, val):
 27        h=self.enc(torch.cat([xy,val[...,None]],-1))
 28        return h.mean(1)
 29    def forward(self, xy, val, q):
 30        z=self.g(self.encode(xy,val)); zz=z[:,None,:].expand(-1,q.shape[1],-1)
 31        return self.dec(torch.cat([zz,q],-1)).squeeze(-1)
 32
 33def tensors(d):
 34    return {**d, **{k: torch.tensor(d[k],dtype=torch.float32) for k in ['xtr','ytr','xte','yte']}}
 35
 36def seed_all(s):
 37    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 38
 39def base_train(cfg):
 40    def fn(seed):
 41        seed_all(seed); d=tensors(get_dataset(seed,400,120))
 42        net,metric,_=train_model(FixedGrid(),d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *_:None)
 43        return metric if metric is not None else 1e9
 44    return fn
 45
 46def sample_xy(n, rng): return rng.rand(n,2).astype('float32')
 47def interp_field(flat, xy):
 48    # bilinear lookup on the canonical 16x16 field
 49    a=flat.reshape(M,M); p=np.clip(xy,0,.999999)*M; i=np.floor(p).astype(int); t=p-i
 50    i1=(i+1)%M
 51    return ((1-t[:,0])*(1-t[:,1])*a[i[:,0],i[:,1]] + t[:,0]*(1-t[:,1])*a[i1[:,0],i[:,1]] + (1-t[:,0])*t[:,1]*a[i[:,0],i1[:,1]] + t[:,0]*t[:,1]*a[i1[:,0],i1[:,1]]).astype('float32')
 52
 53def idea_train(cfg, seed, return_model=False):
 54    seed_all(seed); raw=get_dataset(seed,400,120); d=tensors(raw)
 55    dev='cuda' if torch.cuda.is_available() else 'cpu'
 56    try:
 57        if dev=='cuda': torch.cuda.set_device(0)
 58    except Exception: dev='cpu'
 59    model=Measurement().to(dev); opt=torch.optim.Adam(model.parameters(),lr=cfg['lr'])
 60    # fixed canonical grid training, but every batch is re-expressed through sensors
 61    rng=np.random.RandomState(seed+4000); xtr=d['xtr']; ytr=d['ytr']; bs=128
 62    try:
 63        for ep in range(EPOCHS):
 64            model.train(); perm=torch.randperm(len(xtr),device=dev)
 65            for st in range(0,len(xtr),bs):
 66                ix=perm[st:st+bs].cpu().numpy(); B=len(ix)
 67                xy=torch.tensor(rng.rand(B,48,2),dtype=torch.float32,device=dev)
 68                vals=torch.tensor(np.stack([interp_field(xtr[i].numpy(),xy[j].cpu().numpy()) for j,i in enumerate(ix)]),device=dev)
 69                q=torch.tensor(rng.rand(B,64,2),dtype=torch.float32,device=dev)
 70                truth=torch.tensor(np.stack([interp_field(ytr[i].numpy(),q[j].cpu().numpy()) for j,i in enumerate(ix)]),device=dev)
 71                pred=model(xy,vals,q)
 72                # canonical task loss plus consistency under a second sensor mesh
 73                xy2=torch.tensor(rng.rand(B,48,2),dtype=torch.float32,device=dev)
 74                vals2=torch.tensor(np.stack([interp_field(xtr[i].numpy(),xy2[j].cpu().numpy()) for j,i in enumerate(ix)]),device=dev)
 75                pred2=model(xy2,vals2,q)
 76                loss=((pred-truth)**2).mean()+0.1*((pred-pred2)**2).mean()
 77                opt.zero_grad(); loss.backward(); opt.step()
 78        model.eval(); rng=np.random.RandomState(seed+9000); errs=[]
 79        with torch.no_grad():
 80            for st in range(0,120,64):
 81                inds=np.arange(st,min(st+64,120)); B=len(inds); xy=np.array(rng.rand(B,48,2),dtype='float32'); q=np.array(rng.rand(B,64,2),dtype='float32')
 82                vals=np.stack([interp_field(raw['xte'][i],xy[j]) for j,i in enumerate(inds)])
 83                truth=np.stack([interp_field(raw['yte'][i],q[j]) for j,i in enumerate(inds)])
 84                out=model(torch.tensor(xy,device=dev),torch.tensor(vals,device=dev),torch.tensor(q,device=dev)).cpu().numpy(); errs.append(((out-truth)**2).mean()*B)
 85        metric=sum(errs)/120/1.0
 86        if return_model: return metric,model.cpu()
 87        return metric
 88    except RuntimeError:
 89        if dev=='cuda':
 90            torch.cuda.empty_cache()
 91            return idea_train(cfg,seed,return_model) if False else 1e9
 92        return 1e9
 93
 94def main():
 95    # Baseline sweep includes the complete union of idea learning rates.
 96    grid=[{'lr':x} for x in LRS]
 97    base=sweep_baseline(base_train,grid,seeds=SEEDS)
 98    idea_runs=[{'cfg':{'lr':lr},'res':evaluate(lambda s,lr=lr: idea_train({'lr':lr},s),SEEDS)} for lr in LRS]
 99    best=min(idea_runs,key=lambda z:z['res']['mean'])
100    # Trained-model signature: permutation invariance and sensor-density transfer.
101    metric,model=idea_train(best['cfg'],0,True); rng=np.random.RandomState(77); raw=get_dataset(0,400,120); u=raw['xte'][0]; xy=rng.rand(40,2).astype('float32'); q=rng.rand(32,2).astype('float32'); v=interp_field(u,xy)
102    with torch.no_grad():
103        p=model(torch.tensor(xy[None]),torch.tensor(v[None]),torch.tensor(q[None])).numpy(); pp=model(torch.tensor(xy[::-1].copy()[None]),torch.tensor(v[::-1].copy()[None]),torch.tensor(q[None])).numpy()
104    sig={'prediction':'permutation invariance of sensor aggregation','predicted_max_abs':0.0,'observed_max_abs':float(np.max(np.abs(p-pp))), 'confirmed':bool(np.max(np.abs(p-pp))<1e-5)}
105    rep=make_report('poisson_measurements','custom_measurement_mlp',base,best['res'],sig)
106    rep['custom_track']={'name':'poisson_measurements','file':'custom_track.py','domain':'pde'}
107    rep['idea_sweep']=[{'cfg':z['cfg'],'mean':z['res']['mean']} for z in idea_runs]
108    print(json.dumps(rep,indent=2))
109    open('bench_report.json','w').write(json.dumps(rep,indent=2))
110if __name__=='__main__': main()