MCPcopy Create free account
hub / github.com/HazyResearch/spacetime / backward

Method backward

model/functional/toeplitz.py:87–98  ·  view source on GitHub ↗
(ctx, grad)

Source from the content-addressed store, hash-verified

85
86 @staticmethod
87 def backward(ctx, grad):
88 u_f, v_f = ctx.saved_tensors
89 n = grad.shape[-1]
90 g_expand = F.pad(grad.flip(-1), (0, n))
91 g_f = torch.fft.rfft(g_expand, n=2*n, dim=-1)
92 gu_f = g_f * u_f
93 gv_f = g_f * v_f
94 d_u = torch.fft.irfft(gv_f, n=2*n, dim=-1)[..., :n]
95 d_v = torch.fft.irfft(gu_f, n=2*n, dim=-1)[..., :n]
96 d_u = d_u.flip(-1)
97 d_v = d_v.flip(-1)
98 return d_u, d_v
99
100class TriangularToeplitzMultPadded(torch.autograd.Function):
101 @staticmethod

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected