Add a few tests. Debug KNN.

This commit is contained in:
Austin Garrett
2018-03-20 19:41:54 -04:00
parent ffae381e7f
commit 6559f815be
4 changed files with 92 additions and 58 deletions
+24 -7
View File
@@ -4,7 +4,7 @@ import torch
from torch.autograd import Variable
import numpy as np
from pointcnn.core import XConv
from pointcnn.core import XConv, knn_indices_func
class BasicTests(unittest.TestCase):
""" Basic test cases """
@@ -19,12 +19,29 @@ class BasicTests(unittest.TestCase):
C_out = 32
N_neighbors = 100
model = XConv(C_in, C_out, D, N_neighbors)
test_p = Variable(torch.from_numpy(np.random.rand(N,D).astype(np.float32)))
test_P = Variable(torch.from_numpy(np.random.rand(N,N_neighbors,D).astype(np.float32)))
test_F = Variable(torch.from_numpy(np.random.rand(N,N_neighbors,C_in).astype(np.float32)))
test_out = model(test_p, test_P, test_F)
self.assertEqual(test_out.size(), (N, C_out))
model = XConv(C_in, C_out, D, N_neighbors).cuda()
p = Variable(torch.from_numpy(np.random.rand(N,D).astype(np.float32))).cuda()
P = Variable(torch.from_numpy(np.random.rand(N,N_neighbors,D).astype(np.float32))).cuda()
F = Variable(torch.from_numpy(np.random.rand(N,N_neighbors,C_in).astype(np.float32))).cuda()
out = model(p, P, F)
self.assertEqual(out.size(), (N, C_out))
def test_knn(self):
P = np.array([[[0,0],
[0,0.95],
[1,0],
[1,1]]])
ps = P[:,[0,3],:]
P = Variable(torch.from_numpy(P))
ps = Variable(torch.from_numpy(ps))
out = knn_indices_func(ps, P, 2).numpy()
target = np.array([[[1,2],
[2,1]]])
self.assertTrue(np.array_equal(target, out))
if __name__ == "__main__":
unittest.main()