import sys, json, random import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, sweep_baseline, make_report class PHCell(nn.Module): """Explicit-Euler port-Hamiltonian hidden transition.""" def __init__(self, inp=3, hidden=64, dt=0.05, damping=0.2, coupling=0.15): super().__init__() if hidden % 2: raise ValueError('hidden must be even') self.hidden, self.q, self.dt, self.coupling = hidden, hidden // 2, dt, coupling self.input = nn.Linear(inp, hidden) self.state = nn.Linear(hidden, hidden, bias=False) self.formation = nn.Parameter(torch.randn(self.q, self.q) * 0.03) self.log_damping = nn.Parameter(torch.full((hidden,), np.log(np.expm1(damping)))) def forward(self, x, h, return_terms=False): force = torch.tanh(self.input(x) + self.state(h)) B = self.formation ga, gb = force[..., :self.q], force[..., self.q:] # [0,-B^T; B,0] acting on gradient; use force as input-dependent gradient. jforce = torch.cat((-gb @ B, ga @ B.T), dim=-1) grad = h + force damping = torch.nn.functional.softplus(self.log_damping) dh = self.coupling * jforce - damping * grad hn = h + self.dt * dh if return_terms: return hn, 0.5 * h.square().sum(-1), dh, damping return hn class PHModel(nn.Module): def __init__(self, out_dim=1, hidden=64, dt=0.05, damping=0.2, coupling=0.15): super().__init__(); self.cell=PHCell(3, hidden, dt, damping, coupling); self.head=nn.Linear(hidden,out_dim) def forward(self, x): seq=x.view(x.shape[0],-1,3) h=torch.zeros(x.shape[0],self.cell.hidden,device=x.device,dtype=x.dtype) for k in range(seq.shape[1]): h=self.cell(seq[:,k],h) return self.head(h) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def run_one(seed, lr, idea): seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=200) shape=tuple(d['xtr'].shape[1:]) net=PHModel() if idea else make_model('rnn_small',shape,1) net, metric, history=train_model(net,d,epochs=25,lr=lr,batch=128) return float(metric),net,d def signature(seed=0, lr=0.003): metric, net, d=run_one(seed,lr,True); net.eval() device=next(net.parameters()).device; x=d['xte'][:64].to(device) with torch.no_grad(): seq=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],64,device=device) before=[]; after=[]; powers=[] for k in range(seq.shape[1]): old=h; h,e,dh,_=net.cell(seq[:,k],h,return_terms=True) before.append(float(e.mean())); after.append(float((0.5*h.square().sum(-1)).mean())) powers.append(float((old*dh).sum(-1).mean())) inc=float(max(np.asarray(after)-np.asarray(before))) return {'energy_first_mean':before[0],'energy_last_mean':after[-1], 'max_step_energy_increase':inc,'mean_continuous_power':float(np.mean(powers)), 'prediction':'bounded hidden energy under damped port-Hamiltonian updates', 'confirmed':bool(inc <= 1e-2 and np.isfinite(metric))} def main(): # Shared union: every idea lr is also swept by the baseline. grid = [{'lr': 0.001}, {'lr': 0.003}, {'lr': 0.009}] base_block = sweep_baseline( lambda cfg: lambda seed: run_one(seed, cfg['lr'], False)[0], grid=grid) # Evaluate every idea setting on the same eight paired seeds; choose its best. idea_grid = [] for cfg in grid: vals = [run_one(seed, cfg['lr'], True)[0] for seed in range(8)] idea_grid.append({'cfg': cfg, 'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals)}) best_idea = min(idea_grid, key=lambda z: z['mean']) idea_res = {'best_cfg': best_idea['cfg'], 'sweep': [ {'cfg': z['cfg'], 'mean': z['mean']} for z in idea_grid], 'per_seed': best_idea['per_seed'], 'mean': best_idea['mean'], 'std': best_idea['std'], 'n': best_idea['n']} sig = signature(0, best_idea['cfg']['lr']) report = make_report('dynamics', 'rnn_small', base_block, idea_res, extra={'mechanism_signature': sig, 'track_justification': 'Dynamics is structurally matched: the idea claims stability/control of repeated state propagation.'}) with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()