Adjoint-Weak Fractional Residuals / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, itertools
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn.functional as F
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, make_report
  8
  9SEEDS = list(range(8))
 10EPOCHS = 16
 11BATCH = 128
 12# Union is shared by baseline and idea: baseline is evaluated at every lr tried by idea.
 13LRS = [1e-3, 3e-3, 1e-2]
 14ALPHAS = [0.55, 0.8, 1.05]
 15
 16def device():
 17    return 'cuda' if torch.cuda.is_available() else 'cpu'
 18
 19def train(d, lr, weak, alpha=0.8, seed=0):
 20    torch.manual_seed(seed); np.random.seed(seed)
 21    try: dev = device(); net = make_model('rnn_small', d['input_shape'], d['out_dim']).to(dev)
 22    except Exception:
 23        dev = 'cpu'; net = make_model('rnn_small', d['input_shape'], d['out_dim']).to(dev)
 24    x, y = d['xtr'].to(dev), d['ytr'].to(dev)
 25    opt = torch.optim.Adam(net.parameters(), lr=lr)
 26    # Smooth compact-ish Gaussian local test kernel and its right-adjoint
 27    # Grünwald fractional derivative, represented by the transpose matrix.
 28    n, alpha_n = 8, alpha
 29    h = 1.0/(n-1)
 30    w = np.empty(n); w[0] = 1.
 31    for k in range(1,n): w[k] = w[k-1] * (-(alpha_n-k+1)/k)
 32    A = np.zeros((n,n), dtype=np.float32)
 33    for i in range(n): A[i,:i+1] = h**(-alpha_n)*w[:i+1][::-1]
 34    t = np.linspace(0,1,n); phi = np.exp(-0.5*((t-.62)/.24)**2).astype(np.float32)
 35    q = np.full(n,h,dtype=np.float32); q[[0,-1]] = h/2
 36    # The residual is based on the learned model's predicted trajectory window.
 37    # Weak form transfers D from predicted noisy trajectory to A^T(q*phi).
 38    weak_kernel = torch.tensor(A.T @ (q*phi), device=dev)
 39    identity_kernel = torch.tensor(q*phi, device=dev)
 40    for ep in range(EPOCHS):
 41        gen = torch.Generator(device='cpu').manual_seed(seed*1000+ep)
 42        order = torch.randperm(len(x), generator=gen, device='cpu')
 43        net.train()
 44        for ids in order.split(BATCH):
 45            xb, yb = x[ids], y[ids]
 46            pred = net(xb)
 47            data_loss = F.mse_loss(pred, yb)
 48            if weak:
 49                # Use observed window states as the local field and predicted next
 50                # state as the boundary/source term; this makes the intervention
 51                # a genuine training loss, not a post-processing readout.
 52                traj = xb.view(-1,8,3)[:,:,0]
 53                weak_time = (traj * weak_kernel).sum(1, keepdim=True)
 54                weak_id = (traj * identity_kernel).sum(1, keepdim=True)
 55                # dynamics residual: weak fractional temporal change against
 56                # predicted terminal state (scaled consistently across systems).
 57                residual = weak_time - 0.25*weak_id - 0.05*pred
 58                loss = data_loss + 0.15 * residual.square().mean()
 59            else:
 60                # Standard strong local residual: pointwise fractional derivative
 61                # of the measured trajectory at the final collocation point.
 62                traj = xb.view(-1,8,3)[:,:,0]
 63                strong = (traj * torch.tensor(A[-1],device=dev)).sum(1,keepdim=True)
 64                residual = strong - 0.25*traj[:,-1:] - 0.05*pred
 65                loss = data_loss + 0.15 * residual.square().mean()
 66            opt.zero_grad(); loss.backward(); opt.step()
 67    net.eval()
 68    with torch.no_grad():
 69        pred = net(d['xte'].to(dev)); mse = F.mse_loss(pred,d['yte'].to(dev)).item()
 70        # Model-behaviour signature: prediction vs observed trajectory-derived
 71        # strong/weak features, measured on trained weights.
 72        te = d['xte'].to(dev).view(-1,8,3)[:,:,0]
 73        wk = (te*weak_kernel).sum(1,keepdim=True)
 74        sk = (te*torch.tensor(A[-1],device=dev)).sum(1,keepdim=True)
 75        corr_w = float(torch.corrcoef(torch.stack([pred[:,0],wk[:,0]]))[0,1].cpu())
 76        corr_s = float(torch.corrcoef(torch.stack([pred[:,0],sk[:,0]]))[0,1].cpu())
 77    return mse, {'weak_corr':corr_w, 'strong_corr':corr_s}
 78
 79def permutation(diffs):
 80    diffs=np.asarray(diffs); obs=diffs.mean(); count=0; total=2**len(diffs)
 81    for bits in itertools.product([-1,1], repeat=len(diffs)):
 82        if np.mean(diffs*np.asarray(bits)) <= obs+1e-12: count += 1
 83    return count/total
 84
 85def block(results, best_lr):
 86    vals=[r['mse'] for r in results if r['lr']==best_lr]
 87    return {'best_hyperparams':{'lr':best_lr,'epochs':EPOCHS}, 'sweep':results,
 88            'full':{'per_seed':vals,'mean':float(np.mean(vals)), 'std':float(np.std(vals))}}
 89
 90def main():
 91    base_all=[]; idea_all=[]; signatures=[]
 92    # Baseline sweep covers the union of all idea learning rates.
 93    for lr in LRS:
 94        for s in SEEDS:
 95            d=get_dataset('dynamics',s,n_train=400,n_test=200)
 96            m, sig=train(d,lr,False,seed=s); base_all.append({'lr':lr,'seed':s,'mse':m})
 97    means={lr:np.mean([z['mse'] for z in base_all if z['lr']==lr]) for lr in LRS}
 98    best_lr=min(means,key=means.get)
 99    base_block=block(base_all,best_lr)
100    for lr in LRS:
101        for s in SEEDS:
102            d=get_dataset('dynamics',s,n_train=400,n_test=200)
103            m,sig=train(d,lr,True,alpha=0.8,seed=s); idea_all.append({'lr':lr,'seed':s,'mse':m}); signatures.append(sig)
104    im={lr:np.mean([z['mse'] for z in idea_all if z['lr']==lr]) for lr in LRS}; idea_lr=min(im,key=im.get)
105    vals=[z['mse'] for z in idea_all if z['lr']==idea_lr]
106    idea={'best_hyperparams':{'lr':idea_lr,'alpha':0.8,'epochs':EPOCHS},'sweep':idea_all,
107          'per_seed':vals,'mean':float(np.mean(vals)),'std':float(np.std(vals))}
108    sig={'trained_model_behavior':{'mean_weak_prediction_correlation':float(np.mean([x['weak_corr'] for x in signatures])), 'mean_strong_prediction_correlation':float(np.mean([x['strong_corr'] for x in signatures]))}, 'prediction':'weak residual should couple to trajectory while reducing derivative noise amplification','confirmed':bool(np.mean([x['weak_corr'] for x in signatures]) >= np.mean([x['strong_corr'] for x in signatures]))}
109    rep=make_report('dynamics','rnn_small',base_block,idea,extra=sig)
110    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
111if __name__=='__main__': main()