Aluode/PerceptionLabPortable
0
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 