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

Method forward

pointops2/functions/pointops.py:374–406  ·  view source on GitHub ↗

input: q: (N, h, hdim), index_q: (M), k: (N, h, hdim), index_k: (M), table_q: (L, h, hdim, 3), table_k: (L, h, hdim, 3), rel_idx: (M, 3) output: output: [M, h]

(ctx, q, index_q, k, index_k, table_q, table_k, rel_idx)

Source from the content-addressed store, hash-verified

372class DotProdWithIdx_v2(Function):
373 @staticmethod
374 def forward(ctx, q, index_q, k, index_k, table_q, table_k, rel_idx):
375 """
376 input: q: (N, h, hdim), index_q: (M), k: (N, h, hdim), index_k: (M), table_q: (L, h, hdim, 3), table_k: (L, h, hdim, 3), rel_idx: (M, 3)
377 output: output: [M, h]
378 """
379 assert q.is_contiguous() and index_q.is_contiguous() and k.is_contiguous() and index_k.is_contiguous() and table_q.is_contiguous() and table_k.is_contiguous() and rel_idx.is_contiguous()
380
381 N, h, hdim = q.shape
382 M = index_q.shape[0]
383 L = table_q.shape[0]
384 assert table_k.shape[0] == L and index_k.shape[0] == M
385
386 # obtain the mapping from block_idx to m_idx
387 rel_idx_merge = rel_idx[:, 0] + rel_idx[:, 1] * L + rel_idx[:, 2] * (L ** 2) #[M, ]
388 sorted_values, sort_indices = torch.sort(rel_idx_merge)
389 _, counts = torch.unique_consecutive(sorted_values, return_counts=True)
390 rel_idx_offsets = torch.cumsum(counts, dim=-1) #[T,]
391 rel_idx_offsets = torch.cat([torch.zeros(1, dtype=torch.long).cuda(), rel_idx_offsets], 0) #[T+1,]
392 n_max = counts.max()
393 T = counts.shape[0]
394
395 # print("M: {}, L: {}, n_max: {}, T: {}".format(M, L, n_max, T))
396 # print("rel_idx_merge.shape: {}, sorted_values.shape: {}".format(rel_idx_merge.shape, sorted_values.shape))
397 # print("counts.shape: {}".format(counts.shape))
398
399 output = torch.cuda.FloatTensor(M, h).zero_()
400 # pointops_cuda.dot_prod_with_idx_forward_cuda(N, M, h, hdim, q, index, table, rel_idx, output)
401 pointops_cuda.dot_prod_with_idx_forward_cuda_v2(N, M, h, hdim, n_max, T, q, index_q, k, index_k, table_q, table_k, rel_idx, rel_idx_offsets.int(), sort_indices.int(), output)
402
403 ctx.n_max = n_max
404 ctx.T = T
405 ctx.save_for_backward(q, index_q, k, index_k, table_q, table_k, rel_idx, rel_idx_offsets, sort_indices)
406 return output
407
408 @staticmethod
409 def backward(ctx, grad_output):

Callers

nothing calls this directly

Calls 1

cudaMethod · 0.80

Tested by

no test coverage detected