MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / DWT1DForward

Class DWT1DForward

layers/DWT_Decomposition.py:128–182  ·  view source on GitHub ↗

Performs a 1d DWT Forward decomposition of an image Args: J (int): Number of levels of decomposition wave (str or pywt.Wavelet or tuple(ndarray)): Which wavelet to use. Can be: 1) a string to pass to pywt.Wavelet constructor 2) a pywt.

Source from the content-addressed store, hash-verified

126###############################################################################################
127
128class DWT1DForward(nn.Module):
129 """ Performs a 1d DWT Forward decomposition of an image
130
131 Args:
132 J (int): Number of levels of decomposition
133 wave (str or pywt.Wavelet or tuple(ndarray)): Which wavelet to use.
134 Can be:
135 1) a string to pass to pywt.Wavelet constructor
136 2) a pywt.Wavelet class
137 3) a tuple of numpy arrays (h0, h1)
138 mode (str): 'zero', 'symmetric', 'reflect' or 'periodization'. The
139 padding scheme
140 """
141
142 def __init__(self, J=1, wave='db1', mode='zero', use_amp=False):
143 super().__init__()
144 self.use_amp = use_amp
145 if isinstance(wave, str):
146 wave = pywt.Wavelet(wave)
147 if isinstance(wave, pywt.Wavelet):
148 h0, h1 = wave.dec_lo, wave.dec_hi
149 else:
150 assert len(wave) == 2
151 h0, h1 = wave[0], wave[1]
152
153 # Prepare the filters - this makes them into column filters
154 filts = prep_filt_afb1d(h0, h1)
155 self.register_buffer('h0', filts[0])
156 self.register_buffer('h1', filts[1])
157 self.J = J
158 self.mode = mode
159
160 def forward(self, x):
161 """ Forward pass of the DWT.
162
163 Args:
164 x (tensor): Input of shape :math:`(N, C_{in}, L_{in})`
165
166 Returns:
167 (yl, yh)
168 tuple of lowpass (yl) and bandpass (yh) coefficients.
169 yh is a list of length J with the first entry
170 being the finest scale coefficients.
171 """
172 assert x.ndim == 3, "Can only handle 3d inputs (N, C, L)"
173 highs = []
174 x0 = x
175 mode = mode_to_int(self.mode)
176
177 # Do a multilevel transform
178 for j in range(self.J):
179 x0, x1 = AFB1D.apply(x0, self.h0, self.h1, mode, self.use_amp)
180 highs.append(x1)
181
182 return x0, highs
183
184
185class DWT1DInverse(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected