Tinax is a small, typed library of explicit productivity primitives for JAX, Flax NNX, Optax, Orbax, Grain, Chex, and Safetensors workflows.
It provides stable policies for array and RNG ownership, trace-budgeted compilation and batching, hardened autodiff, bounded debug observation, NNX graph copies, explicit stdlib application boundaries, deterministic input pipelines, complete checkpoints, multi-device parallelism, and weight interchange. Tested ecosystem recipes live under examples/ without stable API guarantees.
- Python 3.12, 3.13, or 3.14
pip install tinaxInstall a JAX accelerator distribution when needed:
pip install "tinax[gpu]"
pip install "tinax[tpu]"See the installation guide for platform and accelerator details.
import numpy as np
from tinax.array import from_numpy, inspect_array, to_numpy
host = np.arange(8, dtype=np.float32)
device = from_numpy(host, copy=True)
info = inspect_array(device)
round_trip = to_numpy(device, writable=False)Importing tinax alone is inert. Import the module that owns the behavior you need.
Apache-2.0. See LICENSE.