MCPcopy Create free account
hub / github.com/OpenImagingLab/4DSloMo / forward

Method forward

pointops2/functions/pointops.py:270–290  ·  view source on GitHub ↗

input: attn: (M, h), v: (N, h, C//h), index0: (M), index1: (M) output: output: [L, h, C//h]

(ctx, attn, v, index0, index1)

Source from the content-addressed store, hash-verified

268class AttentionStep2_v2(Function):
269 @staticmethod
270 def forward(ctx, attn, v, index0, index1):
271 """
272 input: attn: (M, h), v: (N, h, C//h), index0: (M), index1: (M)
273 output: output: [L, h, C//h]
274 """
275 assert attn.is_contiguous() and v.is_contiguous() and index0.is_contiguous() and index1.is_contiguous()
276
277 L = int(index0.max().item()) + 1
278
279 M, h = attn.shape
280 N, h, C_div_h = v.shape
281 C = int(C_div_h * h)
282
283 output = torch.cuda.FloatTensor(L, h, C//h).zero_()
284 pointops_cuda.attention_step2_forward_cuda(N, M, h, C, attn, v, index0, index1, output)
285 ctx.M = M
286
287 # print("attn[:5,:5]: ", attn[:5, :5])
288
289 ctx.save_for_backward(attn, v, index0, index1)
290 return output
291
292 @staticmethod
293 def backward(ctx, grad_output):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected