"""Small differentiable equilibrium projection for F(z; u)=z^3+a z-u.""" import torch def newton_project(u, a, z0, steps=12, damping=1.0): """Fixed-step Newton projection; gradients can flow through the iterations.""" z = z0 for _ in range(steps): f = z**3 + a*z - u j = 3*z**2 + a z = z - damping*f/j return z class ImplicitCubic(torch.autograd.Function): """Solve F=0 forward and use dz/du=1/J, dz/da=-z/J backward.""" @staticmethod def forward(ctx, u, a, z0, steps=20): with torch.no_grad(): z = newton_project(u, a, z0, steps=steps) ctx.save_for_backward(z, a) return z @staticmethod def backward(ctx, grad_out): z, a = ctx.saved_tensors j = 3*z**2 + a return grad_out/j, (-grad_out*z/j), None, None def implicit_project(u, a, z0=None, steps=20): if z0 is None: z0 = torch.zeros_like(u) return ImplicitCubic.apply(u, a, z0, steps)