import torch from skew_midpoint import make_metric, skew, midpoint_step torch.set_default_dtype(torch.float64) torch.manual_seed(11) d = 4 L = torch.tril(torch.randn(d, d)).requires_grad_() with torch.no_grad(): L.diagonal().copy_(torch.tensor([1.2, 1.0, 0.8, 1.1])) class ANet(torch.nn.Module): def __init__(self): super().__init__() self.w = torch.nn.Parameter(torch.randn(d, d)) self.v = torch.nn.Parameter(torch.randn(d)) def forward(self, x): return self.w * torch.tanh(x[:, None] + self.v[None, :]) net = ANet() e = torch.randn(d, requires_grad=True) dt = 0.35 z = midpoint_step(e, dt, L, net, newton_iters=10) M = make_metric(L) mid = (z + e) / 2 R = M @ ((z - e) / dt) - skew(net(mid)) @ mid H = lambda x: 0.5 * x @ M @ x loss = (z ** 2).sum() + 0.1 * H(z) grads = torch.autograd.grad(loss, tuple(net.parameters()) + (L,)) print({ 'residual_norm': torch.linalg.norm(R).item(), 'energy_change': abs((H(z) - H(e)).item()), 'parameter_gradient_norms': [torch.linalg.norm(q).item() for q in grads], 'finite': bool(torch.isfinite(z).all() and torch.isfinite(R).all()) })