Adjoint-Weak Fractional Residuals / stage2_bench.py
Mechanism confirmed, baseline not beaten
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()