| 108 | |
| 109 | class TestGatherPoints(unittest.TestCase): |
| 110 | def test(self, b=10, n=100, max_size=20): |
| 111 | if not torch.cuda.is_available(): |
| 112 | return |
| 113 | |
| 114 | size = torch.randint(max_size-1, size=[b, 3]) + 1 # (b, 3) |
| 115 | total_cells = torch.prod(size, dim=-1) # (b,) |
| 116 | grid_idxs = (torch.rand(b, n) * total_cells.unsqueeze(1)).floor().long() |
| 117 | valid_mask = torch.randn(b, n) > 0.5 |
| 118 | |
| 119 | stime = timer() |
| 120 | all_cell2pidx_gt = naive.gather_points( |
| 121 | grid_idxs=grid_idxs, |
| 122 | total_cells=total_cells, |
| 123 | valid_mask=valid_mask, |
| 124 | ) |
| 125 | time_python = timer() - stime |
| 126 | |
| 127 | stime = timer() |
| 128 | _all_cell2pidx = pr_cpp.gather_points( |
| 129 | grid_idxs, |
| 130 | total_cells, |
| 131 | valid_mask, |
| 132 | ) |
| 133 | time_cpp = timer() - stime |
| 134 | |
| 135 | # compute cell_counts |
| 136 | max_num_cells = total_cells.max() |
| 137 | cell_counts = torch.zeros(b, max_num_cells, dtype=torch.int) |
| 138 | for bidx in range(b): |
| 139 | for pidx in range(n): |
| 140 | if valid_mask[bidx, pidx]: |
| 141 | cell_counts[bidx, grid_idxs[bidx,pidx]] += 1 |
| 142 | |
| 143 | grid_idxs = grid_idxs.cuda() |
| 144 | total_cells = total_cells.cuda() |
| 145 | valid_mask = valid_mask.cuda() |
| 146 | cell_counts = cell_counts.cuda() |
| 147 | |
| 148 | stime = timer() |
| 149 | all_cell2pidx = pr_cuda.gather_points_v1( |
| 150 | grid_idxs, |
| 151 | total_cells, |
| 152 | valid_mask, |
| 153 | cell_counts, |
| 154 | ) |
| 155 | time_cuda_v1 = timer() - stime |
| 156 | |
| 157 | stime = timer() |
| 158 | membank, cell_start_idxs = pr_cuda.gather_points_v2( |
| 159 | grid_idxs, |
| 160 | total_cells, |
| 161 | valid_mask, |
| 162 | cell_counts, |
| 163 | ) # membank: (b, n), cell_start_idxs: (b, n_cell+1) |
| 164 | time_cuda_v2 = timer() - stime |
| 165 | |
| 166 | # contruct all_cell2pidx from membank |
| 167 | all_cell2pidx_v2 = [] |