Skip to content

Commit ae71c2a

Browse files
authored
Merge pull request #6 from pyatsysh/docs/auto-differentiable-jax
Docs: auto-differentiable (JAX) namespace — API reference + tutorial
2 parents fe24e3b + 4749c6e commit ae71c2a

12 files changed

Lines changed: 266 additions & 0 deletions

File tree

source-docs/index.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,5 +41,6 @@ Code documentation for :code:`equadratures` can be found by clicking on the link
4141

4242
Core classes <core/index>
4343
Secondary classes <secondary/index>
44+
Auto-differentiable (JAX) <jax/index>
4445
Theory <theory/index>
4546
Tutorials <tutorials/index>

source-docs/jax/basis.txt

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
Basis and design matrix
2+
==================
3+
4+
.. automodule:: equadratures.jax.basis
5+
:members:

source-docs/jax/index.txt

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
:orphan:
2+
3+
.. _jax:
4+
5+
============================
6+
Auto-differentiable (JAX)
7+
============================
8+
9+
The :code:`equadratures.jax` namespace is a JAX-native, differentiable backend for
10+
``equadratures``. Quadrature, orthogonal-polynomial approximation, uncertainty
11+
quantification and learnable polynomial kernels are all ``jax.grad``-able,
12+
``jax.jit``-able and ``jax.vmap``-able, so models and kernels can be trained by
13+
gradient descent. It is a *parallel namespace*: the classic NumPy API is
14+
unchanged.
15+
16+
Install the optional dependency with ``pip install equadratures[jax]`` (add
17+
``equadratures[jax-learn]`` for the Optax-based kernel training). See the
18+
:doc:`tutorial <../tutorials/tutorials/Auto_Differentiable_Equadratures>` for a
19+
worked introduction.
20+
21+
Code documentation for the individual modules is linked below.
22+
23+
24+
.. toctree::
25+
:maxdepth: 1
26+
27+
parameter
28+
recurrence
29+
quadrature
30+
polynomials
31+
basis
32+
poly
33+
kernels
34+
35+
36+
.. toctree:
37+
:hidden:
38+
39+
../documentation

source-docs/jax/kernels.txt

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
Learnable polynomial kernels
2+
==================
3+
4+
.. automodule:: equadratures.jax.kernel
5+
:members:

source-docs/jax/parameter.txt

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
Parameter
2+
==================
3+
4+
.. automodule:: equadratures.jax.parameter
5+
:members:

source-docs/jax/poly.txt

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
Polynomial model
2+
==================
3+
4+
.. automodule:: equadratures.jax.poly
5+
:members:

source-docs/jax/polynomials.txt

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
Orthonormal polynomials
2+
==================
3+
4+
.. automodule:: equadratures.jax.polynomial
5+
:members:

source-docs/jax/quadrature.txt

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
Gauss quadrature
2+
==================
3+
4+
.. automodule:: equadratures.jax.quadrature
5+
:members:

source-docs/jax/recurrence.txt

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
Recurrence coefficients
2+
==================
3+
4+
.. automodule:: equadratures.jax.recurrence
5+
:members:

source-docs/requirements.txt

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,3 +10,7 @@ sphinx-gallery
1010
sphinxcontrib-applehelp==1.0.2
1111
sphinxcontrib-bibtex
1212
sphinxcontrib-devhelp==1.0.2
13+
14+
# for the equadratures.jax (auto-differentiable) docs + tutorial
15+
jax
16+
optax

0 commit comments

Comments
 (0)