asdex: Automatic Sparse Differentiation in JAX
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 $m \times n$ Jacobian requires $n$ forward-mode or $m$ 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 ...
Description / Details
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.
Source: arXiv:2610.12336v1 - http://arxiv.org/abs/2610.12336v1 PDF: https://arxiv.org/pdf/2610.12336v1 Original Link: http://arxiv.org/abs/2610.12336v1
Please sign in to join the discussion.
No comments yet. Be the first to share your thoughts!
Oct 9, 2026
Mathematics
Mathematics
0