(self)
| 168 | ) |
| 169 | |
| 170 | def test(self): |
| 171 | print(f"{self.__class__.__name__}: test") |
| 172 | in_channels, out_channels, D = 2, 3, 2 |
| 173 | coords, feats, labels = data_loader(in_channels) |
| 174 | feats = feats.double() |
| 175 | feats.requires_grad_() |
| 176 | input = SparseTensor(feats, coordinates=coords) |
| 177 | # Initialize context |
| 178 | conv = MinkowskiConvolution( |
| 179 | in_channels, out_channels, kernel_size=3, stride=2, bias=True, dimension=D |
| 180 | ) |
| 181 | conv = conv.double() |
| 182 | output = conv(input) |
| 183 | print(output) |
| 184 | |
| 185 | self.assertEqual(input.coordinate_map_key.get_tensor_stride(), [1, 1]) |
| 186 | self.assertEqual(output.coordinate_map_key.get_tensor_stride(), [2, 2]) |
| 187 | |
| 188 | if torch.cuda.is_available(): |
| 189 | input_gpu = SparseTensor(feats, coordinates=coords, device="cuda") |
| 190 | conv_gpu = conv.cuda() |
| 191 | output_gpu = conv_gpu(input_gpu) |
| 192 | self.assertTrue(torch.allclose(output_gpu.F.var(0).cpu(), output.F.var(0))) |
| 193 | self.assertTrue( |
| 194 | torch.allclose(output_gpu.F.mean(0).cpu(), output.F.mean(0)) |
| 195 | ) |
| 196 | |
| 197 | # kernel_map = input.coords_man.kernel_map( |
| 198 | # 1, 2, stride=2, kernel_size=3) |
| 199 | # print(kernel_map) |
| 200 | |
| 201 | # Check backward |
| 202 | fn = MinkowskiConvolutionFunction() |
| 203 | |
| 204 | conv = conv.cpu() |
| 205 | self.assertTrue( |
| 206 | gradcheck( |
| 207 | fn, |
| 208 | ( |
| 209 | input.F, |
| 210 | conv.kernel, |
| 211 | conv.kernel_generator, |
| 212 | conv.convolution_mode, |
| 213 | input.coordinate_map_key, |
| 214 | output.coordinate_map_key, |
| 215 | input.coordinate_manager, |
| 216 | ), |
| 217 | ) |
| 218 | ) |
| 219 | |
| 220 | for i in range(LEAK_TEST_ITER): |
| 221 | input = SparseTensor(feats, coordinates=coords) |
| 222 | conv(input).F.sum().backward() |
| 223 | if i % 1000 == 0: |
| 224 | print(i) |
| 225 | |
| 226 | def test_analytic(self): |
| 227 | print(f"{self.__class__.__name__}: test") |
nothing calls this directly
no test coverage detected