Skew-Midpoint Neural Dynamics / skew_midpoint.py
Mechanism confirmed, baseline not beaten
1import torch
2
3
4def make_metric(L, eps=1e-4):
5 return L @ L.transpose(-1, -2) + eps * torch.eye(L.shape[-1], device=L.device, dtype=L.dtype)
6
7
8def skew(A):
9 return A - A.transpose(-1, -2)
10
11
12def midpoint_step(e, dt, L, A_fn, B=None, u=None, newton_iters=6, damping=1.0, eps=1e-4):
13 """Differentiable implicit-midpoint step for M(z-e)/dt=J(mid)mid+B u.
14
15 This MVP handles one latent vector (not a batch); A_fn maps [d] to [d,d].
16 Newton iterations remain in the autograd graph so losses can train parameters.
17 """
18 M = make_metric(L, eps)
19 z = e.clone()
20 for _ in range(newton_iters):
21 if not z.requires_grad:
22 z = z.requires_grad_(True)
23 mid = (z + e) / 2
24 J = skew(A_fn(mid))
25 forcing = torch.zeros_like(e) if B is None or u is None else B @ u
26 r = M @ ((z - e) / dt) - J @ mid - forcing
27 jac_rows = []
28 for i in range(e.numel()):
29 jac_rows.append(torch.autograd.grad(r[i], z, retain_graph=True, create_graph=True)[0])
30 jac = torch.stack(jac_rows)
31 z = z - damping * torch.linalg.solve(jac, r)
32 return z
33
34
35def linear_midpoint_step(e, dt, J):
36 n = e.numel()
37 I = torch.eye(n, dtype=e.dtype, device=e.device)
38 return torch.linalg.solve(I - .5 * dt * J, (I + .5 * dt * J) @ e)