num_levels a wavelet decomposition of an image. Args: im: A numpy or TF tensor of single or double precision floats of size (batch_size, width, height) num_levels: The number of levels (or scales) of the wavelet decomposition to apply. A value of 0 returns a "wavelet decomposi
(im, num_levels, wavelet_type)
| 292 | |
| 293 | |
| 294 | def construct(im, num_levels, wavelet_type): |
| 295 | """num_levels a wavelet decomposition of an image. |
| 296 | |
| 297 | Args: |
| 298 | im: A numpy or TF tensor of single or double precision floats of size |
| 299 | (batch_size, width, height) |
| 300 | num_levels: The number of levels (or scales) of the wavelet decomposition to |
| 301 | apply. A value of 0 returns a "wavelet decomposition" that is just the |
| 302 | image. |
| 303 | wavelet_type: The kind of wavelet to use, see generate_filters(). |
| 304 | |
| 305 | Returns: |
| 306 | A wavelet decomposition of `im` that has `num_levels` levels (not including |
| 307 | the coarsest residual level) and is of type `wavelet_type`. This |
| 308 | decomposition is represented as a tuple of 3-tuples, with the final element |
| 309 | being a tensor: |
| 310 | ((band00, band01, band02), (band10, band11, band12), ..., resid) |
| 311 | Where band** and resid are TF tensors. Each element of these nested tuples |
| 312 | is of shape [batch_size, width * 2^-(level+1), height * 2^-(level+1)], |
| 313 | though the spatial dimensions may be off by 1 if width and height are not |
| 314 | factors of 2. The residual image is of the same (rough) size as the last set |
| 315 | of bands. The floating point precision of these tensors matches that of |
| 316 | `im`. |
| 317 | """ |
| 318 | if len(im.shape) != 3: |
| 319 | raise ValueError( |
| 320 | 'Expected `im` to have a rank of 3, but is of size {}'.format(im.shape)) |
| 321 | im = torch.as_tensor(im) |
| 322 | if num_levels == 0: |
| 323 | return (im,) |
| 324 | max_num_levels = get_max_num_levels(im.shape) |
| 325 | assert max_num_levels >= num_levels |
| 326 | filters = generate_filters(wavelet_type) |
| 327 | pyr = [] |
| 328 | for _ in range(num_levels): |
| 329 | # import pdb |
| 330 | # pdb.set_trace() |
| 331 | hi = _downsample(im, filters.analysis_hi, 0, 1) |
| 332 | lo = _downsample(im, filters.analysis_lo, 0, 0) |
| 333 | pyr.append((_downsample(hi, filters.analysis_hi, 1, |
| 334 | 1), _downsample(lo, filters.analysis_hi, 1, 1), |
| 335 | _downsample(hi, filters.analysis_lo, 1, 0))) |
| 336 | im = _downsample(lo, filters.analysis_lo, 1, 0) |
| 337 | pyr.append(im) |
| 338 | pyr = tuple(pyr) |
| 339 | return pyr |
| 340 | |
| 341 | |
| 342 | def collapse(pyr, wavelet_type): |
nothing calls this directly
no test coverage detected