Skip to main content
Ctrl+K
JAX  documentation - Home JAX  documentation - Home

Getting started

  • Installation

Documentation

  • Newly documented
  • JAX 101: expressing computations
    • Arrays and jax.numpy
    • Transformations: grad and vmap
    • Pytrees
    • Pseudorandom numbers
    • Stateful computations
    • Errors
    • Convolutions
    • Default dtypes and the X64 flag
    • Type promotion semantics
    • Rank promotion warning
  • JAX 201: performance and scaling
    • Just-in-time compilation
    • Ahead-of-time lowering and compilation
    • Control flow and logical operators with jit
    • Data placement
    • Distributed arrays and automatic parallelization
    • Manual parallelism with shard_map
    • External callbacks
    • Benchmarking and profiling
    • Debugging runtime values
    • Debugging slow JAX tracing and XLA compilation
    • Matmul precision
    • Controlling XLA from JAX
    • GPU memory allocation
    • Memory spaces and host offloading
  • JAX 301: advanced autodiff and extending JAX
    • The Autodiff Cookbook with JVP and VJP
    • First-class VJPs
    • Autodiff and sharding
    • Custom derivative rules with hijax primitives
    • Custom JVPs or VJPs with custom_jvp and custom_vjp
    • Autodiff with refs
    • Gradient checkpointing with jax.checkpoint (jax.remat)
    • Defining new JAX types with hijax
  • JAX 401: kernels and FFI
    • Pallas: custom kernels in JAX
    • Writing High-Performance GPU Kernels with CuTe DSL and JAX
    • Foreign function interface (FFI)
  • JAX 501: systems topics
    • Introduction to multi-controller JAX (aka multi-process/multi-host JAX)
    • Distributed data loading
    • Fault Tolerant Distributed JAX
    • Security considerations
    • Exporting and serializing staged-out computations
    • Shape polymorphism
    • Persistent compilation cache
    • Transfer guard
  • JAX 601: internals
    • The jaxpr language
    • Primitives
    • Autodidax: JAX core from scratch
    • Autodidax2, part 1: JAX from scratch, again
  • 🔪 JAX - The Sharp Bits 🔪
  • Pallas: a JAX kernel language
    • Pallas TPU
      • Quickstart: TPU
      • Writing TPU kernels with Pallas
      • TPU Pipelining
      • Matrix Multiplication
      • Scalar Prefetch and Block-Sparse Computation
      • Distributed Computing in Pallas for TPUs
      • Pallas Core-specific Programming
      • SparseCore Kernel Writing
      • Pseudo-Random Number Generation
      • TPU Hardware Reference
    • Pallas Mosaic GPU
      • Quickstart: GPU
      • Mosaic GPU Pipelining
      • Writing high-performance matrix multiplication kernels for Blackwell
      • Collective matrix multiplication
      • Writing Mosaic GPU kernels with Pallas
    • Pallas Quickstart
    • Grids and BlockSpecs
    • Software Pipelining
    • API reference
      • Pallas TPU (TensorCore)
      • Pallas MGPU
      • Triton
    • Pallas Design Notes
      • Pallas Design
      • Pallas Async Operations
    • Pallas Changelog

Resources

  • Additional guides
    • Colocated Python
    • The Training Cookbook
    • GPU performance tips
  • API Reference
    • jax.numpy module
      • jax.numpy.fft.fft
      • jax.numpy.fft.fft2
      • jax.numpy.fft.fftfreq
      • jax.numpy.fft.fftn
      • jax.numpy.fft.fftshift
      • jax.numpy.fft.hfft
      • jax.numpy.fft.ifft
      • jax.numpy.fft.ifft2
      • jax.numpy.fft.ifftn
      • jax.numpy.fft.ifftshift
      • jax.numpy.fft.ihfft
      • jax.numpy.fft.irfft
      • jax.numpy.fft.irfft2
      • jax.numpy.fft.irfftn
      • jax.numpy.fft.rfft
      • jax.numpy.fft.rfft2
      • jax.numpy.fft.rfftfreq
      • jax.numpy.fft.rfftn
    • jax.scipy module
      • jax.scipy.stats.bernoulli.logpmf
      • jax.scipy.stats.bernoulli.pmf
      • jax.scipy.stats.bernoulli.cdf
      • jax.scipy.stats.bernoulli.ppf
    • jax.lax module
    • jax.random module
    • jax.sharding module
    • jax.ad_checkpoint module
    • jax.debug module
    • jax.dlpack module
    • jax.distributed module
    • jax.dtypes module
    • jax.ffi module
    • jax.flatten_util module
    • jax.image module
    • jax.nn module
      • jax.nn.initializers module
    • jax.ops module
    • jax.profiler module
    • jax.ref module
    • jax.stages module
    • jax.test_util module
    • jax.tree module
    • jax.tree_util module
    • jax.typing module
    • jax.export module
    • jax.extend module
      • jax.extend.backend module
      • jax.extend.core module
      • jax.extend.linear_util module
      • jax.extend.lowering module
      • jax.extend.mlir module
      • jax.extend.pallas module
      • jax.extend.random module
      • jax.extend.xla module
    • jax.example_libraries module
      • jax.example_libraries.optimizers module
      • jax.example_libraries.stax module
    • jax.experimental module
      • jax.experimental.checkify module
      • jax.experimental.compilation_cache module
      • jax.experimental.custom_partitioning module
      • jax.experimental.jet module
      • jax.experimental.key_reuse module
      • jax.experimental.mesh_utils module
      • jax.experimental.multihost_utils module
      • jax.experimental.pallas module
        • Pallas TPU (TensorCore)
        • Pallas MGPU
        • Triton
      • jax.experimental.random module
      • jax.experimental.serialize_executable module
      • jax.experimental.sparse module
        • jax.experimental.sparse.BCOO
        • jax.experimental.sparse.bcoo_broadcast_in_dim
        • jax.experimental.sparse.bcoo_concatenate
        • jax.experimental.sparse.bcoo_dot_general
        • jax.experimental.sparse.bcoo_dot_general_sampled
        • jax.experimental.sparse.bcoo_dynamic_slice
        • jax.experimental.sparse.bcoo_extract
        • jax.experimental.sparse.bcoo_fromdense
        • jax.experimental.sparse.bcoo_gather
        • jax.experimental.sparse.bcoo_multiply_dense
        • jax.experimental.sparse.bcoo_multiply_sparse
        • jax.experimental.sparse.bcoo_update_layout
        • jax.experimental.sparse.bcoo_reduce_sum
        • jax.experimental.sparse.bcoo_reshape
        • jax.experimental.sparse.bcoo_slice
        • jax.experimental.sparse.bcoo_sort_indices
        • jax.experimental.sparse.bcoo_squeeze
        • jax.experimental.sparse.bcoo_sum_duplicates
        • jax.experimental.sparse.bcoo_todense
        • jax.experimental.sparse.bcoo_transpose
      • jax.experimental.xla_metadata module
    • jax.Array.addressable_shards
    • jax.Array.all
    • jax.Array.any
    • jax.Array.argmax
    • jax.Array.argmin
    • jax.Array.argpartition
    • jax.Array.argsort
    • jax.Array.astype
    • jax.Array.at
    • jax.Array.byteswap
    • jax.Array.choose
    • jax.Array.clip
    • jax.Array.compress
    • jax.Array.committed
    • jax.Array.conj
    • jax.Array.conjugate
    • jax.Array.copy
    • jax.Array.copy_to_host_async
    • jax.Array.cumprod
    • jax.Array.cumsum
    • jax.Array.device
    • jax.Array.diagonal
    • jax.Array.dot
    • jax.Array.dtype
    • jax.Array.flat
    • jax.Array.flatten
    • jax.Array.global_shards
    • jax.Array.imag
    • jax.Array.is_fully_addressable
    • jax.Array.is_fully_replicated
    • jax.Array.item
    • jax.Array.itemsize
    • jax.Array.max
    • jax.Array.mean
    • jax.Array.min
    • jax.Array.nbytes
    • jax.Array.ndim
    • jax.Array.nonzero
    • jax.Array.prod
    • jax.Array.ptp
    • jax.Array.ravel
    • jax.Array.real
    • jax.Array.repeat
    • jax.Array.reshape
    • jax.Array.round
    • jax.Array.searchsorted
    • jax.Array.shape
    • jax.Array.sharding
    • jax.Array.size
    • jax.Array.sort
    • jax.Array.squeeze
    • jax.Array.std
    • jax.Array.sum
    • jax.Array.swapaxes
    • jax.Array.take
    • jax.Array.to_device
    • jax.Array.trace
    • jax.Array.transpose
    • jax.Array.var
    • jax.Array.view
    • jax.Array.T
    • jax.Array.mT
  • API compatibility
  • Python and NumPy version support policy
  • Developer notes
    • Contributing to JAX
    • Building from source
    • Investigating a regression
    • JAX Enhancement Proposals (JEPs)
      • 263: JAX PRNG Design
      • 2026: Custom JVP/VJP rules for JAX-transformable functions
      • 4008: Custom VJP and `nondiff_argnums` update
      • 4410: Omnistaging
      • 9263: Typed keys & pluggable RNGs
      • 9407: Design of Type Promotion Semantics for JAX
      • 9419: Jax and Jaxlib versioning
      • 10657: Sequencing side-effects in JAX
      • 11830: `jax.remat` / `jax.checkpoint` new implementation
      • 12049: Type Annotation Roadmap for JAX
      • 14273: `shard_map` (`shmap`) for simple per-device code
      • 15856: `jax.extend`, an extensions module
      • 17111: Efficient transposition of `shard_map` (and other maps)
      • 18137: Scope of JAX NumPy & SciPy Wrappers
      • 25516: Effort-based versioning
      • 28661: Supporting the `__jax_array__` protocol
      • 28845: Stateful Randomness in JAX
    • JAX Internal Implementation Notes
      • Handling of closed-over constants
  • Extension guides
    • Writing custom Jaxpr interpreters in JAX
    • jax.extend module
      • jax.extend.backend module
      • jax.extend.core module
      • jax.extend.linear_util module
      • jax.extend.lowering module
      • jax.extend.mlir module
      • jax.extend.pallas module
      • jax.extend.random module
      • jax.extend.xla module
    • Building on JAX
  • About the project
  • Frequently asked questions (FAQ)
  • Change log
  • Glossary of terms
  • Configuration Options
  • API Reference
  • jax.experimental module
  • .rst

jax.experimental module

Contents

  • Experimental Modules

jax.experimental module#

jax.experimental.optix has been moved into its own Python package (deepmind/optax).

jax.experimental.ann has been moved into jax.lax.

Experimental Modules#

  • jax.experimental.checkify module
  • jax.experimental.compilation_cache module
  • jax.experimental.custom_partitioning module
  • jax.experimental.jet module
  • jax.experimental.key_reuse module
  • jax.experimental.mesh_utils module
  • jax.experimental.multihost_utils module
  • jax.experimental.pallas module
  • jax.experimental.random module
  • jax.experimental.serialize_executable module
  • jax.experimental.sparse module
  • jax.experimental.xla_metadata module

previous

jax.example_libraries.stax module

next

jax.experimental.checkify module

Contents
  • Experimental Modules

By The JAX authors

© Copyright 2024, The JAX Authors.