MCPcopy Create free account
hub / github.com/BorealisAI/scaleformer / construct

Function construct

robust_loss_pytorch/wavelet.py:294–339  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

292
293
294def 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
342def collapse(pyr, wavelet_type):

Callers

nothing calls this directly

Calls 3

get_max_num_levelsFunction · 0.85
generate_filtersFunction · 0.85
_downsampleFunction · 0.85

Tested by

no test coverage detected