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.
| 126 | ############################################################################################### |
| 127 | |
| 128 | class 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 | |
| 185 | class DWT1DInverse(nn.Module): |