diff --git a/modules/pointops/src/knnquery/knnquery_cuda.cpp b/modules/pointops/src/knnquery/knnquery_cuda.cpp index 568f136..9ec5e75 100755 --- a/modules/pointops/src/knnquery/knnquery_cuda.cpp +++ b/modules/pointops/src/knnquery/knnquery_cuda.cpp @@ -1,5 +1,6 @@ #include #include +#include #include #include #include "knnquery_cuda_kernel.h" @@ -7,6 +8,13 @@ void knnquery_cuda(int m, int nsample, at::Tensor xyz_tensor, at::Tensor new_xyz_tensor, at::Tensor offset_tensor, at::Tensor new_offset_tensor, at::Tensor idx_tensor, at::Tensor dist2_tensor) { + TORCH_CHECK( + nsample >= 1 && nsample <= KNNQUERY_MAX_NEIGHBORS, + "knnquery nsample must be between 1 and ", + KNNQUERY_MAX_NEIGHBORS, + "; got ", + nsample + ); const float *xyz = xyz_tensor.data_ptr(); const float *new_xyz = new_xyz_tensor.data_ptr(); const int *offset = offset_tensor.data_ptr(); diff --git a/modules/pointops/src/knnquery/knnquery_cuda_kernel.cu b/modules/pointops/src/knnquery/knnquery_cuda_kernel.cu index 83762bc..764dafb 100755 --- a/modules/pointops/src/knnquery/knnquery_cuda_kernel.cu +++ b/modules/pointops/src/knnquery/knnquery_cuda_kernel.cu @@ -83,8 +83,8 @@ __global__ void knnquery_cuda_kernel(int m, int nsample, const float *__restrict float new_y = new_xyz[1]; float new_z = new_xyz[2]; - float best_dist[100]; - int best_idx[100]; + float best_dist[KNNQUERY_MAX_NEIGHBORS]; + int best_idx[KNNQUERY_MAX_NEIGHBORS]; for(int i = 0; i < nsample; i++){ best_dist[i] = 1e10; best_idx[i] = start; diff --git a/modules/pointops/src/knnquery/knnquery_cuda_kernel.h b/modules/pointops/src/knnquery/knnquery_cuda_kernel.h index 3c0aedf..41296f2 100755 --- a/modules/pointops/src/knnquery/knnquery_cuda_kernel.h +++ b/modules/pointops/src/knnquery/knnquery_cuda_kernel.h @@ -4,6 +4,8 @@ #include #include +constexpr int KNNQUERY_MAX_NEIGHBORS = 256; + void knnquery_cuda(int m, int nsample, at::Tensor xyz_tensor, at::Tensor new_xyz_tensor, at::Tensor offset_tensor, at::Tensor new_offset_tensor, at::Tensor idx_tensor, at::Tensor dist2_tensor); #ifdef __cplusplus diff --git a/tests/test_knnquery_capacity.py b/tests/test_knnquery_capacity.py new file mode 100644 index 0000000..f752f88 --- /dev/null +++ b/tests/test_knnquery_capacity.py @@ -0,0 +1,57 @@ +import unittest + +import torch + +from modules.pointops.functions import pointops + + +def make_points(count): + x = torch.linspace(0.0, 1.0, count, device="cuda", dtype=torch.float32) + zeros = torch.zeros_like(x) + xyz = torch.stack((x, zeros, zeros), dim=1).contiguous() + offset = torch.cuda.IntTensor([count]) + return xyz, offset + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA is required") +class KNNQueryCapacityTest(unittest.TestCase): + def test_supports_256_neighbors_without_invalid_indexes(self): + xyz, offset = make_points(256) + new_xyz = xyz[:1].contiguous() + new_offset = torch.cuda.IntTensor([1]) + + indexes, distances = pointops.knnquery( + 256, xyz, new_xyz, offset, new_offset + ) + torch.cuda.synchronize() + + self.assertEqual(tuple(indexes.shape), (1, 256)) + self.assertEqual(tuple(distances.shape), (1, 256)) + self.assertGreaterEqual(indexes.min().item(), 0) + self.assertLess(indexes.max().item(), xyz.shape[0]) + + features = torch.arange( + 256, device="cuda", dtype=torch.float32 + ).reshape(256, 1).contiguous() + output = pointops.interpolation( + xyz, new_xyz, features, offset, new_offset, k=256 + ) + torch.cuda.synchronize() + + self.assertEqual(tuple(output.shape), (1, 1)) + self.assertTrue(torch.isfinite(output).all().item()) + + def test_rejects_neighbor_count_above_capacity(self): + xyz, offset = make_points(257) + new_xyz = xyz[:1].contiguous() + new_offset = torch.cuda.IntTensor([1]) + + with self.assertRaisesRegex( + RuntimeError, + r"knnquery nsample must be between 1 and 256; got 257", + ): + pointops.knnquery(257, xyz, new_xyz, offset, new_offset) + + +if __name__ == "__main__": + unittest.main()