Differentiable Physics-Equilibrium Projection / implicit_layer.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 1"""Small differentiable equilibrium projection for F(z; u)=z^3+a z-u."""
 2import torch
 3
 4
 5def newton_project(u, a, z0, steps=12, damping=1.0):
 6    """Fixed-step Newton projection; gradients can flow through the iterations."""
 7    z = z0
 8    for _ in range(steps):
 9        f = z**3 + a*z - u
10        j = 3*z**2 + a
11        z = z - damping*f/j
12    return z
13
14
15class ImplicitCubic(torch.autograd.Function):
16    """Solve F=0 forward and use dz/du=1/J, dz/da=-z/J backward."""
17    @staticmethod
18    def forward(ctx, u, a, z0, steps=20):
19        with torch.no_grad():
20            z = newton_project(u, a, z0, steps=steps)
21        ctx.save_for_backward(z, a)
22        return z
23
24    @staticmethod
25    def backward(ctx, grad_out):
26        z, a = ctx.saved_tensors
27        j = 3*z**2 + a
28        return grad_out/j, (-grad_out*z/j), None, None
29
30
31def implicit_project(u, a, z0=None, steps=20):
32    if z0 is None:
33        z0 = torch.zeros_like(u)
34    return ImplicitCubic.apply(u, a, z0, steps)