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
| 905 | |
| 906 | |
| 907 | class 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 | |
| 958 | class SFB1D(Function): |
nothing calls this directly
no outgoing calls
no test coverage detected