Unrolling Implicit Layers Is a Memory Nightmare — Here's How to Avoid It
I spent about three weeks debugging a custom autograd pass for a deep equilibrium model before I accepted that unrolling was going to eat my GPU memory and make training absurdly slow. The fix is implicit differentiation, but not the textbook version most people implement. It's more of a practical tradeoff between storing intermediate values and solving an adjoint equation, and the details matter more than the high-level description. Let's just get the definition out of the way, since it keeps getting stated in ways that obscure what's actually happening. Circuit Training Implicit Differentiation refers to computing gradients through a neural network that contains implicit layers — layers whose output is defined as the fixed point of a function rather than through a direct forward pass. The "circuit training" part is somewhat overloaded; some people use it to describe training procedures where the same fixed-point circuit is unrolled multiple times, others use it more broadly for any training loop that solves an implicit equation during the forward pass and then differentiates through that solve. The core setup: you have an implicit layer defined by the equation x = F(x, ), where F is some neural network and are learnable parameters. The output x is the fixed point. A naive approach would find x by iterating until convergence, store every iterate, and then backpropagate through the entire unrolled sequence. This uses O(k) memory where k is the number of iterations. That's not viable past about 50 iterations on anything resembling a real model.
Implicit differentiation avoids that by computing the gradient of the loss L with respect to without storing the intermediate iterates. You solve for the fixed point x once, then solve a separate linear system to get dx/d. The adjoint equation is (I - F/x)¹ L/x, and you never form the full inverse explicitly. This drops memory usage from O(k) to roughly O(1) for the forward pass, though you do pay for an extra solve per backward step. The practical catch is that F/x is rarely tractable to form as a full matrix. What you actually do is implement a matrix-free linear solve. You give an iterative solver — conjugate gradient or GMRES — a function that computes the product of (I - F/x) with an arbitrary vector. That product is computed via reverse-mode automatic differentiation on F itself. PyTorch's torch.func.vjp does exactly this if you're using functional API transforms, or you can roll your own with torch.autograd.functional.vjp. The solver converges in far fewer steps than the fixed-point iteration, usually under ten iterations for well-conditioned problems. Here's where I ran into something that isn't covered in the papers. If your implicit layer uses a heavy regularization term or a contractive mapping that's only barely contractive, the spectral radius of F/x can sit uncomfortably close to 1. The conjugate gradient solver then struggles. I had a Deep Equilibrium Model where the loss landscape caused the fixed point to drift toward a region where the Jacobian had eigenvalues near 0.98, and CG needed hundreds of iterations to converge, which effectively killed training speed. The workaround was to add an explicit contraction enforcement term — (||x - F(x, )||²) — to the loss during the early training phases. It's ugly but it kept the spectral radius down and the solver usable. Once the model stabilized, you can reduce . This is documented nowhere explicitly; it's just what happens when you try to run these things on real data instead of toy examples.
Another thing nobody emphasizes enough: the choice of forward solver matters for the backward pass. Broyden's method or a simple fixed-point iteration with a good initial guess will find the fixed point faster, but if you restart the fixed-point solve differently between the forward and the matrix-vector product computation inside the adjoint solve, you can end up at slightly different fixed points. The gradient becomes inconsistent. I learned this when my training loss dropped while the validation loss spiked, and the discrepancy was traced back to the fact that my forward pass used 30 fixed-point iterations while the adjoint solve's internal matrix products were triggering a different number of iterations due to a slightly different stopping criterion. Aligning the tolerances and iteration counts between forward and backward resolved it. For the implementation, here's what I'd actually recommend rather than what the papers suggest. Use a library that already handles the implicit layer abstraction if you can. PyTorch's torch.nn.modules implicitly supports this pattern, and libraries like implicit-layer-toolbox or the DEQ implementations from rubinstein et al. handle the vjp-based adjoint solve correctly. Writing this yourself is fine for understanding the mechanics, but the edge cases around numerical stability in the linear solve are easy to get wrong. If you're implementing from scratch, structure it as follows. Define a custom autograd Function. In the forward method, solve for the fixed point x using your preferred iterative method and store only x and the parameters . In the backward method, compute the vjp of F with respect to its first argument at (x, ), then use an iterative linear solver to compute (I - J)¹ times the upstream gradient, where J is the vjp operator. Multiply the result by the vjp of F with respect to to get dL/d. That's it. Don't try to form J explicitly. Don't try to invert (I - J) directly.
Get the Full Details

There are limitations that are worth stating plainly. Implicit layers with implicit differentiation do not help when your model has many sequential implicit layers stacked on top of each other. Each layer requires its own adjoint solve, and the memory cost compounds. For a single implicit layer in an otherwise standard architecture, the approach is solid. For a stack of ten implicit layers, you're better off reconsidering the architecture or accepting the memory cost of unrolling. Similarly, if the fixed point is not unique or is difficult to find reliably — which happens frequently with poorly conditioned mappings — the whole method becomes unstable. You'll get NaN gradients without much warning, because the linear solve will silently diverge or converge to the wrong solution. A common alternative when implicit differentiation starts causing more problems than it solves is to just use explicit unrolling with checkpointing. Reversible networks or activation checkpointing can reduce the memory cost of storing intermediate values to something manageable, and the gradients are numerically more stable because you're differentiating through a finite unrolled computation rather than through an iterative solver. For models that don't strictly require the fixed-point formulation, this is often the more pragmatic choice. The one area where implicit differentiation with circuit training genuinely earns its keep is when you need to train a model whose depth is effectively infinite — things like neural ODEs, equilibrium models, or controllers that solve an optimization problem at each time step. In those cases, the alternative of unrolling is impossible by definition, and the adjoint method is the only option. Just be aware that "the only option" doesn't mean "the easy option." It means you're trading memory for computational overhead in the backward pass, and you're responsible for making sure that the linear solve is stable and accurate enough for your use case.
If you want to experiment, start small. Train a single implicit layer on MNIST with a contractive mapping, verify that the gradient matches a finite-difference approximation, and then scale up. I found that the finite-difference check catches about 90% of implementation bugs before they become expensive to debug at full scale. The remaining 10% tend to be the numerical edge cases I mentioned above, and those only surface when you're actually training on real data with a real loss landscape.