asdex: Automatic Sparse Differentiation in JAX
Organizations: BIFOLD – Berlin Institute for the Foundations of Learning and Data, Berlin, Germany · Machine Learning Group, Technical University of Berlin, Berlin, Germany · LVMT, ENPC, Institut Polytechnique de Paris, Univ Gustave Eiffel, Marne-la-Vallée, France
Abstract
Many tasks in scientific computing and machine learning require the Jacobian or Hessian matrix of a function. Automatic differentiation (AD) computes these derivatives to machine precision, but materializing a dense Jacobian requires forward-mode or reverse-mode AD passes, one per column or row. For a large class of functions, each output depends on only a few inputs, making the derivative matrix sparse. Automatic sparse differentiation (ASD) exploits this structure in four steps: detection of the input-agnostic sparsity pattern, coloring of a graph to group columns or rows that can share an AD pass, compressed differentiation to compute a compressed derivative matrix with one AD pass per color, and finally decompression into the original sparsity pattern. The number of colors, and hence of AD passes, is often independent of the problem dimension: a banded Jacobian with contiguous bands, for instance, only ever requires colors, regardless of its size. asdex offers the first standalone ASD toolkit in the popular JAX ecosystem. With asdex.jacobian and asdex.hessian, it provides sparse drop-in replacements for jax.jacobian and jax.hessian.
Figures & tables
| Package | Detection | Coloring |
|---|---|---|
| sparsejac 1 1 1 https://github.com/mfschubert/sparsejac | None | Distance-1 (from networkx ) |
| sparsediffax 2 2 2 https://github.com/gdalle/sparsediffax | None | Distance-2 & star (from SMC) |
| jax-nansparse 3 3 3 https://github.com/nardi/jax-nansparse | NaN tracing | None |
| jax2sympy 4 4 4 https://github.com/johnviljoen/jax2sympy | jaxpr symbolic conversion | None |
| asdex | jaxpr tracing | Distance-2 & star (based on SMC) |