Skip to main content

jax-progress

Progress meters for JAX loops, scans, and Diffrax solves.

Features

  • Tqdm progress bars for JAX loops (scan, while_loop).
  • Support for vmap with correct progress tracking (skips batched updates, tracks n slowest processes).
  • Support for shard_map with device-level progress tracking.
  • diffrax compatible 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 vmap and shard_map for multi-level parallelism. See examples/ 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)

Source distribution for jax-progress 0.1.0
File Size Uploaded
jax_progress-0.1.0.tar.gz 12.1 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for jax-progress 0.1.0
File Interpreter ABI Platform
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 log

Release 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

Release history Release notifications | RSS feed

This release

0.1.0 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