Skip to main content
GitHub Actions Documentation Status PyPI Versions PyPI - Project PyPI - License Black code style

Differentiable and GPU-enabled fast wavelet transforms in JAX.

Features

  • 1d analysis and synthesis transforms are implemented in src/jaxwt/conv_fwt.py.

  • 2d analysis and synthesis transform are part of the src/jaxwt/conv_fwt_2d.py module.

  • cwt-function supports 1d continuous wavelet transforms.

  • The WaveletPacket object supports 1d wavelet packet transforms.

  • WaveletPacket2d implements two-dimensional wavelet packet transforms.

This toolbox extends PyWavelets . We additionally provide GPU and gradient support via a PyTorch backend.

Installation

To install Jax, head over to https://github.com/google/jax#installation and follow the procedure described there. Afterward, type pip install jaxwt to install the Jax-Wavelet-Toolbox. You can uninstall it later by typing pip uninstall jaxwt.

Documentation

The documentation is available at: https://jax-wavelet-toolbox.readthedocs.io .

Transform Examples:

One-dimensional fast wavelet transform:

import pywt
import numpy as np;
import jax.numpy as jnp
import jaxwt as jwt
# generate an input of even length.
data = jnp.array([0., 1, 2, 3, 4, 5, 6, 7, 7, 6, 5, 4, 3, 2, 1, 0])
wavelet = pywt.Wavelet('haar')

# compare the forward fwt coefficients
print(pywt.wavedec(np.array(data), wavelet, mode='zero', level=2))
print(jwt.wavedec(data, wavelet, mode='zero', level=2))

# invert the fwt.
print(jwt.waverec(jwt.wavedec(data, wavelet, mode='zero', level=2),
                  wavelet))

Two-dimensional fast wavelet transform:

import pywt, scipy.datasets
import jaxwt as jwt
import jax.numpy as jnp
face = jnp.transpose(
    scipy.datasets.face(), [2, 0, 1]).astype(jnp.float64)
transformed = jwt.wavedec2(face, pywt.Wavelet("haar"),
                           level=2, mode="reflect")
reconstruction = jwt.waverec2(transformed, pywt.Wavelet("haar"))
jnp.max(jnp.abs(face - reconstruction))

Testing

Unit tests are handled by nox. Clone the repository and run it with the following:

$ pip install nox
$ git clone https://github.com/v0lta/Jax-Wavelet-Toolbox
$ cd Jax-Wavelet-Toolbox
$ nox -s test

Goals

  • In the spirit of Jax, the aim is to be 100% pywt compatible. Whenever possible, interfaces should be the same results identical.

64-Bit floating-point numbers

If you need 64-bit floating point support, set the Jax config flag:

from jax.config import config
config.update("jax_enable_x64", True)

Citation

If you use this work in a scientific context, please cite:

@phdthesis{handle:20.500.11811/9245,
  urn: https://nbn-resolving.org/urn:nbn:de:hbz:5-63361,
  author = {{Moritz Wolter}},
  title = {Frequency Domain Methods in Recurrent Neural Networks for Sequential Data Processing},
  school = {Rheinische Friedrich-Wilhelms-Universität Bonn},
  year = 2021,
  month = jul,
  url = {https://hdl.handle.net/20.500.11811/9245}
}

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

jaxwt-0.0.8.tar.gz (21.2 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

jaxwt-0.0.8-py3-none-any.whl (20.2 kB view details)

Uploaded Python 3

File details

Details for the file jaxwt-0.0.8.tar.gz.

File metadata

  • Download URL: jaxwt-0.0.8.tar.gz
  • Upload date:
  • Size: 21.2 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/4.0.2 CPython/3.9.12

File hashes

Hashes for jaxwt-0.0.8.tar.gz
Algorithm Hash digest
SHA256 10e25d8317611b294d582d13f6137500ce7cd282b0c8f56403aa359df72e6730
MD5 3c18139af280718d73fae6f2a28f391c
BLAKE2b-256 bf70ad460c09b4f1b28a33e536967f2fb76bd0dd30d2d32d008d3cf112c94c2d

See more details on using hashes here.

File details

Details for the file jaxwt-0.0.8-py3-none-any.whl.

File metadata

  • Download URL: jaxwt-0.0.8-py3-none-any.whl
  • Upload date:
  • Size: 20.2 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/4.0.2 CPython/3.9.12

File hashes

Hashes for jaxwt-0.0.8-py3-none-any.whl
Algorithm Hash digest
SHA256 8d74ab2c3f5d20f9290934d8c837d8f136e090b7c5033783a6dc300f38277d6c
MD5 e20f376b24d57349526f1f075b7ea3c9
BLAKE2b-256 c7c8f80f0d509f87ff96df6c7316168b72b2d6fd0ec595b8f8954d1dea2723a1

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page