Trade off speed (20-25% faster) for more memory (20-25%)
| 112 | return d_u, d_v |
| 113 | |
| 114 | class TriangularToeplitzMultPaddedFast(torch.autograd.Function): |
| 115 | """ Trade off speed (20-25% faster) for more memory (20-25%) """ |
| 116 | |
| 117 | @staticmethod |
| 118 | def forward(ctx, u, v): |
| 119 | n = u.shape[-1] |
| 120 | u_f = torch.fft.rfft(u, n=n, dim=-1) |
| 121 | v_f = torch.fft.rfft(v, n=n, dim=-1) |
| 122 | |
| 123 | ctx.save_for_backward(u_f, v_f) |
| 124 | |
| 125 | uv_f = u_f * v_f |
| 126 | output = torch.fft.irfft(uv_f, n=n, dim=-1) |
| 127 | output[..., n//2:].zero_() |
| 128 | return output |
| 129 | |
| 130 | @staticmethod |
| 131 | def backward(ctx, grad): |
| 132 | u_f, v_f = ctx.saved_tensors |
| 133 | n = grad.shape[-1] |
| 134 | g_expand = F.pad(grad[..., :n//2].flip(-1), (0, n//2)) |
| 135 | g_f = torch.fft.rfft(g_expand, n=n, dim=-1) |
| 136 | gu_f = g_f * u_f |
| 137 | gv_f = g_f * v_f |
| 138 | d_u = torch.fft.irfft(gv_f, n=n, dim=-1) |
| 139 | d_v = torch.fft.irfft(gu_f, n=n, dim=-1) |
| 140 | d_u[..., n//2:].zero_() |
| 141 | d_v[..., n//2:].zero_() |
| 142 | d_u[..., :n//2] = d_u[..., :n//2].flip(-1) # TODO |
| 143 | d_v[..., :n//2] = d_v[..., :n//2].flip(-1) # TODO |
| 144 | return d_u, d_v |
| 145 | |
| 146 | # triangular_toeplitz_multiply = triangular_toeplitz_multiply_ |
| 147 | triangular_toeplitz_multiply = TriangularToeplitzMult.apply |
nothing calls this directly
no outgoing calls
no test coverage detected