cs.MSOct 8, 2026

asdex: Automatic Sparse Differentiation in JAX

Authors: Adrian Hill, Guillaume Dalle

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 m×nm \times n Jacobian requires nn forward-mode or mm 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 bb contiguous bands, for instance, only ever requires bb 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

Explore similar work

Sep 17, 2026physics.chem-ph

Truncated automatic sparse differentiation for machine learning interatomic potentials

Machine learning interatomic potentials (MLIPs) learn the mapping from atomic positions to potential energy. The forces, the negative gradient of this energy, drive molecular dynamics and are readily obtained using automatic differentiation. Higher-order derivatives, most notably the Hessian, describe collective motion and allow the direct prediction of experimental observables, but are considered computationally inaccessible for large systems. We suggest a solution: in physical systems, interactions decay with distance, and most MLIPs build on this locality through message passing up to a finite receptive field. This implies both sparsity of higher-order derivatives and their decay with distance. This structure can be exploited using automatic sparse differentiation (ASD). We explain how to compute the sparsity pattern for MLIP derivatives and demonstrate that, for multiple foundation MLIPs, ASD computes full Hessians of large porous materials exactly, but with modest speedups at best. The larger gains come from truncated ASD: discarding small, but nonzero, Hessian entries between distant atoms yields order-of-magnitude speedups with negligible impact on predicted observables.
May 11, 2026cs.LG

jNO: A JAX Library for Neural Operator and Foundation Model Training

jNO (jax Neural Operators) is a JAX-native library for neural operators and foundation models with unified support for both data-driven and physics-informed training. Its core design is a tracing system in which domains, model calls, residuals, supervised losses, and diagnostics are written in one symbolic language and compiled into one optimization pipeline. This allows users to move between operator regression, mesh-aware residual evaluation, and PDE-constrained training without restructuring the surrounding code. jNO also supports multi-model compositions, fine-grained control at parameter level (model, optimizer, and learning rate), hyperparameter tuning, and JAX-native workflows for translated PDE foundation-model families. The source repository is available at https://github.com/FhG-IISB/jNO.
Sep 10, 2026cs.LG

AdamX: Cosine similarity meets gradient descent

We introduce AdamX, a first-order optimizer that incorporates cosine similarity as an adaptive mechanism for controlling update magnitudes. The proposed method is scalable, model-agnostic, and straightforward to integrate into existing training pipelines. We further introduce a variance rectification scheme that promotes smoother optimization during the early stages of training. Overall, we provide empirical evidence that AdamX achieves competitive convergence rates across a range of benchmark datasets and architectures. Performance is evaluated in terms of the number of epochs required to reach predefined performance thresholds under a fixed hyperparameter budget. Code and Experiments available at: https://github.com/FranciscoCaldas/adamX.