Skew-Midpoint Neural Dynamics / nonlinear_check.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import torch
 2from skew_midpoint import make_metric, skew, midpoint_step
 3
 4torch.set_default_dtype(torch.float64)
 5torch.manual_seed(11)
 6d = 4
 7L = torch.tril(torch.randn(d, d)).requires_grad_()
 8with torch.no_grad():
 9    L.diagonal().copy_(torch.tensor([1.2, 1.0, 0.8, 1.1]))
10
11class ANet(torch.nn.Module):
12    def __init__(self):
13        super().__init__()
14        self.w = torch.nn.Parameter(torch.randn(d, d))
15        self.v = torch.nn.Parameter(torch.randn(d))
16
17    def forward(self, x):
18        return self.w * torch.tanh(x[:, None] + self.v[None, :])
19
20net = ANet()
21e = torch.randn(d, requires_grad=True)
22dt = 0.35
23z = midpoint_step(e, dt, L, net, newton_iters=10)
24M = make_metric(L)
25mid = (z + e) / 2
26R = M @ ((z - e) / dt) - skew(net(mid)) @ mid
27H = lambda x: 0.5 * x @ M @ x
28loss = (z ** 2).sum() + 0.1 * H(z)
29grads = torch.autograd.grad(loss, tuple(net.parameters()) + (L,))
30print({
31    'residual_norm': torch.linalg.norm(R).item(),
32    'energy_change': abs((H(z) - H(e)).item()),
33    'parameter_gradient_norms': [torch.linalg.norm(q).item() for q in grads],
34    'finite': bool(torch.isfinite(z).all() and torch.isfinite(R).all())
35})