import torch def make_metric(L, eps=1e-4): return L @ L.transpose(-1, -2) + eps * torch.eye(L.shape[-1], device=L.device, dtype=L.dtype) def skew(A): return A - A.transpose(-1, -2) def midpoint_step(e, dt, L, A_fn, B=None, u=None, newton_iters=6, damping=1.0, eps=1e-4): """Differentiable implicit-midpoint step for M(z-e)/dt=J(mid)mid+B u. This MVP handles one latent vector (not a batch); A_fn maps [d] to [d,d]. Newton iterations remain in the autograd graph so losses can train parameters. """ M = make_metric(L, eps) z = e.clone() for _ in range(newton_iters): if not z.requires_grad: z = z.requires_grad_(True) mid = (z + e) / 2 J = skew(A_fn(mid)) forcing = torch.zeros_like(e) if B is None or u is None else B @ u r = M @ ((z - e) / dt) - J @ mid - forcing jac_rows = [] for i in range(e.numel()): jac_rows.append(torch.autograd.grad(r[i], z, retain_graph=True, create_graph=True)[0]) jac = torch.stack(jac_rows) z = z - damping * torch.linalg.solve(jac, r) return z def linear_midpoint_step(e, dt, J): n = e.numel() I = torch.eye(n, dtype=e.dtype, device=e.device) return torch.linalg.solve(I - .5 * dt * J, (I + .5 * dt * J) @ e)