MCPcopy Create free account
hub / github.com/apple/ml-pointersect / test

Method test

tests/pointersect/pr/cuda/test_cuda.py:110–202  ·  view source on GitHub ↗
(self, b=10, n=100, max_size=20)

Source from the content-addressed store, hash-verified

108
109class 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 = []

Callers 1

test2Method · 0.95

Calls 1

detachMethod · 0.45

Tested by

no test coverage detected