Skew-Midpoint Neural Dynamics / skew_midpoint.py

Mechanism confirmed, baseline not beaten

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