Skip to main content

jaxtyping

IMPORTANT: this relaxes the python constraint to python3.8 to be installable with eztils, but will throw a RuntimeError if you actually try to run with python3.8.

Type annotations and runtime type-checking for:

  1. shape and dtype of JAX arrays; (Now also supports PyTorch, NumPy, and TensorFlow!)
  2. PyTrees.

For example:

from jaxtyping import Array, Float, PyTree

# Accepts floating-point 2D arrays with matching dimensions
def matrix_multiply(x: Float[Array, "dim1 dim2"],
                    y: Float[Array, "dim2 dim3"]
                  ) -> Float[Array, "dim1 dim3"]:
    ...

def accepts_pytree_of_ints(x: PyTree[int]):
    ...

def accepts_pytree_of_arrays(x: PyTree[Float[Array, "batch c1 c2"]]):
    ...

Installation

pip install jaxtyping

Requires Python 3.9+.

JAX is an optional dependency, required for a few JAX-specific types. If JAX is not installed then these will not be available, but you may still use jaxtyping to provide shape/dtype annotations for PyTorch/NumPy/TensorFlow/etc.

The annotations provided by jaxtyping are compatible with runtime type-checking packages, so it is common to also install one of these. The two most popular are typeguard (which exhaustively checks every argument) and beartype (which checks random pieces of arguments).

Documentation

Available at https://docs.kidger.site/jaxtyping.

Finally

See also: other libraries in the JAX ecosystem

Equinox: neural networks.

Optax: first-order gradient (SGD, Adam, ...) optimisers.

Diffrax: numerical differential equation solvers.

Lineax: linear solvers and linear least squares.

Eqxvision: computer vision models.

sympy2jax: SymPy<->JAX conversion; train symbolic expressions via gradient descent.

Levanter: scalable+reliable training of foundation models (e.g. LLMs).

Disclaimer

This is not an official Google product.

Metadata

Release files for ezjaxtyping 0.2.20

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for ezjaxtyping 0.2.20
File Size Uploaded
ezjaxtyping-0.2.20.tar.gz 17.5 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for ezjaxtyping 0.2.20
File Interpreter ABI Platform
ezjaxtyping-0.2.20-py3-none-any.whl Python 3 none any Details

Total release size: 42.6 kB

Release files / ezjaxtyping-0.2.20.tar.gz

Download URL ezjaxtyping-0.2.20.tar.gz
Size 17.5 kB
Tags Source
SHA-256 checksum
How to use checksums
aaa2396bdb515516903fb1aa73a22d53e8d79e5bdd68ac16b47805930692e291
BLAKE2b-256 checksum
How to use checksums
0086fe8f5de46329268db82473a79f7b5ca8027c5eb11e4bbf1048e86dc5a3a3
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via python-httpx/0.24.1

Release files / ezjaxtyping-0.2.20-py3-none-any.whl

Download URL ezjaxtyping-0.2.20-py3-none-any.whl
Size 25.1 kB
Tags Python 3
SHA-256 checksum
How to use checksums
466f483aec7265fe60cf622592861d64fb73b2f0d1bbc7cdbcef63e22e422116
BLAKE2b-256 checksum
How to use checksums
a1b2bb7d6534c83419d42410ace10faacc116646442cb19b4371b53e36f8ca84
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via python-httpx/0.24.1

Release history Release notifications | RSS feed

This release

0.2.20 This release

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page