Some faster/better kernels. Tensors with points are now kept in (b, c, ...) format as this is easier for pytorch

This commit is contained in:
erikwijmans
2018-02-10 20:32:52 -05:00
parent 8bce353da4
commit 65a127f3d2
19 changed files with 372 additions and 296 deletions
+71 -93
View File
@@ -7,7 +7,59 @@ import pytorch_utils as pt_utils
from typing import List
class PointnetSAModuleMSG(nn.Module):
class _PointnetSAModuleBase(nn.Module):
def __init__(self):
super().__init__()
self.npoint = None
self.groupers = None
self.mlps = None
def forward(self, xyz: torch.Tensor,
points: torch.Tensor = None) -> (torch.Tensor, torch.Tensor):
r"""
Parameters
----------
xyz : torch.Tensor
(B, N, 3) tensor of the xyz coordinates of the points
point : torch.Tensor
(B, N, C) tensor of the descriptors of the the points
Returns
-------
new_xyz : torch.Tensor
(B, npoint, 3) tensor of the new points' xyz
new_points : torch.Tensor
(B, npoint, \sum_k(mlps[k][-1])) tensor of the new_points descriptors
"""
new_points_list = []
xyz_flipped = xyz.transpose(1, 2).contiguous()
new_xyz = pointnet2_utils.gather_points(
xyz_flipped,
pointnet2_utils.furthest_point_sample(xyz, self.npoint)
).transpose(1, 2).contiguous() if self.npoint is not None else None
for i in range(len(self.groupers)):
new_points = self.groupers[i](
xyz, new_xyz, points
) # (B, C, npoint, nsample)
new_points = self.mlps[i](
new_points
) # (B, mlp[-1], npoint, nsample)
new_points = F.max_pool2d(
new_points, kernel_size=[1, new_points.size(3)]
) # (B, mlp[-1], npoint, 1)
new_points = new_points.squeeze(-1) # (B, mlp[-1], npoint)
new_points_list.append(new_points)
return new_xyz, torch.cat(new_points_list, dim=1)
class PointnetSAModuleMSG(_PointnetSAModuleBase):
r"""Pointnet set abstrction layer with multiscale grouping
Parameters
@@ -48,49 +100,13 @@ class PointnetSAModuleMSG(nn.Module):
pointnet2_utils.QueryAndGroup(radius, nsample, use_xyz=use_xyz)
)
mlp_spec = mlps[i]
if use_xyz:
mlp_spec[0] += 3
self.mlps.append(pt_utils.SharedMLP(mlp_spec, bn=bn))
def forward(self, xyz: torch.Tensor,
points: torch.Tensor = None) -> (torch.Tensor, torch.Tensor):
r"""
Parameters
----------
xyz : torch.Tensor
(B, N, 3) tensor of the xyz coordinates of the points
point : torch.Tensor
(B, N, C) tensor of the descriptors of the the points
Returns
-------
new_xyz : torch.Tensor
(B, npoint, 3) tensor of the new points' xyz
new_points : torch.Tensor
(B, npoint, \sum_k(mlps[k][-1])) tensor of the new_points descriptors
"""
new_points_list = []
new_xyz = pointnet2_utils.gather_points(
xyz, pointnet2_utils.furthest_point_sample(xyz, self.npoint)
)
for i in range(len(self.groupers)):
new_points = self.groupers[i](xyz, new_xyz, points)
new_points = self.mlps[i](new_points.permute(0, 3, 1, 2)
) # (B, mlp[-1], npoint, nsample)
new_points = F.max_pool2d(
new_points, kernel_size=[1, new_points.size(3)]
) # (B, mlp[-1], npoint, 1)
new_points = new_points.squeeze(-1) # (B, mlp[-1], npoint)
new_points = new_points.transpose(
1, 2
).contiguous() # (B, npoint, mlp[-1])
new_points_list.append(new_points)
return new_xyz, torch.cat(new_points_list, dim=-1)
class PointnetSAModule(nn.Module):
class PointnetSAModule(_PointnetSAModuleBase):
r"""Pointnet set abstrction layer
Parameters
@@ -119,57 +135,22 @@ class PointnetSAModule(nn.Module):
):
super().__init__()
self.npoint = npoint
self.groupers = nn.ModuleList()
self.mlps = nn.ModuleList()
if self.npoint is not None:
assert radius is not None
assert nsample is not None
self.grouper = pointnet2_utils.QueryAndGroup(
radius, nsample, use_xyz=use_xyz
self.groupers.append(
pointnet2_utils.QueryAndGroup(radius, nsample, use_xyz=use_xyz)
)
else:
self.grouper = pointnet2_utils.GroupAll(use_xyz=use_xyz)
self.groupers.append(pointnet2_utils.GroupAll(use_xyz=use_xyz))
self.mlp = pt_utils.SharedMLP(mlp, bn=bn)
if use_xyz:
mlp[0] += 3
def forward(self, xyz: torch.Tensor,
points: torch.Tensor = None) -> (torch.Tensor, torch.Tensor):
r"""
Parameters
----------
xyz : torch.Tensor
(B, N, 3) tensor of the xyz coordinates of the points
point : torch.Tensor
(B, N, C) tensor of the descriptors of the the points
Returns
-------
new_xyz : torch.Tensor
(B, npoint, 3) tensor of the new points' xyz
new_points : torch.Tensor
(B, npoint, mlp[-1]) tensor of the new_points descriptors
"""
if self.npoint is not None:
new_xyz = pointnet2_utils.gather_points(
xyz, pointnet2_utils.furthest_point_sample(xyz, self.npoint)
)
else:
new_xyz = xyz.data.new([[[0, 0, 0]]]).expand(xyz.size(0), 1, 3)
new_points = self.grouper(
xyz, new_xyz, points
) # (B, npoint, nsample, 3 + C)
new_points = self.mlp(new_points.permute(0, 3, 1, 2)
) # (B, mlp[-1], npoint, nsample)
new_points = F.max_pool2d(
new_points, kernel_size=[1, new_points.size(3)]
) # (B, mlp[-1], npoint, 1)
new_points = new_points.squeeze(-1) # (B, mlp[-1], npoint)
new_points = new_points.transpose(1, 2
).contiguous() # (B, npoint, mlp[-1])
return new_xyz, new_points
self.mlps.append(pt_utils.SharedMLP(mlp, bn=bn))
class PointnetFPModule(nn.Module):
@@ -199,14 +180,14 @@ class PointnetFPModule(nn.Module):
known : torch.Tensor
(B, m, 3) tensor of the xyz positions of the known points
unknow_feats : torch.Tensor
(B, n, C1) tensor of the features to be propigated to
(B, C1, n) tensor of the features to be propigated to
known_feats : torch.Tensor
(B, m, C2) tensor of features to be propigated
(B, C2, m) tensor of features to be propigated
Returns
-------
new_points : torch.Tensor
(B, n, mlp[-1]) tensor of the features of the unknown points
(B, mlp[-1], n) tensor of the features of the unknown points
"""
dist, idx = pointnet2_utils.three_nn(unknown, known)
@@ -219,17 +200,14 @@ class PointnetFPModule(nn.Module):
)
if unknow_feats is not None:
new_points = torch.cat([interpolated_feats, unknow_feats],
dim=-1) #(B, n, C2 + C1)
dim=1) #(B, C2 + C1, n)
else:
new_points = interpolated_feats
new_points = new_points.unsqueeze(-1).transpose(
1, 2
) #(B, C2 + C1, n, 1)
new_points = new_points.unsqueeze(-1)
new_points = self.mlp(new_points)
return new_points.squeeze(-1).transpose(1, 2
).contiguous() #(B, n, mlp[-1])
return new_points.squeeze(-1)
if __name__ == "__main__":