Power-Preserving Formation GNN / bench_power_formation.py
Failed on benchmark
1import sys, json, random
2import numpy as np
3import torch
4from torch import nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
8
9
10class PHCell(nn.Module):
11 """Explicit-Euler port-Hamiltonian hidden transition."""
12 def __init__(self, inp=3, hidden=64, dt=0.05, damping=0.2, coupling=0.15):
13 super().__init__()
14 if hidden % 2: raise ValueError('hidden must be even')
15 self.hidden, self.q, self.dt, self.coupling = hidden, hidden // 2, dt, coupling
16 self.input = nn.Linear(inp, hidden)
17 self.state = nn.Linear(hidden, hidden, bias=False)
18 self.formation = nn.Parameter(torch.randn(self.q, self.q) * 0.03)
19 self.log_damping = nn.Parameter(torch.full((hidden,), np.log(np.expm1(damping))))
20
21 def forward(self, x, h, return_terms=False):
22 force = torch.tanh(self.input(x) + self.state(h))
23 B = self.formation
24 ga, gb = force[..., :self.q], force[..., self.q:]
25 # [0,-B^T; B,0] acting on gradient; use force as input-dependent gradient.
26 jforce = torch.cat((-gb @ B, ga @ B.T), dim=-1)
27 grad = h + force
28 damping = torch.nn.functional.softplus(self.log_damping)
29 dh = self.coupling * jforce - damping * grad
30 hn = h + self.dt * dh
31 if return_terms:
32 return hn, 0.5 * h.square().sum(-1), dh, damping
33 return hn
34
35
36class PHModel(nn.Module):
37 def __init__(self, out_dim=1, hidden=64, dt=0.05, damping=0.2, coupling=0.15):
38 super().__init__(); self.cell=PHCell(3, hidden, dt, damping, coupling); self.head=nn.Linear(hidden,out_dim)
39 def forward(self, x):
40 seq=x.view(x.shape[0],-1,3)
41 h=torch.zeros(x.shape[0],self.cell.hidden,device=x.device,dtype=x.dtype)
42 for k in range(seq.shape[1]): h=self.cell(seq[:,k],h)
43 return self.head(h)
44
45
46def seed_all(seed):
47 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
48 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
49
50
51def run_one(seed, lr, idea):
52 seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=200)
53 shape=tuple(d['xtr'].shape[1:])
54 net=PHModel() if idea else make_model('rnn_small',shape,1)
55 net, metric, history=train_model(net,d,epochs=25,lr=lr,batch=128)
56 return float(metric),net,d
57
58
59def signature(seed=0, lr=0.003):
60 metric, net, d=run_one(seed,lr,True); net.eval()
61 device=next(net.parameters()).device; x=d['xte'][:64].to(device)
62 with torch.no_grad():
63 seq=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],64,device=device)
64 before=[]; after=[]; powers=[]
65 for k in range(seq.shape[1]):
66 old=h; h,e,dh,_=net.cell(seq[:,k],h,return_terms=True)
67 before.append(float(e.mean())); after.append(float((0.5*h.square().sum(-1)).mean()))
68 powers.append(float((old*dh).sum(-1).mean()))
69 inc=float(max(np.asarray(after)-np.asarray(before)))
70 return {'energy_first_mean':before[0],'energy_last_mean':after[-1],
71 'max_step_energy_increase':inc,'mean_continuous_power':float(np.mean(powers)),
72 'prediction':'bounded hidden energy under damped port-Hamiltonian updates',
73 'confirmed':bool(inc <= 1e-2 and np.isfinite(metric))}
74
75
76def main():
77 # Shared union: every idea lr is also swept by the baseline.
78 grid = [{'lr': 0.001}, {'lr': 0.003}, {'lr': 0.009}]
79 base_block = sweep_baseline(
80 lambda cfg: lambda seed: run_one(seed, cfg['lr'], False)[0], grid=grid)
81
82 # Evaluate every idea setting on the same eight paired seeds; choose its best.
83 idea_grid = []
84 for cfg in grid:
85 vals = [run_one(seed, cfg['lr'], True)[0] for seed in range(8)]
86 idea_grid.append({'cfg': cfg, 'mean': float(np.mean(vals)),
87 'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals)})
88 best_idea = min(idea_grid, key=lambda z: z['mean'])
89 idea_res = {'best_cfg': best_idea['cfg'], 'sweep': [
90 {'cfg': z['cfg'], 'mean': z['mean']} for z in idea_grid],
91 'per_seed': best_idea['per_seed'], 'mean': best_idea['mean'],
92 'std': best_idea['std'], 'n': best_idea['n']}
93 sig = signature(0, best_idea['cfg']['lr'])
94 report = make_report('dynamics', 'rnn_small', base_block, idea_res,
95 extra={'mechanism_signature': sig,
96 'track_justification':
97 'Dynamics is structurally matched: the idea claims stability/control of repeated state propagation.'})
98 with open('bench_report.json', 'w') as f:
99 json.dump(report, f, indent=2)
100 print(json.dumps(report, indent=2))
101
102if __name__ == '__main__':
103 main()