Skip to main content
Parajax

Automatic vectorization and parallelization of JAX-based functions

Documentation CI Codecov Ruff ty uv Publish PyPI PyPI - Python Version

Features

Parajax provides simple decorators for automatically mapping JAX functions over array batches and across multiple devices.

  • 🪄 Automatic vectorization — map functions over arbitrary broadcastable batch dimensions
  • 🚀 Automatic parallelization — distribute batched computations across CPUs, GPUs, or TPUs
  • 📐 NumPy-style broadcasting — infer vectorization structure directly from input shapes
  • 📦 Memory-aware batching — optionally evaluate vectorized functions in smaller batches
  • 🧩 Composable with JAX — use naturally with jax.jit and other JAX transformations
  • 🎯 Simple decorator-based API — add vectorization or parallelism without restructuring your function

Installation

pip install parajax

Vectorization

@vectorize turns a function that operates on individual values, vectors, matrices, or other array objects into one that automatically operates over arbitrary broadcastable leading dimensions.

For a scalar function, no configuration is needed:

import jax
import jax.numpy as jnp

from parajax import vectorize


@vectorize
def soft_threshold(x, threshold):
    return jax.lax.cond(
        jnp.abs(x) > threshold,
        lambda: jnp.sign(x) * (jnp.abs(x) - threshold),
        lambda: 0.0,
    )


x = jnp.arange(12).reshape(3, 4)
y = soft_threshold(x, threshold=2.0)

assert y.shape == (3, 4)

Note that the soft_threshold function will also continue to work on scalar inputs, as well as with inputs of any other shape that are compatible (broadcastable) with each other.

For functions that are defined to operate on arrays, ndim specifies how many dimensions belong to each individual input. Any additional leading dimensions are treated as batch dimensions.

For example, a matrix-vector product operates on a rank-2 matrix and a rank-1 vector:

@vectorize(ndim=(2, 1))
def matvec(A, x):
    return A @ x

The unvectorized function therefore expects:

A: (m, n)
x: (n,)

but the decorated function also accepts arbitrary broadcastable batch dimensions:

A = jnp.ones((100, 1, 3, 4))
x = jnp.ones((50, 4))

y = matvec(A, x)

assert y.shape == (100, 50, 3)

Here the batch shapes (100, 1) and (50,) are broadcast to (100, 50). At each point in that batch, the original function receives a matrix with shape (3, 4) and a vector with shape (4,).

Specifying ndim

A single integer applies the same core dimensionality to every argument:

@vectorize(ndim=1)
def dot(x, y):
    return x @ y

A sequence specifies the core dimensionality of each parameter:

@vectorize(ndim=(2, 1))
def matvec(A, x):
    return A @ x

A mapping is useful when only some parameters should be vectorized:

@vectorize(ndim={"A": 2, "x": 1})
def matvec(A, x, *, scale=1.0):
    return scale * (A @ x)

Parameters omitted from a mapping are passed unchanged to the underlying function.

None can also be used explicitly to mark a parameter as static:

@vectorize(ndim=(1, 1, None))
def distance(x, y, metric): ...

Conceptually, ndim separates each array shape into:

batch dimensions + core dimensions

For example, with ndim=2:

(..., m, n)
 ^^^  ^^^^
batch core

Parajax automatically broadcasts the batch dimensions and maps the original function over them.

Batched execution

By default, vectorize processes the complete broadcast batch at once, equivalently to using jax.vmap.

For computations where memory use is more important than maximum vectorization, set batch_size:

@vectorize(ndim=1, batch_size=32)
def expensive_function(x): ...

The same vectorized operation is then evaluated in batches of at most 32 elements along each mapped dimension.

Parallelization

@parallelize distributes a batched JAX function across all available devices.

import multiprocessing

import jax
import jax.numpy as jnp

from parajax import parallelize


jax.config.update("jax_num_cpu_devices", multiprocessing.cpu_count())
# Only needed on CPU to make multiple CPU devices available to JAX.


@parallelize
def square(xs):
    return xs**2


xs = jnp.arange(12_345)
ys = square(xs)

Invocations of square are automatically distributed across the available devices. Input sizes do not need to be divisible by the number of devices.

Documentation

See the documentation for the complete API reference and additional examples.

Metadata

Release files for parajax 0.4.0

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

Source distribution (sdist)

Source distribution for parajax 0.4.0
File Size Uploaded
parajax-0.4.0.tar.gz 8.3 kB Details

Built distribution (wheel)

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

Total release size: 18.0 kB

Release files / parajax-0.4.0.tar.gz

Download URL parajax-0.4.0.tar.gz
Size 8.3 kB
Tags Source
SHA-256 checksum
How to use checksums
6d61915e8896dfaa6cc7ced9f717db2c8a90347d173b85ec79d5755c8dc737c6
BLAKE2b-256 checksum
How to use checksums
b2a2bfc68712fddcebe0462d1150c388cd3bcc9da87f2b2a658a616f4b3e6464
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via uv/0.12.11 {"installer":{"name":"uv","version":"0.12.11","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}

Release files / parajax-0.4.0-py3-none-any.whl

Download URL parajax-0.4.0-py3-none-any.whl
Size 9.7 kB
Tags Python 3
SHA-256 checksum
How to use checksums
4eb227e8ef63ee14a9ef7ec1ec2336e9d5c410e128373714cf2a9fac47b9a97c
BLAKE2b-256 checksum
How to use checksums
08665f5b6c56217e450b367a54663343f2287eb40a7e3dad435d4a3f35da1787
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via uv/0.12.11 {"installer":{"name":"uv","version":"0.12.11","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}

Release history Release notifications | RSS feed

0.4.1

2 release files

This release

0.4.0 This release

2 release files

0.3.5

2 release files

0.3.4

2 release files

0.3.3

2 release files

0.3.2

2 release files

0.3.1

2 release files

0.3.0

2 release files

0.2.4

2 release files

0.2.3

2 release files

0.2.2

2 release files

0.2.1

2 release files

0.2.0

2 release files

0.1.3

2 release files

0.1.2

2 release files

0.1.1

2 release files

0.1.0

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