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)
| 372 | class 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): |