Automatic Functional Differentiation in JAX
Min Lin
Abstract
We extend JAX with the capability to automatically differentiate higher-order functions (functionals and operators). By representing functions as a generalization of arrays, we seamlessly use JAX's existing primitive system to implement higher-order functions. We present a set of primitive operators that serve as foundational building blocks for constructing several key types of functionals. For every introduced primitive operator, we derive and implement both linearization and transposition rules, aligning with JAX's internal protocols for forward and reverse mode automatic differentiation. This enhancement allows for functional differentiation in the same syntax traditionally use for functions. The resulting functional gradients are themselves functions ready to be invoked in python. We showcase this tool's efficacy and simplicity through applications where functional derivatives are indispensable. The source code of this work is released at https://github.com/sail-sg/autofd .
Ask about this paper
Your agent reads all of it.
Lune indexed this paper to the last equation, along with the top-tier papers that cite it. Ask a question and the answer quotes them.
Your agent calls
Luneget_paper_fulltext
Free to start. No credit card required.
Terminal
Install the CLIlune papers fulltext 5e51a1e4-9942-4c61-b7e8-193091c0aa9dCited by top-tier papers1
Ask how each one uses itBuilds on4
- Fourier Neural Operator for Parametric Partial Differential EquationsZongyi Li, Nikola Borislavov Kovachki, Kamyar Azizzadenesheli, Burigede Liu et al.ICLR 2021 · 3,911 citations
- JAX MD: A Framework for Differentiable PhysicsSamuel S. Schoenholz, Ekin Dogus CubukNeurIPS 2020 · 195 citations
- You Only Linearize Once: Tangents Transpose to GradientsAlexey Radul, Adam Paszke, Roy Frostig, Matthew J. Johnson et al.POPL 2023 · 13 citations
- 𝜆ₛ: computable semantics for differentiable programming with higher-order functions and datatypesBenjamin Sherman, Jesse Michel, Michael CarbinPOPL 2021 · 11 citations
Related papers
- Provably correct, asymptotically efficient, higher-order reverse-mode automatic differentiationFaustyna Krawiec, Simon Peyton Jones, Neel Krishnaswami, Tom Ellis et al.POPL 2022 · 27 citations
- A simple differentiable programming languageMartín Abadi, Gordon D. PlotkinPOPL 2020 · 49 citations
- JAX Autodiff from a Linear Logic PerspectiveGiulia Giusti, Michele PaganiPOPL 2026
- SoftJAX & SoftTorch: Empowering Automatic Differentiation Libraries with Informative GradientsAnselm Paulus, Andreas René Geist, Vit Musil, Sebastian Hoffmann et al.ICML 2026 · 3 citations
- Compositional Taylor expansion in cartesian differential categoriesAymeric WalchLICS 2025
