Efficient and Modular Implicit Differentiation

Mathieu Blondel, Quentin Berthet, Marco Cuturi, Roy Frostig, Stephan Hoyer, Felipe Llinares-López, Fabian Pedregosa, Jean-Philippe Vert

Introduction

Automatic differentiation (autodiff) is now an inherent part of machine learning software. It allows to express complex computations by composing elementary ones in creative ways and removes the tedious burden of computing their derivatives by hand. In parallel, the differentiation of optimization problem solutions has found many applications. A classical example is bi-level optimization, which typically involves computing the derivatives of a nested optimization problem in order to solve an outer one. Examples of applications in machine learning include hyper-parameter optimization , neural networks , and meta-learning . Another line of active research involving differentiation of optimization problem solutions are optimization layers , which can be used to encourage structured outputs, and implicit deep networks , which have a smaller memory footprint than backprop-trained networks.

Since optimization problem solutions typically do not enjoy an explicit formula in terms of their inputs, autodiff cannot be used directly to differentiate these functions. In recent years, two main approaches have been developed to circumvent this problem. The first one consists of unrolling the iterations of an optimization algorithm and using the final iteration as a proxy for the optimization problem solution . This allows to explicitly construct a computational graph relating the algorithm output to the inputs, on which autodiff can then be used transparently. However, this requires a reimplementation of the algorithm using the autodiff system, and not all algorithms are necessarily autodiff friendly. Moreover, forward-mode autodiff has time complexity that scales linearly with the number of variables and reverse-mode autodiff has memory complexity that scales linearly with the number of algorithm iterations. In contrast, a second approach consists in implicitly relating an optimization problem solution to its inputs using optimality conditions. In a machine learning context, such implicit differentiation has been used for stationarity conditions , KKT conditions and the proximal gradient fixed point . An advantage of implicit differentiation is that a solver reimplementation is not needed, allowing to build upon decades of state-of-the-art software. Although implicit differentiation has a long history in numerical analysis , so far, it remained difficult to use for practitioners, as it required a case-by-case tedious mathematical derivation and implementation. CasADi allows to differentiate various optimization and root finding problem algorithms provided by the library. However, it does not allow to easily add implicit differentiation on top of existing solvers from optimality conditions expressed by the user, as we do. A recent tutorial explains how to implement implicit differentiation in JAX . However, the tutorial requires the user to take care of low-level technical details and does not cover a large catalog of optimality condition mappings as we do. Other work attempts to address this issue by adding implicit differentiation on top of cvxpy . This works by reducing all convex optimization problems to a conic program and using conic programming’s optimality conditions to derive an implicit differentiation formula. While this approach is very generic, solving a convex optimization problem using a conic programming solver—an ADMM-based splitting conic solver in the case of cvxpy—is rarely state-of-the-art for every problem instance.

In this work, we ambition to achieve for optimization problem solutions what autodiff did for computational graphs. We propose automatic implicit differentiation, a simple approach to add implicit differentiation on top of any existing solver. In this approach, the user defines directly in Python a mapping function FF capturing the optimality conditions of the problem solved by the algorithm. Once this is done, we leverage autodiff of FF combined with the implicit function theorem to automatically differentiate the optimization problem solution. Our approach is generic, yet it can exploit the efficiency of state-of-the-art solvers. It therefore combines the benefits of implicit differentiation and autodiff. To summarize, we make the following contributions.

We describe our framework and its JAX implementation (https://github.com/google/jaxopt/). Our framework significantly lowers the barrier to use implicit differentiation, thanks to the seamless integration in JAX, with low-level details all abstracted away.

We instantiate our framework on a large catalog of optimality conditions (Table 1), recovering existing schemes and obtaining new ones, such as the mirror descent fixed point based one.

On the theoretical side, we provide new bounds on the Jacobian error when the optimization problem is only solved approximately, and empirically validate them.

We implement four illustrative applications, demonstrating our framework’s ease of use.

Beyond our software implementation in JAX, we hope this paper provides a self-contained blueprint for creating an efficient and modular implementation of implicit differentiation in other frameworks.

Automatic implicit differentiation

Contrary to autodiff through unrolled algorithm iterations, implicit differentiation typically involves a manual, sometimes complicated, mathematical derivation. For instance, numerous works use Karush–Kuhn–Tucker (KKT) conditions in order to relate a constrained optimization problem’s solution to its inputs, and to manually derive a formula for its derivatives. The derivation and implementation in these works are typically case-by-case.

In this work, we propose a generic way to easily add implicit differentiation on top of existing solvers. In our approach, the user defines directly in Python a mapping function FF capturing the optimality conditions of the problem solved by the algorithm. We provide reusable building blocks to easily express such FF. The provided FF is then plugged into our Python decorator @custom_root, which we append on top of the solver declaration we wish to differentiate. Under the hood, we combine the implicit function theorem and autodiff of FF to automatically differentiate the optimization problem solution. A simple illustrative example is given in Figure 1.

Differentiating a root.

Computing ∂x⋆(θ)\partial x^{\star}(\theta) therefore boils down to the resolution of the linear system of equations

When (1) is a one-dimensional root finding problem (d=1d=1), (3) becomes particularly simple since we then have ∇x⋆(θ)=B⊤/A\nabla x^{\star}(\theta)=B^{\top}/A, where AA is a scalar value.

We will show that existing and new implicit differentiation methods all reduce to this simple principle. We call our approach automatic implicit differentiation as the user can freely express the optimization solution to be differentiated through the optimality conditions FF. Our approach is efficient as it can be added on top of any state-of-the-art solver and modular as the optimality condition specification is decoupled from the implicit differentiation mechanism. This contrasts with existing works, where the derivation and implementation are specific to each optimality condition.

Differentiating a fixed point.

We will encounter numerous applications where x⋆(θ)x^{\star}(\theta) is instead implicitly defined through a fixed point:

In this case, when TT is continuously differentiable, using the chain rule, we have

Computing JVPs and VJPs.

In most practical scenarios, it is not necessary to explicitly form the Jacobian matrix, and instead it is sufficient to left-multiply or right-multiply by ∂1F\partial_{1}F and ∂2F\partial_{2}F. These are called vector-Jacobian product (VJP) and Jacobian-vector product (JVP), and are useful for integrating x⋆(θ)x^{\star}(\theta) with reverse-mode and forward-mode autodiff, respectively. Oftentimes, FF will be explicitly defined. In this case, computing the VJP or JVP can be done via autodiff. In some cases, FF may itself be implicitly defined, for instance when FF involves the solution of a variational problem. In this case, computing the VJP or JVP will itself involve implicit differentiation.

The right-multiplication (JVP) between J=∂x⋆(θ)J=\partial x^{\star}(\theta) and a vector vv, JvJv, can be computed efficiently by solving A(Jv)=BvA(Jv)=Bv. The left-multiplication (VJP) of v⊤v^{\top} with JJ, v⊤Jv^{\top}J, can be computed by first solving A⊤u=vA^{\top}u=v. Then, we can obtain v⊤Jv^{\top}J by v⊤J=u⊤AJ=u⊤Bv^{\top}J=u^{\top}AJ=u^{\top}B. Note that when BB changes but AA and vv remain the same, we do not need to solve A⊤u=vA^{\top}u=v once again. This allows to compute the VJP w.r.t. different variables while solving only one linear system.

To solve these linear systems, we can use the conjugate gradient method when AA is symmetric positive semi-definite and GMRES or BiCGSTAB otherwise. These algorithms are all matrix-free: they only require matrix-vector products. Thus, all we need from FF is its JVPs or VJPs. An alternative to GMRES/BiCGSTAB is to solve the normal equation AA⊤u=AvAA^{\top}u=Av using conjugate gradient. This can be implemented using JAX’s transpose routine jax.linear_transpose . In case of non-invertibility, a common heuristic is to solve a least squares min⁡J∥AJ−B∥2\min_{J}\|AJ-B\|^{2} instead.

Pre-processing and post-processing mappings.

Oftentimes, the goal is not to differentiate θ\theta per se, but the parameters of a function producing θ\theta. One example of such pre-processing is to convert the parameters to be differentiated from one form to another canonical form, such as a quadratic program or a conic program . Another example is when x⋆(θ)x^{\star}(\theta) is used as the output of a neural network layer, in which case θ\theta is produced by the previous layer. Likewise, x⋆(θ)x^{\star}(\theta) will often not be the final output we want to differentiate. One example of such post-processing is when x⋆(θ)x^{\star}(\theta) is the solution of a dual program and we apply the dual-primal mapping to recover the solution of the primal program. Another example is the application of a loss function, in order to reduce x⋆(θ)x^{\star}(\theta) to a scalar value. We leave the differentiation of such pre/post-processing mappings to the autodiff system, allowing to compose functions in complex ways.

Implementation details.

When a solver function is decorated with @custom_root, we use jax.custom_jvp and jax.custom_vjp to automatically add custom JVP and VJP rules to the function, overriding JAX’s default behavior. As mentioned above, we use linear system solvers based on matrix-vector products and therefore we only need access to FF through the JVP or VJP with ∂1F\partial_{1}F and ∂2F\partial_{2}F. This is done by using jax.jvp and jax.vjp, respectively. Note that, as in Figure 1, the definition of FF will often include a gradient mapping ∇1f(x,θ)\nabla_{1}f(x,\theta). Thankfully, JAX supports second-order derivatives transparently. For convenience, our library also provides a @custom_fixed_point decorator, for adding implicit differentiation on top of a solver, given a fixed point iteration TT; see code examples in Appendix B.

2 Examples

We now give various examples of mapping FF or fixed point iteration TT, recovering existing implicit differentiation methods and creating new ones. Each choice of FF or TT implies different trade-offs in terms of computational oracles; see Table 1. Source code examples are given in Appendix B.

The simplest example is to differentiate through the implicit function

We then have ∂1F(x,θ)=∇12f(x,θ)\partial_{1}F(x,\theta)=\nabla^{2}_{1}f(x,\theta) and ∂2F(x,θ)=∂2∇1f(x,θ)\partial_{2}F(x,\theta)=\partial_{2}\nabla_{1}f(x,\theta), the Hessian of ff in its first argument and the Jacobian in the second argument of ∇1f(x,θ)\nabla_{1}f(x,\theta). In practice, we use autodiff to compute Jacobian products automatically. Equivalently, we can use the gradient descent fixed point

for all η>0\eta>0. Using (5), it is easy to check that we obtain the same linear system since η\eta cancels out.

KKT conditions.

As a more advanced example, we now show that the KKT conditions, manually differentiated in several works , fit our framework. As we will see, the key will be to group the optimal primal and dual variables as our x⋆(θ)x^{\star}(\theta). Let us consider the general problem

Proximal gradient fixed point.

Unfortunately, not all algorithms return both primal and dual solutions. Moreover, if the objective contains non-smooth terms, proximal gradient descent may be more efficient. We now discuss its fixed point . Let x⋆(θ)x^{\star}(\theta) be implicitly defined as

To implicitly differentiate x⋆(θ)x^{\star}(\theta), we use the fixed point mapping [69, p.150]

for any step size η>0\eta>0. The proximity operator is 11-Lipschitz continuous . By Rademacher’s theorem, it is differentiable almost everywhere. If, in addition, it is continuously differentiable in a neighborhood of (x⋆(θ),θ)(x^{\star}(\theta),\theta) and if I−∂1T(x⋆(θ),θ)I-\partial_{1}T(x^{\star}(\theta),\theta) is invertible, then our framework to differentiate x⋆(θ)x^{\star}(\theta) applies. Similar assumptions are made in . Many proximity operators enjoy a closed form and can easily be differentiated, as discussed in Appendix C. An implementation is given in Figure 2.

Projected gradient fixed point.

As a special case, when g(x,θ)g(x,\theta) is the indicator function IC(θ)(x)I_{{\mathcal{C}}(\theta)}(x), where C(θ){\mathcal{C}}(\theta) is a convex set depending on θ\theta, we obtain

The proximity operator proxg{\text{prox}}_{g} becomes the Euclidean projection onto C(θ){\mathcal{C}}(\theta)

and (16) becomes the projected gradient fixed point

Compared to the KKT conditions, this fixed point is particularly suitable when the projection enjoys a closed form. We discuss how to compute the JVP / VJP for a wealth of convex sets in Appendix C.

Current limitations.

While we have not observed issues in practice, we note that the approach developed in this section theoretically only applies to settings where the implicit function theorem is valid, namely, where optimality conditions satisfy the differentiability and invertibility conditions stated in §2.1. While this covers a wide range of situations even for non-smooth optimization problems (e.g., under mild assumptions the solution of a Lasso regression can be differentiated a.e. with respect to the regularization parameter, see Appendix E), an interesting direction for future work is to extend the framework to handle cases where the differentiability and invertibility conditions are not satisfied, using, e.g., the theory of nonsmooth implicit function theorems .

Jacobian precision guarantees

In practice, either by the limitations of finite precision arithmetic or because we perform a finite number of iterations, we rarely reach the exact solution x⋆(θ)x^{\star}(\theta). Instead, we reach an approximate solution x^\hat{x} and apply the implicit differentiation equation (3) at this approximate solution. This motivates the need for precision guarantees of this approach. We introduce the following formalism.

It holds by construction that J(x⋆(θ),θ)=∂x⋆(θ)J(x^{\star}(\theta),\theta)=\partial x^{\star}(\theta). Computing J(x^,θ)J(\hat{x},\theta) for an approximate solution x^\hat{x} of x⋆(θ)x^{\star}(\theta) therefore allows to approximate the true Jacobian ∂x⋆(θ)\partial x^{\star}(\theta). In practice, an algorithm used to solve (1) depends on θ\theta. Note however that, what we compute is not the Jacobian of x^(θ)\hat{x}(\theta), unlike works differentiating through unrolled algorithm iterations, but an estimate of ∂x⋆(θ)\partial x^{\star}(\theta). We therefore use the notation x^\hat{x}, leaving the dependence on θ\theta implicit.

We develop bounds of the form ∥J(x^,θ)−∂x⋆(θ)∥<C∥x^−x⋆(θ)∥\|J(\hat{x},\theta)-\partial x^{\star}(\theta)\|<C\|\hat{x}-x^{\star}(\theta)\|, hence showing that the error on the estimated Jacobian is at most of the same order as that of x^\hat{x} as an approximation of x⋆(θ)x^{\star}(\theta). These bounds are based on the following main theorem, whose proof is included in Appendix D.

AA is well-conditioned, Lipschitz: ∥A(x,θ)v∥≥α∥v∥\|A(x,\theta)v\|\geq\alpha\|v\| , ∥A(x,θ)−A(x⋆(θ),θ)∥op≤γ∥x−x⋆(θ)∥\|A(x,\theta)-A(x^{\star}(\theta),\theta)\|_{\textnormal{op}}\leq\gamma\|x-x^{\star}(\theta)\|.

BB is bounded and Lipschitz: ∥B(x⋆(θ),θ)∥≤R\|B(x^{\star}(\theta),\theta)\|\leq R , ∥B(x,θ)−B(x⋆(θ),θ)∥≤β∥x−x⋆(θ)∥\|B(x,\theta)-B(x^{\star}(\theta),\theta)\|\leq\beta\|x-x^{\star}(\theta)\|.

Under these conditions, when ∥x^−x⋆(θ)∥≤ε\|\hat{x}-x^{\star}(\theta)\|\leq\varepsilon, we have

This result is inspired by [52, Theorem 7.2], that is concerned with the stability of solutions to inverse problems. As a difference, we consider that A(⋅,θ)A(\cdot,\theta) is uniformly well-conditioned, rather than only at x⋆(θ)x^{\star}(\theta). This does not affect the first order in ε\varepsilon of this bound, and makes it valid for all x^\hat{x}. Our goal with Theorem 1 is to provide a result that works for general FF but can be tailored to specific cases.

In particular, for the gradient descent fixed point (9), this yields

By specializing Theorem 1 for this fixed point, we obtain Jacobian precision guarantees with conditions directly on ff rather than FF; see Corollary 1 in Appendix D. These guarantees hold for instance for the dataset distillation experiment in Section 4. Our analysis reveals in particular that Jacobian estimation by implicit differentiation gains a factor of t\mathbf{t} compared to automatic differentiation, after tt iterations of gradient descent in the strongly-convex setting [1, Proposition 3.2]. While our guarantees concern the Jacobian of x⋆(θ)x^{\star}(\theta), we note that other studies give guarantees on hypergradients (i.e., the gradient of an outer objective).

We illustrate these results on ridge regression, where x⋆(θ)=argmin⁡x∥Φx−y∥22+∑iθixi2x^{\star}(\theta)=\operatorname*{argmin}_{x}\|\Phi x-y\|^{2}_{2}+\sum_{i}\theta_{i}x_{i}^{2}. This problem has the merit that the solution x⋆(θ)x^{\star}(\theta) and its Jacobian ∂x⋆(θ)\partial x^{\star}(\theta) are available in closed form. By running gradient descent for tt iterations, we obtain an estimate x^\hat{x} of x⋆(θ)x^{\star}(\theta) and an estimate J(x^,θ)J(\hat{x},\theta) of ∂x⋆(θ)\partial x^{\star}(\theta); cf. Definition 1. By doing so for different numbers of iterations tt, we can graph the relation between the error ∥x⋆(θ)−x^∥2\|x^{\star}(\theta)-\hat{x}\|_{2} and the error ∥∂x⋆(θ)−J(x^,θ)∥2\|\partial x^{\star}(\theta)-J(\hat{x},\theta)\|_{2}, as shown in Figure 3, empirically validating Theorem 1. The results in Figure 3 were obtained using the diabetes dataset from , with other datasets yielding a qualitatively similar behavior. We derive similar guarantees in Corollary 2 in Appendix D for proximal gradient descent.

Experiments

In this section, we demonstrate the ease of solving bi-level optimization problems with our framework. We also present an application to the sensitivity analysis of molecular dynamics.

While KKT conditions can be used to differentiate x⋆(θ)x^{\star}(\theta), a more direct way is to use the projected gradient fixed point (19). The projection onto C{\mathcal{C}} can be easily computed by row-wise projections on the simplex. The projection’s Jacobian enjoys a closed form (Appendix C). Another way to differentiate x⋆(θ)x^{\star}(\theta) is using the mirror descent fixed point (28). Under the KL geometry, projections correspond to a row-wise softmax. They are therefore easy to compute and differentiate. Figure 4 compares the runtime performance of implicit differentiation vs. unrolling for the latter two fixed points.

2 Dataset distillation

In this problem, and unlike in the general hyperparameter optimization setup, both the inner and outer problems are high-dimensional, making it an ideal test-bed for gradient-based bi-level optimization methods. For this experiment, we use the MNIST dataset. The number of parameters in the inner problem is p=282=784p=28^{2}=784. while the number of parameters of the outer loss is k×p=7840k\times p=7840. We solve this problem using gradient descent on both the inner and outer problem, with the gradient of the outer loss computed using implicit differentiation, as described in §2. This is fundamentally different from the approach used in the original paper, where they used differentiation of the unrolled iterates instead. For the same solver, we found that the implicit differentiation approach was 4 times faster than the original one. The obtained distilled images θ\theta are visualized in Figure 5.

3 Task-driven dictionary learning

We illustrate this on breast cancer survival prediction from gene expression data. We frame it as a binary classification problem to discriminate patients who survive longer than 5 years (m1=200m_{1}=200) vs patients who die within 5 years of diagnosis (m0=99m_{0}=99), from p=1,000p=1,000 gene expression values. As shown in Table 2, solving (22) (Task-driven DictL) reaches a classification performance competitive with state-of-the-art L1L_{1} or L2L_{2} regularized logistic regression with 100 times fewer variables.

4 Sensitivity analysis of molecular dynamics

Conclusion

We proposed in this paper an approach for automatic implicit differentiation, allowing the user to freely express the optimality conditions of the optimization problem whose solutions are to be differentiated, directly in Python. The applicability of our approach to a large catalog of optimality conditions is shown in the non-exhaustive list of Table 1, and illustrated by the ease with which we can solve bi-level and sensitivity analysis problems.

References

Appendix A More examples of optimality criteria and fixed points

To demonstrate the generality of our approach, we describe in this section more optimality mapping FF or fixed point iteration TT.

We define the Bregman projection of yy onto C(θ)⊆dom⁡(φ){\mathcal{C}}(\theta)\subseteq\operatorname*{dom}(\varphi) by

Newton fixed point.

Let xx be a root of G(⋅,θ)G(\cdot,\theta), i.e., G(x,θ)=0G(x,\theta)=0. The fixed point iteration of Newton’s method for root-finding is

Using (5), we get A=−∂1F(x,θ)=ηIA=-\partial_{1}F(x,\theta)=\eta I. Similarly,

Newton’s method for optimization is obtained by choosing G(x,θ)=∇1f(x,θ)G(x,\theta)=\nabla_{1}f(x,\theta), which gives

It is easy to check that we recover the same linear system as for the gradient descent fixed point (9). A practical implementation can pre-compute an LU decomposition of ∂1G(x,θ)\partial_{1}G(x,\theta), or a Cholesky decomposition if ∂1G(x,θ)\partial_{1}G(x,\theta) is positive semi-definite.

Proximal block coordinate descent fixed point.

We now consider the case when x⋆(θ)x^{\star}(\theta) is implicitly defined as the solution

where g1,…,gmg_{1},\dots,g_{m} are possibly non-smooth functions operating on subvectors (blocks) x1,…,xmx_{1},\dots,x_{m} of xx. In this case, we can use for i∈[m]i\in[m] the fixed point

where η1,…,ηm\eta_{1},\dots,\eta_{m} are block-wise step sizes. Clearly, when the step sizes are shared, i.e., η1=⋯=ηm=η\eta_{1}=\dots=\eta_{m}=\eta, this fixed point is equivalent to the proximal gradient fixed point (16) with g(x,θ)=∑i=1ngi(xi,θ)g(x,\theta)=\sum_{i=1}^{n}g_{i}(x_{i},\theta).

Quadratic programming.

We now show how to use the KKT conditions discussed in §2.2 to differentiate quadratic programs, recovering Optnet as a special case. To give some intuition, let us start with a simple equality-constrained quadratic program (QP)

In matrix notation, this can be rewritten as

We can write the solution of the linear system (38) as the root x=(z,ν)x=(z,\nu) of a function F(x,θ)F(x,\theta). More generally, the QP can also include inequality constraints

In matrix notation, this can be written as

While x=(z,ν,λ)x=(z,\nu,\lambda) is no longer the solution of a linear system, it is the root of a function F(x,θ)F(x,\theta) and therefore fits our framework. With our framework, no derivation is needed. We simply define ff, HH and GG directly in Python.

Conic programming.

We now show that the differentiation of conic linear programs , at the heart of differentiating through cvxpy layers , easily fits our framework. Consider the problem

where N=p+m+1N=p+m+1. Following , we can use the homogeneous self-dual embedding to reduce the process of solving (44) to finding a root of the residual map

The key oracle whose JVP/VJP we need is therefore Π\Pi, which is studied in . The projection onto a few cones is available in our library and can be used to express FF.

Frank-Wolfe.

where C(θ){\mathcal{C}}(\theta) is a convex polytope, i.e., it is the convex hull of vertices v1(θ),…,vm(θ)v_{1}(\theta),\dots,v_{m}(\theta). The Frank-Wolfe algorithm requires a linear minimization oracle (LMO)

and is a popular algorithm when this LMO is easier to compute than the projection onto C(θ){\mathcal{C}}(\theta). However, since this LMO is piecewise constant, its Jacobian is null almost everywhere. Inspired by SparseMAP , which corresponds to the case when ff is a quadratic, we rewrite (48) as

where V(θ)V(\theta) is a d×md\times m matrix gathering the vertices v1(θ),…,vm(θ)v_{1}(\theta),\dots,v_{m}(\theta). We then have x⋆(θ)=V(θ)p⋆(θ)x^{\star}(\theta)=V(\theta)p^{\star}(\theta). Since we have reduced (48) to minimization over the simplex, we can use the projected gradient fixed point to obtain

We can therefore compute the derivatives of p⋆(θ)p^{\star}(\theta) by implicit differentiation and the derivatives of x⋆(θ)x^{\star}(\theta) by product rule. Frank-Wolfe implementations typically maintain the convex weights of the vertices, which we use to get an approximation of p⋆(θ)p^{\star}(\theta). Moreover, it is well-known that after tt iterations, at most tt vertices are visited. We can leverage this sparsity to solve a smaller linear system. Moreover, in practice, we only need to compute VJPs of x⋆(θ)x^{\star}(\theta).

Appendix B Code examples

Our library provides several reusable optimality condition mappings FF or fixed points TT. We nevertheless demonstrate the ease of writing some of them from scratch.

As a more advanced example, we now describe how to implement the KKT conditions (13). The stationarity, primal feasibility and complementary slackness conditions read

Using jax.vjp to compute vector-Jacobian products, this can be implemented as

Similar mappings FF can be written if the optimization problem contains only equality constraints or only inequality constraints.

Mirror descent fixed point.

Letting η=1\eta=1 and denoting θ=(θf,θproj)\theta=(\theta_{f},\theta_{{\text{proj}}}), the fixed point (28) is

Although not considered in this example, the mapping ∇φ\nabla\varphi could also depend on θ\theta if necessary.

B.2 Code examples for experiments

We now sketch how to implement our experiments using our framework. In the following, jnp is short for jax.numpy. In all experiments, we only show how to compute gradients with the outer objective. We can then use these gradients with gradient-based solvers to solve the outer objective.

Task-driven dictionary learning experiment.

Dataset distillation experiment.

Molecular dynamics experiment.

Appendix C Jacobian products

Our library provides numerous reusable building blocks. We describe in this section how to compute their Jacobian products. As a general guideline, whenever a projection enjoys a closed form, we leave the Jacobian product to the autodiff system.

We describe in this section how to compute the Jacobian products of the projections (in the Euclidean and KL senses) onto various convex sets. When the convex set does not depend on any variable, we simply denote it C{\mathcal{C}} instead of C(θ){\mathcal{C}}(\theta).

Box constraints.

Probability simplex.

When C{\mathcal{C}} is the standard probability simplex, C=△d{\mathcal{C}}=\triangle^{d}, there is no analytical solution for projC(y){\text{proj}}_{\mathcal{C}}(y). Nevertheless, the projection can be computed exactly in O(d)O(d) expected time or O(dlog⁡d)O(d\log d) worst-case time . The Jacobian is given by diag(s)−ss⊤/∥s∥1\text{diag}(s)-ss^{\top}/\|s\|_{1}, where s∈{0,1}ds\in\{0,1\}^{d} is a vector indicating the support of projC(y){\text{proj}}_{\mathcal{C}}(y) . The projection in the KL sense, on the other hand, enjoys a closed form: it reduces to the usual softmax projCφ(y)=exp⁡(y)/∑j=1dexp⁡(yj){\text{proj}}^{\varphi}_{\mathcal{C}}(y)=\exp(y)/\sum_{j=1}^{d}\exp(y_{j}).

Box sections.

The root can be found, e.g., by bisection. The gradient ∇x⋆(θ)\nabla x^{\star}(\theta) is given by ∇x⋆(θ)=B⊤/A\nabla x^{\star}(\theta)=B^{\top}/A and the Jacobian ∂z⋆(θ)\partial z^{\star}(\theta) is obtained by application of the chain rule on LL.

Norm balls.

Affine sets.

where A†A^{\dagger} is the Moore-Penrose pseudoinverse of AA. The second equality holds if p<dp<d and AA is full rank. A practical implementation can pre-compute a factorization of the Gram matrix AA⊤AA^{\top}. Alternatively, we can also use the KKT conditions.

Hyperplanes and half spaces.

Transportation and Birkhoff polytopes.

Order simplex.

Polyhedra.

C.2 Jacobian products of proximity operators

We provide several proximity operators, including for the lasso (soft thresholding), elastic net and group lasso (block soft thresholding). All satisfy closed form expressions and can be differentiated automatically via autodiff. For more advanced proximity operators, which do not enjoy a closed form, recent works have derived their Jacobians. The Jacobians of fused lasso and OSCAR were derived in . For general total variation, the Jacobians were derived in .

Appendix D Jacobian precision proofs

To simplify notations, we note A⋆≔A(x⋆,θ)A_{\star}\coloneqq A(x^{\star},\theta) and A^≔A(x^,θ)\hat{A}\coloneqq A(\hat{x},\theta), and similarly for BB and JJ. We have by definition of the Jacobian estimate function A⋆J⋆=B⋆A_{\star}J_{\star}=B_{\star} and A^J^=B^\hat{A}\hat{J}=\hat{B}. Therefore we have

For any invertible matrices M1,M2M_{1},M_{2}, it holds that M1−1−M2−1=M1−1(M2−M1)M2−1M_{1}^{-1}-M_{2}^{-1}=M_{1}^{-1}(M_{2}-M_{1})M_{2}^{-1}, so

As a consequence, the second term in J(x^,θ)−∂x⋆(θ)J(\hat{x},\theta)-\partial x^{\star}(\theta) can be upper bounded and we obtain

Let ff be such that f(⋅,θ)f(\cdot,\theta) is twice differentiable and α\alpha-strongly convex and ∇12f(⋅,θ)\nabla_{1}^{2}f(\cdot,\theta) is γ\gamma-Lipschitz (in the operator norm) and ∂2∇1f(x,θ)\partial_{2}\nabla_{1}f(x,\theta) is β\beta-Lipschitz and bounded in norm by RR. The estimated Jacobian evaluated at x^\hat{x} is then given by

This follows from Theorem 1, applied to this specific A(x,θ)A(x,\theta) and B(x,θ)B(x,\theta). ∎

For proximal gradient descent, where T(x,θ)=proxηg(x−η∇1f(x,θ),θ)T(x,\theta)={\text{prox}}_{\eta g}(x-\eta\nabla_{1}f(x,\theta),\theta), this yields

We now focus in the case of proximal gradient descent on an objective f(x,θ)+g(x)f(x,\theta)+g(x), where gg is smooth and does not depend on θ\theta. This is the case in our experiments in §4.3. Recent work also exploits local smoothness of solutions to derive similar bounds [13, Theorem 13]

First, let us note that proxηg(y,θ){\text{prox}}_{\eta g}(y,\theta) does not depend on θ\theta, since gg itself does not depend on θ\theta, and is therefore equal to classical proximity operator of ηg\eta g which, with a slight overload of notations, we denote as proxηg(y){\text{prox}}_{\eta g}(y) (with a single argument). In other words,

Regarding the first claim (expression of the estimated Jacobian evaluated at x^\hat{x}), we first have that proxηg(y){\text{prox}}_{\eta g}(y) is the solution to (x′−y)+η∇g(x′)=0(x^{\prime}-y)+\eta\nabla g(x^{\prime})=0 in x′x^{\prime} - by first-order condition for a smooth convex function. We therefore have that

the first II and inverse being functional identity and inverse, and the second IdI_{d} and inverse being in the matrix sense, by inverse rule for Jacobians ∂h(z)=[∂h−1(h(z))]−1\partial h(z)=[\partial h^{-1}(h(z))]^{-1} (applied to the prox).

As a consequence, we have, for Γη(x,θ)=∇2g(proxηg(x−η∇1f(x,θ))\Gamma_{\eta}(x,\theta)=\nabla^{2}g({\text{prox}}_{\eta g}(x-\eta\nabla_{1}f(x,\theta)) that

In the following, we modify slightly the notation of both AA and BB, writing

Appendix E The Lasso case

Our approach to differentiate the solution of a root equation F(x,θ)=0F(x,\theta)=0 is valid as long as the smooth implicit function theorem holds, namely, as long as FF is continuously differentiable near a solution (x0,θ0)(x_{0},\theta_{0}) and ∇1F(x0,θ0)\nabla_{1}F(x_{0},\theta_{0}) is invertible. While the first assumption is easy to check when FF is continuously differentiable everywhere, it does not always hold when this is not the case, e.g., when FF involves the proximity operator of a non-smooth function. In such cases, one may therefore have to study theoretically the properties of the function FF near the solutions (x(θ),θ)(x(\theta),\theta) to justify differentiation using the smooth implicit function theorem. Here, we develop such an analysis to justify the use of our approach to differentiate the solution of a Lasso regression problem with respect to the regularization parameter. We note that the smooth implicit function theorem has already been used for this problem ; here we justify why it is a valid approach, even though FF itself is not continuously differentiable everywhere. More precisely, we consider the Lasso problem:

We first show that FηF_{\eta} is continuously differentiable in a neighborhood of (x⋆(θ),θ)(x^{\star}(\theta),\theta), for any θ\theta that is not a kink. Since the ST operator is continuously differentiable everywhere except on the closed set:

We thus need to show that (x⋆(θ)+ηeθγ,ηeθ)∉S(x^{\star}(\theta)+\eta e^{\theta}\gamma,\eta e^{\theta})\notin{\mathcal{S}}. We first see that, for any i∈[1,d]i\in[1,d] such that x⋆(θ)i≠0x^{\star}(\theta)_{i}\neq 0,

The case x⋆(θ)i≠0x^{\star}(\theta)_{i}\neq 0 requires more care, since the property γi∈\gamma_{i}\in is not sufficient to show that ∣x⋆(θ)i+ηeθγi∣=ηeθ∣γi∣|x^{\star}(\theta)_{i}+\eta e^{\theta}\gamma_{i}|=\eta e^{\theta}|\gamma_{i}| is not equal to ηeθ\eta e^{\theta}: we need to show that, in fact, ∣γi∣<1|\gamma_{i}|<1. For that purpose, let E(θ)={i∈[1,d] : ∣γi∣=1}{\mathcal{E}}(\theta)=\left\{i\in[1,d]\,:\,|\gamma_{i}|=1\right\}. Denoting ΦE(θ)\Phi_{{\mathcal{E}}(\theta)} the matrix made of the columns of Φ\Phi in E(θ){\mathcal{E}}(\theta), we know that, under the assumptions of Theorem 2, with probability one the matrix ΦE(θ)⊤ΦE(θ)\Phi_{{\mathcal{E}}(\theta)}^{\top}\Phi_{{\mathcal{E}}(\theta)} is invertible and the lasso problem has a unique solution given by x⋆(θ)E(θ)C=0x^{\star}(\theta)_{{\mathcal{E}}(\theta)^{C}}=0 and

where s(θ)=sign(Φ⊤(b−Φx⋆(θ))))∈{−1,0,1}ds(\theta)={\text{sign}}(\Phi^{\top}\left(b-\Phi x^{\star}(\theta))\right))\in\{-1,0,1\}^{d} . Furthermore, we know that E(θ){\mathcal{E}}(\theta) is constant between two successive kinks , so if x⋆(θ)x^{\star}(\theta) is not a kink then there is a neighborhood [θ1,θ2][\theta_{1},\theta_{2}] of θ\theta such as E(θ′)=E(θ){\mathcal{E}}(\theta^{\prime})={\mathcal{E}}(\theta) and s(θ′)E(θ′)=s(θ)E(θ)s(\theta^{\prime})_{{\mathcal{E}}(\theta^{\prime})}=s(\theta)_{{\mathcal{E}}(\theta)}, for any θ′∈[θ1,θ2]\theta^{\prime}\in[\theta_{1},\theta_{2}]. Let us now assume that Φ\Phi is such that for any E⊂[1,d]{\mathcal{E}}\subset[1,d] and s∈{−1,1}∣E∣s\in\{-1,1\}^{|{\mathcal{E}}|}, ΦE⊤ΦE\Phi_{\mathcal{E}}^{\top}\Phi_{\mathcal{E}} is invertible and (ΦE⊤ΦE)−1s(\Phi_{\mathcal{E}}^{\top}\Phi_{\mathcal{E}})^{-1}s has no coordinate equal to zero. This happens with probability one under the assumptions of Theorem 2 since the set of singular matrices is measure zero. Then we see from (71) that, for θ′∈[θ1,θ2]\theta^{\prime}\in[\theta_{1},\theta_{2}] and i∈E(θ)i\in{\mathcal{E}}(\theta), x⋆(θ′)ix^{\star}(\theta^{\prime})_{i} is an affine and non-constant function of eθ′e^{\theta^{\prime}}. Since in addition x⋆(θ1)ix^{\star}(\theta_{1})_{i} and x⋆(θ2)ix^{\star}(\theta_{2})_{i} are either both nonnegative or nonpositive, then necessarily x⋆(θ)ix^{\star}(\theta)_{i} is positive or negative, respectively. In other words, we have shown that ∣γi∣=1  ⟹  x⋆(θ)i≠0|\gamma_{i}|=1\implies x^{\star}(\theta)_{i}\neq 0, or equivalently that x⋆(θ)i=0  ⟹  ∣γi∣<1x^{\star}(\theta)_{i}=0\implies|\gamma_{i}|<1. From this we deduce that for any i∈[1,d]i\in[1,d] such that x⋆(θ)i=0x^{\star}(\theta)_{i}=0,

This concludes the proof that (x⋆(θ)+ηeθγ,ηeθ)∉S(x^{\star}(\theta)+\eta e^{\theta}\gamma,\eta e^{\theta})\notin{\mathcal{S}}, and therefore that FηF_{\eta} is continuously differentiable in a neighborhood of (x⋆(θ),θ)(x^{\star}(\theta),\theta). The second condition for the smooth implicit theorem to hold, namely, the invertibility of ∇1Fη(x⋆(θ),θ)\nabla_{1}F_{\eta}(x^{\star}(\theta),\theta), is easily obtained by explicit computation [12, 13, Proposition 1] ∎

Appendix F Experimental setup and additional results

Our experiments use JAX , which is Apache2-licensed and scikit-learn , which is BSD-licensed.

For the outer problem, gradient descent was used with a stepsize of 5⋅10−35\cdot 10^{-3} for the first 100100 steps, following a inverse square root decay afterwards up to a total of 150150 steps.

Conjugate gradient was used to solve the linear systems in implicit differentiation for at most 25002500 iterations.

All results reported pertaining CPU runtimes were obtained using an internal compute cluster. GPU results were obtained using a single NVIDIA P100 GPU with 16GB of memory per dataset. For each dataset size, we report the average runtime of an individual iteration in the outer problem, alongside a 90% confidence interval estimated from the corresponding 150150 runtime values.

Additional results

Figure 13 compares the runtime of implicit differentiation and unrolling on GPU. These results highlight a fundamental limitation of the unrolling approach in memory-limited systems such as accelerators, as the inner solver suffered from out-of-memory errors for most problem sizes (p≥2000p\geq 2000 for mirror descent, p≥750p\geq 750 for proximal gradient and block coordinate descent). While it might be possible to ameliorate this limitation by reducing the maximum number of iterations in the inner solver, doing so might lead to additional challenges and require careful tuning.

Figure 14 depicts the validation loss (value of the outer problem objective function) at convergence. It shows that all approaches were able to solve the outer problem, with solutions produced by different approaches being qualitatively indistinguishable from each other across the range of problem sizes considered.

Figure 15 shows the Jacobian error achieved as a function of the solution error, when varying the number of features.

F.2 Task-driven dictionary learning

We downloaded from http://acgt.cs.tau.ac.il/multi_omic_benchmark/download.html a set of breast cancer gene expression data together with survival information generated by the TCGA Research Network (https://www.cancer.gov/tcga) and processed as explained by . The gene expression matrix contains the expression value for p=20,531 genes in m=1,212 samples, from which we keep only the primary tumors (m=1,093). From the survival information, we select the patients who survived at least five years after diagnosis (m1=200m_{1}=200), and the patients who died before five years (m0=99m_{0}=99), resulting in a cohort of m=299m=299 patients with gene expression and binary label. Note that non-selected patients are those who are marked as alive but were not followed for 5 years.

To evaluate different binary classification methods on this cohort, we repeated 10 times a random split of the full cohort into a training (60%), validation (20%) and test (20%) sets. For each split and each method, 1) the method is trained with different parameters on the training set, 2) the parameter that maximizes the classification AUC on the validation set is selected, 3) the method is then re-trained on the union of the training and validation sets with the selected parameter, and 4) we measure the AUC of that model on the test set. We then report, for each method, the mean test AUC over the 10 repeats, together with a 95% confidence interval defined a mean ±\pm 1.96 ×\times standard error of the mean.

F.3 Dataset Distillation

For the inner problem, we used gradient descent with backtracking line-search, while for the outer problem we used gradient descent with momentum and a fixed step-size. The momentum parameter was set to 0.90.9 while the step-size was set to 11.

Figure 5 was produced after 4000 iterations of the outer loop on CPU (Intel(R) Xeon(R) Platinum P-8136 CPU @ 2.00GHz), which took 1h55. Unrolled differentiation took instead 8h:05 (44 times more) to run the same number of iterations. As can be seen in Figure 16, the output is the same in both approaches.

F.4 Molecular dynamics

Our experimental setup is adapted from the JAX-MD example notebook available at https://github.com/google/jax-md/blob/master/notebooks/meta_optimization.ipynb.

We emphasize that calculating the gradient of the total energy objective, f(x,θ)=∑ijU(xi,j,θ)f(x,\theta)=\sum_{ij}U(x_{i,j},\theta), with respect to the diameter θ\theta of the smaller particles, ∇1f(x,θ)\nabla_{1}f(x,\theta), does not require implicit differentiation or unrolling. This is because ∇1f(x,θ)=0\nabla_{1}f(x,\theta)=0 at x=x⋆(θ)x=x^{\star}(\theta):

This is known as Danskin’s theorem or envelope theorem. Thus instead, we consider sensitivities of position ∂x⋆(θ)\partial x^{\star}(\theta) directly, which does require implicit differentiation or unrolling.

Our results comparing implicit and unrolled differentiation for calculating the sensitivity of position are shown in Figure 17. We use BiCGSTAB to perform the tangent linear solve. Like in the original JAX-MD experiment, we use k=128k=128 particles in m=2m=2 dimensions.