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
| 183 | |
| 184 | |
| 185 | class 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 |