Tesseract-JAX is a lightweight extension to Tesseract Core that makes Tesseracts look and feel like regular JAX primitives, and makes them jittable, differentiable, and composable.
Read the docs | Explore the examples | Report an issue | Talk to the community | Contribute
The API of Tesseract-JAX consists of a single function, apply_tesseract(tesseract_client, inputs), which is fully traceable by JAX. This enables end-to-end autodifferentiation and JIT compilation of Tesseract-based pipelines:
@jax.jit
def vector_sum(x, y):
res = apply_tesseract(vectoradd_tesseract, {"a": {"v": x}, "b": {"v": y}})
return res["vector_add"]["result"].sum()
jax.grad(vector_sum)(x, y) # 🎉Note
Before proceeding, make sure you have a working installation of Docker and a modern Python installation (Python 3.10+).
Important
For more detailed installation instructions, please refer to the Tesseract Core documentation.
-
Install Tesseract-JAX:
$ pip install tesseract-jax
-
Build an example Tesseract:
$ git clone https://github.com/pasteurlabs/tesseract-jax $ tesseract build tesseract-jax/examples/simple/vectoradd_jax
-
Use it as part of a JAX program via the JAX-native
apply_tesseractfunction:import jax import jax.numpy as jnp from tesseract_core import Tesseract from tesseract_jax import apply_tesseract # Load the Tesseract t = Tesseract.from_image("vectoradd_jax") t.serve() # Run it with JAX x = jnp.ones((1000,)) y = jnp.ones((1000,)) def vector_sum(x, y): res = apply_tesseract(t, {"a": {"v": x}, "b": {"v": y}}, vmap_method="sequential") return res["vector_add"]["result"].sum() vector_sum(x, y) # success! # You can also use it with JAX transformations like JIT and grad vector_sum_jit = jax.jit(vector_sum) vector_sum_jit(x, y) vector_sum_grad = jax.grad(vector_sum) vector_sum_grad(x, y) # vmap requires an explicit vmap_method — "sequential" is safe but slow # while "auto_experimental" or "expand_dims" is more efficient for Tesseracts that support batching. # See https://docs.pasteurlabs.ai/projects/tesseract-jax/latest/content/vmap-methods.html vector_sum_vmap = jax.vmap(vector_sum) vector_sum_vmap(x.reshape(10, 100), y.reshape(10, 100))
Tip
Now you're ready to jump into our examples for more ways to use Tesseract-JAX.
- Additional required endpoints: Tesseract-JAX requires the
abstract_evalTesseract endpoint to be defined to enable JAX tracing and FFI dispatch. To run a Tesseract that has noabstract_evalendpoint, call it directly through the Tesseract client instead. Additionally, many gradient transformations likejax.gradrequirevector_jacobian_productto be defined.
Tip
When creating a new Tesseract based on a JAX function, use tesseract init --recipe jax to define all required endpoints automatically, including abstract_eval and vector_jacobian_product.
- Non-array outputs come from
abstract_eval: anOutputSchemafield that is not an array, such as astror abool, cannot enter a traced computation. As such, non-array outputs returned byapplyare ignored and their value is instead taken fromabstract_eval.apply_tesseractwarns when the two differ; passcheck_static_outputs=Falseor setTESSERACT_JAX_CHECK_STATIC_OUTPUTS=0to skip the comparison. A field whose value depends on the input values belongs in the schema as an array.
Tesseract-JAX is licensed under the Apache License 2.0 and is free to use, modify, and distribute (under the terms of the license).
Tesseract is a registered trademark of Pasteur Labs, Inc. and may not be used without permission.