.. title:: Talk: asdex
.. only:: html
asdex: Sparse AD in JAX
=======================
Speaker: Adrian Hill, BIFOLD, Technische Universität Berlin
Abstract
^^^^^^^^
JAX is increasingly popular in scientific computing, yet its API for
materializing Jacobian and Hessian matrices is strictly dense, with
``jax.jacfwd``, ``jax.jacrev``, and ``jax.hessian`` all performing one AD pass
per input or output dimension. For the large sparse systems that arise in
nonlinear solvers and second-order optimization, this is prohibitive in both
memory and compute. Sparse automatic differentiation (i.e., compressed
evaluation) avoids both bottlenecks by exploiting the sparsity of the derivative
matrices, but no general-purpose implementation exists for JAX. We present
``asdex``, an open-source package that fills this gap with sparse drop-in
replacements for JAX's Jacobian and Hessian functions.
``asdex`` detects Jacobian sparsity patterns by abstract interpretation of
jaxprs, JAX's functional intermediate representation. Propagating index sets
through the program yields conservative global patterns without modifying user
code. Since ``jax.grad`` is a function transformation producing an ordinary
jaxpr, detecting the Jacobian sparsity of the gradient yields Hessian sparsity
patterns for free. Greedy distance-2 and star coloring algorithms then group
structurally orthogonal columns or rows of the detected pattern, so that one
batched JVP or VJP per group recovers every nonzero. Following the design of
``DifferentiationInterface.jl`` in Julia, ``asdex`` separates one-time
preparation from repeated evaluation, amortizing the cost of sparsity detection
and coloring. Prepared derivative functions are plain JAX functions that compose
with ``jax.jit`` and ``jax.vmap`` and run on CPU, GPU, and TPU.