jax-progress
Progress meters for JAX loops, scans, and Diffrax solves.
Features
- Tqdm progress bars for JAX loops (
scan,while_loop). - Support for
vmapwith correct progress tracking (skips batched updates, tracks n slowest processes). - Support for
shard_mapwith device-level progress tracking. diffraxcompatible progress meter.
Installation
pip install jax-progress
Usage
Basic vmap example
import jax
import jax.numpy as jnp
from jax_progress import TqdmProgressMeter
# Limit to 3 progress bars (shows 3 slowest tasks)
pbar = TqdmProgressMeter(total=100, max_bars=3)
def task(data):
state = pbar.init(vmapped_element=data)
def body(carry, x):
return pbar.step(carry, progress=1), x
state, _ = jax.lax.scan(body, state, data)
pbar.close(state)
return data.sum()
# Run 10 tasks in parallel, but only show 3 slowest
results = jax.vmap(task)(jnp.ones((10, 100)))
shard_map example
from jax.sharding import PartitionSpec as P
from functools import partial
mesh = jax.make_mesh((4,), ('x',))
pbar = TqdmProgressMeter(total=100)
@partial(jax.shard_map, mesh=mesh, in_specs=P('x'), out_specs=P('x'))
def sharded_task(data):
state = pbar.init(spec=P('x'))
def body(carry, x):
return pbar.step(carry, progress=1), x
state, _ = jax.lax.scan(body, state, jnp.arange(100))
pbar.close(state)
return data
results = sharded_task(jnp.ones(4))
Diffrax integration (drop-in replacement)
TqdmProgressMeter can be used as a drop-in replacement for Diffrax's default progress meter:
import diffrax
# Create progress meter with percent_progress=True for Diffrax
pbar = TqdmProgressMeter(total=100, percent_progress=True)
# Use directly in diffeqsolve
sol = diffrax.diffeqsolve(
term, solver, t0=0.0, t1=10.0, dt0=0.01, y0=y0,
stepsize_controller=stepsize_controller,
progress_meter=pbar # Drop-in replacement
)
pbar.terminate()
Note: You can combine
vmapandshard_mapfor multi-level parallelism. Seeexamples/directory for more.
Metadata
Release files for jax-progress 0.1.0
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| jax_progress-0.1.0.tar.gz | 12.1 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| jax_progress-0.1.0-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 21.4 kB
Release files / jax_progress-0.1.0.tar.gz
| Download URL | jax_progress-0.1.0.tar.gz |
|---|---|
| Size | 12.1 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
700ebd006a15cb8e556fac3628a3cfef17071dbafb4c2d20f868d8d4c2b18f40
|
|
BLAKE2b-256 checksum How to use checksums |
1e3c72bc8fc97cdd1b566233ce846ea81b7906d377c468b3837a3c17e77794a3
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/6.1.0 CPython/3.13.7
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Dec 29, 2025.
Transparency logRelease files / jax_progress-0.1.0-py3-none-any.whl
| Download URL | jax_progress-0.1.0-py3-none-any.whl |
|---|---|
| Size | 9.3 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
3fae8ce5045a525ae40288757749e6c9d3830223c9b41d925c50b32c1ce5715e
|
|
BLAKE2b-256 checksum How to use checksums |
3d3ea1d5c39392c0e762da98d9caaf4e5e5eea34e2ecf21f9569165daa54c288
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/6.1.0 CPython/3.13.7
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Dec 29, 2025.
Transparency log