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

Class TriangularToeplitzMultPaddedFast

model/functional/toeplitz.py:114–144  ·  view source on GitHub ↗

Trade off speed (20-25% faster) for more memory (20-25%)

Source from the content-addressed store, hash-verified

112 return d_u, d_v
113
114class 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_
147triangular_toeplitz_multiply = TriangularToeplitzMult.apply

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected