Implement XConv operator. Begin implementing PointCNN.

This commit is contained in:
Austin J. Garrett
2018-03-17 20:27:49 -04:00
parent 805aba436d
commit 50b77758df
10 changed files with 186 additions and 15 deletions
+20 -2
View File
@@ -1,12 +1,30 @@
import unittest
from context import pointcnn
import torch
from torch.autograd import Variable
import numpy as np
from pointcnn.core import XConv
class BasicTests(unittest.TestCase):
""" Basic test cases """
def test_example(self):
def test_xconv_shape(self):
self.assertTrue(True)
np.random.seed(0)
N = 4
D = 3
C_in = 8
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))
if __name__ == "__main__":
unittest.main()