Power-Preserving Formation GNN / bench_power_formation.py

Failed on benchmark

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