CoolFace
Apppublic

Aluode/PerceptionLabPortable

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
_utils.py105 linesDownload Raw Back to pywt
1# Copyright (c) 2017 The PyWavelets Developers
2#                    <https://github.com/PyWavelets/pywt>
3# See COPYING for license details.
4import inspect
5from collections.abc import Iterable
6
7import numpy as np
8
9from ._extensions._pywt import (
10    ContinuousWavelet,
11    DiscreteContinuousWavelet,
12    Modes,
13    Wavelet,
14)
15
16AxisError: type[Exception]
17if np.lib.NumpyVersion(np.__version__) >= '1.25.0':
18    from numpy.exceptions import AxisError
19else:
20    from numpy import AxisError
21
22
23def _as_wavelet(wavelet):
24    """Convert wavelet name to a Wavelet object."""
25    if not isinstance(wavelet, (ContinuousWavelet, Wavelet)):
26        wavelet = DiscreteContinuousWavelet(wavelet)
27    if isinstance(wavelet, ContinuousWavelet):
28        raise ValueError(
29            "A ContinuousWavelet object was provided, but only discrete "
30            "Wavelet objects are supported by this function.  A list of all "
31            "supported discrete wavelets can be obtained by running:\n"
32            "print(pywt.wavelist(kind='discrete'))")
33    return wavelet
34
35
36def _wavelets_per_axis(wavelet, axes):
37    """Initialize Wavelets for each axis to be transformed.
38
39    Parameters
40    ----------
41    wavelet : Wavelet or tuple of Wavelets
42        If a single Wavelet is provided, it will used for all axes.  Otherwise
43        one Wavelet per axis must be provided.
44    axes : list
45        The tuple of axes to be transformed.
46
47    Returns
48    -------
49    wavelets : list of Wavelet objects
50        A tuple of Wavelets equal in length to ``axes``.
51
52    """
53    axes = tuple(axes)
54    if isinstance(wavelet, (str, Wavelet)):
55        # same wavelet on all axes
56        wavelets = [_as_wavelet(wavelet), ] * len(axes)
57    elif isinstance(wavelet, Iterable):
58        # (potentially) unique wavelet per axis (e.g. for dual-tree DWT)
59        if len(wavelet) == 1:
60            wavelets = [_as_wavelet(wavelet[0]), ] * len(axes)
61        else:
62            if len(wavelet) != len(axes):
63                raise ValueError(
64                    "The number of wavelets must match the number of axes "
65                    "to be transformed.")
66            wavelets = [_as_wavelet(w) for w in wavelet]
67    else:
68        raise ValueError("wavelet must be a str, Wavelet or iterable")
69    return wavelets
70
71
72def _modes_per_axis(modes, axes):
73    """Initialize mode for each axis to be transformed.
74
75    Parameters
76    ----------
77    modes : str or tuple of strings
78        If a single mode is provided, it will used for all axes.  Otherwise
79        one mode per axis must be provided.
80    axes : tuple
81        The tuple of axes to be transformed.
82
83    Returns
84    -------
85    modes : tuple of int
86        A tuple of Modes equal in length to ``axes``.
87
88    """
89    axes = tuple(axes)
90    if isinstance(modes, (int, str)):
91        # same wavelet on all axes
92        modes = [Modes.from_object(modes), ] * len(axes)
93    elif isinstance(modes, Iterable):
94        if len(modes) == 1:
95            modes = [Modes.from_object(modes[0]), ] * len(axes)
96        else:
97            # (potentially) unique wavelet per axis (e.g. for dual-tree DWT)
98            if len(modes) != len(axes):
99                raise ValueError("The number of modes must match the number "
100                                  "of axes to be transformed.")
101        modes = [Modes.from_object(mode) for mode in modes]
102    else:
103        raise ValueError("modes must be a str, Mode enum or iterable")
104    return modes
105