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

Class DWT1DInverse

layers/DWT_Decomposition.py:185–240  ·  view source on GitHub ↗

Performs a 1d DWT Inverse reconstruction of an image Args: 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.Wavelet class 3) a tuple of numpy arra

Source from the content-addressed store, hash-verified

183
184
185class DWT1DInverse(nn.Module):
186 """ Performs a 1d DWT Inverse reconstruction of an image
187
188 Args:
189 wave (str or pywt.Wavelet or tuple(ndarray)): Which wavelet to use.
190 Can be:
191 1) a string to pass to pywt.Wavelet constructor
192 2) a pywt.Wavelet class
193 3) a tuple of numpy arrays (h0, h1)
194 mode (str): 'zero', 'symmetric', 'reflect' or 'periodization'. The
195 padding scheme
196 """
197
198 def __init__(self, wave='db1', mode='zero', use_amp=False):
199 super().__init__()
200 self.use_amp = use_amp
201 if isinstance(wave, str):
202 wave = pywt.Wavelet(wave)
203 if isinstance(wave, pywt.Wavelet):
204 g0, g1 = wave.rec_lo, wave.rec_hi
205 else:
206 assert len(wave) == 2
207 g0, g1 = wave[0], wave[1]
208
209 # Prepare the filters
210 filts = prep_filt_sfb1d(g0, g1)
211 self.register_buffer('g0', filts[0])
212 self.register_buffer('g1', filts[1])
213 self.mode = mode
214
215 def forward(self, coeffs):
216 """
217 Args:
218 coeffs (yl, yh): tuple of lowpass and bandpass coefficients, should
219 match the format returned by DWT1DForward.
220
221 Returns:
222 Reconstructed input of shape :math:`(N, C_{in}, L_{in})`
223
224 Note:
225 Can have None for any of the highpass scales and will treat the
226 values as zeros (not in an efficient way though).
227 """
228 x0, highs = coeffs
229 assert x0.ndim == 3, "Can only handle 3d inputs (N, C, L)"
230 mode = mode_to_int(self.mode)
231 # Do a multilevel inverse transform
232 for x1 in highs[::-1]:
233 if x1 is None:
234 x1 = torch.zeros_like(x0)
235
236 # 'Unpad' added signal
237 if x0.shape[-1] > x1.shape[-1]:
238 x0 = x0[..., :-1]
239 x0 = SFB1D.apply(x0, x1, self.g0, self.g1, mode, self.use_amp)
240 return x0
241
242

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected