Skip to main content

jaxwavelets

PyPI CI License: MIT Python codecov

Extending PyWavelets to JAX. Differentiable, JIT-compilable, GPU-ready wavelet transforms.

Built on the mathematical foundations of PyWavelets and validated against it to machine precision. jaxwavelets brings the full PyWavelets API to JAX, enabling automatic differentiation, GPU acceleration, and composability with jax.vmap, jax.jit, and jax.pmap.

Features

Transform Functions
Discrete wavelet dwt, idwt, dwt2, idwt2, dwtn, idwtn
Multilevel wavedec2, waverec2, wavedecn, waverecn
Stationary (undecimated) swt, iswt, swt2, iswt2, swtn, iswtn
Continuous cwt, prepare_cwt, apply_cwt
Fully separable fswavedecn, fswaverecn
Multiresolution analysis mra, imra, mra2, imra2, mran, imran
Wavelet packets wp_decompose, wp_reconstruct, wp_decompose_nd, wp_reconstruct_nd
Thresholding soft_threshold, hard_threshold, garrote_threshold, firm_threshold
Utilities downcoef, upcoef, qmf, orthogonal_filter_bank

Wavelets: haar, db1-20, sym2-20, coif1-5, plus continuous wavelets (Morlet, Mexican hat, Gaussian 1-8, complex Gaussian 1-8, complex Morlet, Shannon, frequency B-spline).

Usage

import jax
import jax.numpy as jnp
import jaxwavelets as wt

# Decompose and reconstruct
x = jnp.ones((64, 64))
coeffs = wt.wavedecn(x, 'db4', level=3)
rec = wt.waverecn(coeffs, 'db4')

# Batch via vmap
from functools import partial
batch = jnp.ones((10, 64, 64))
batch_coeffs = jax.vmap(partial(wt.wavedecn, wavelet='db4', level=3))(batch)

# Differentiate through the transform
grad = jax.grad(lambda x: jnp.sum(wt.waverecn(wt.wavedecn(x, 'db4'), 'db4')))(x)

# JIT-compile for speed
fast = jax.jit(wt.wavedecn, static_argnames=['wavelet', 'mode', 'level'])
coeffs = fast(x, wavelet='db4', level=3)

Performance

JIT-compiled jaxwavelets on CPU vs PyWavelets C:

Transform                       pywt         jaxwavelets (JIT)    ratio
--------------------------------------------------------------------------
dwt 1D (N=4096)                  0.011ms       0.023ms       2.1x
wavedecn 1D (N=4096)             0.065ms       0.046ms       0.7x  ← faster
dwt2 (256x256)                   0.608ms       0.287ms       0.5x  ← faster
wavedecn 2D level=3              0.755ms       0.363ms       0.5x  ← faster
swt 1D level=3 (N=1024)          0.023ms       0.025ms       1.1x
cwt morl 6 scales (N=512)        0.316ms       0.139ms       0.4x  ← faster
cwt cmor 6 scales (N=512)        0.615ms       0.254ms       0.4x  ← faster

On top of this, jaxwavelets supports jax.grad, jax.vmap, jax.pmap, and GPU acceleration.

Installation

pip install jaxwavelets

No runtime dependency on PyWavelets. Filter coefficients are pre-extracted.

Testing

pip install pywt pytest
pytest jaxwavelets/tests/

1189 tests verify numerical agreement with PyWavelets to machine precision.

Composability

Every function operates on a single example. Batching, differentiation, compilation, and distribution compose naturally via JAX transforms:

import jaxwavelets as wt

# Batch over examples
jax.vmap(partial(wt.wavedecn, wavelet='db4'))(batch_of_fields)

# Per-example gradients
jax.vmap(jax.grad(loss_fn))(batch)

# Distribute across devices
jax.pmap(partial(wt.wavedecn, wavelet='db4'))(sharded_data)

# Nest arbitrarily
jax.jit(jax.vmap(jax.grad(
    lambda x: jnp.sum(wt.waverecn(wt.wavedecn(x, 'db4'), 'db4'))
)))(batch)

Coefficients are JAX pytrees, so jax.tree_util.tree_map works directly on them.

Design

  • Pure JAX — no numpy, no C extensions
  • Single-example functions — compose with jax.vmap/jax.pmap/jax.grad/jax.jit
  • Pytree coefficients — all outputs are JAX-compatible pytrees
  • Validated against PyWavelets — machine-precision numerical agreement

Acknowledgements

jaxwavelets extends the PyWavelets library to JAX. PyWavelets provides the mathematical reference implementation and filter coefficient database used for validation.

Release files for jaxwavelets 0.1.13

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

Source distribution (sdist)

Source distribution for jaxwavelets 0.1.13
File Size Uploaded
jaxwavelets-0.1.13.tar.gz 42.9 kB Details

Built distribution (wheel)

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

Total release size:91.1 kB

Release files / jaxwavelets-0.1.13.tar.gz

Download URL jaxwavelets-0.1.13.tar.gz
Size 42.9 kB
Tags Source
SHA-256 checksum
How to use checksums
8c84415da5bd6dbb580d6b9a54794d636f2725ca7f4d3417159f84fa385dc193
BLAKE2b-256 checksum
How to use checksums
4269a4e504670aca2f36caa57abc10bc81ea9e318e9704647222188abc1a4ac9
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.1.0 CPython/3.13.12

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 Apr 16, 2026.

Transparency log

Release files / jaxwavelets-0.1.13-py3-none-any.whl

Download URL jaxwavelets-0.1.13-py3-none-any.whl
Size 48.3 kB
Tags Python 3
SHA-256 checksum
How to use checksums
2d53d9208c723821ee4b4ce2f817152bc7722dbf0cd7a3e543755203b2fc761f
BLAKE2b-256 checksum
How to use checksums
ec5606e5ee6856ccf5a3f2dbc8daf8f516aee83257cda8d3356099003a84028b
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.1.0 CPython/3.13.12

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 Apr 16, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.1.13 This release

2 release files

0.1.10

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