rocm_jax/docs/jax.rst

67 lines
1.2 KiB
ReStructuredText
Raw Normal View History

.. currentmodule:: jax
jax package
===========
Subpackages
-----------
.. toctree::
:maxdepth: 1
jax.numpy
jax.scipy
jax.experimental
jax.lax
2019-08-29 17:51:15 -07:00
jax.nn
jax.ops
2019-02-13 19:31:41 -08:00
jax.random
2019-05-14 21:00:27 -04:00
jax.tree_util
jax.dlpack
2019-07-20 14:40:31 +01:00
Just-in-time compilation (:code:`jit`)
--------------------------------------
.. autofunction:: jit
.. autofunction:: disable_jit
.. autofunction:: xla_computation
.. autofunction:: make_jaxpr
.. autofunction:: eval_shape
2019-07-20 14:40:31 +01:00
Automatic differentiation
-------------------------
.. autofunction:: grad
.. autofunction:: value_and_grad
.. autofunction:: jacfwd
.. autofunction:: jacrev
.. autofunction:: hessian
.. autofunction:: jvp
.. autofunction:: linearize
.. autofunction:: vjp
.. autofunction:: custom_transforms
.. autofunction:: defjvp
.. autofunction:: defjvp_all
.. autofunction:: defvjp
.. autofunction:: defvjp_all
.. autofunction:: custom_gradient
2019-07-20 14:40:31 +01:00
Vectorization (:code:`vmap`)
----------------------------
.. autofunction:: vmap
2019-07-20 14:40:31 +01:00
Parallelization (:code:`pmap`)
------------------------------
2019-07-20 14:40:31 +01:00
.. autofunction:: pmap
2019-11-22 11:03:26 -08:00
.. autofunction:: devices
.. autofunction:: local_devices
.. autofunction:: host_id
.. autofunction:: host_ids
.. autofunction:: device_count
.. autofunction:: local_device_count
.. autofunction:: host_count