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

Class SFB2D

layers/DWT_Decomposition.py:907–955  ·  view source on GitHub ↗

Does a single level 2d wavelet decomposition of an input. Does separate row and column filtering by two calls to :py:func:`pytorch_wavelets.dwt.lowlevel.afb1d` Needs to have the tensors in the right form. Because this function defines its own backward pass, saves on memory by

Source from the content-addressed store, hash-verified

905
906
907class SFB2D(Function):
908 """ Does a single level 2d wavelet decomposition of an input. Does separate
909 row and column filtering by two calls to
910 :py:func:`pytorch_wavelets.dwt.lowlevel.afb1d`
911
912 Needs to have the tensors in the right form. Because this function defines
913 its own backward pass, saves on memory by not having to save the input
914 tensors.
915
916 Inputs:
917 x (torch.Tensor): Input to decompose
918 h0_row: row lowpass
919 h1_row: row highpass
920 h0_col: col lowpass
921 h1_col: col highpass
922 mode (int): use mode_to_int to get the int code here
923
924 We encode the mode as an integer rather than a string as gradcheck causes an
925 error when a string is provided.
926
927 Returns:
928 y: Tensor of shape (N, C*4, H, W)
929 """
930
931 @staticmethod
932 def forward(ctx, low, highs, g0_row, g1_row, g0_col, g1_col, mode):
933 mode = int_to_mode(mode)
934 ctx.mode = mode
935 ctx.save_for_backward(g0_row, g1_row, g0_col, g1_col)
936
937 lh, hl, hh = torch.unbind(highs, dim=2)
938 lo = sfb1d(low, lh, g0_col, g1_col, mode=mode, dim=2)
939 hi = sfb1d(hl, hh, g0_col, g1_col, mode=mode, dim=2)
940 y = sfb1d(lo, hi, g0_row, g1_row, mode=mode, dim=3)
941 return y
942
943 @staticmethod
944 def backward(ctx, dy):
945 dlow, dhigh = None, None
946 if ctx.needs_input_grad[0]:
947 mode = ctx.mode
948 g0_row, g1_row, g0_col, g1_col = ctx.saved_tensors
949 dx = afb1d(dy, g0_row, g1_row, mode=mode, dim=3)
950 dx = afb1d(dx, g0_col, g1_col, mode=mode, dim=2)
951 s = dx.shape
952 dx = dx.reshape(s[0], -1, 4, s[-2], s[-1])
953 dlow = dx[:, :, 0].contiguous()
954 dhigh = dx[:, :, 1:].contiguous()
955 return dlow, dhigh, None, None, None, None, None
956
957
958class SFB1D(Function):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected