MCPcopy Create free account
hub / github.com/VAST-AI-Research/TriplaneGaussian / PointNet_SA_Module_KNN

Class PointNet_SA_Module_KNN

tgs/models/snowflake/utils.py:334–383  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

332
333
334class PointNet_SA_Module_KNN(nn.Module):
335 def __init__(self, npoint, nsample, in_channel, mlp, if_bn=True, group_all=False, use_xyz=True, if_idx=False):
336 """
337 Args:
338 npoint: int, number of points to sample
339 nsample: int, number of points in each local region
340 radius: float
341 in_channel: int, input channel of features(points)
342 mlp: list of int,
343 """
344 super(PointNet_SA_Module_KNN, self).__init__()
345 self.npoint = npoint
346 self.nsample = nsample
347 self.mlp = mlp
348 self.group_all = group_all
349 self.use_xyz = use_xyz
350 self.if_idx = if_idx
351 if use_xyz:
352 in_channel += 3
353
354 last_channel = in_channel
355 self.mlp_conv = []
356 for out_channel in mlp[:-1]:
357 self.mlp_conv.append(Conv2d(last_channel, out_channel, if_bn=if_bn))
358 last_channel = out_channel
359 self.mlp_conv.append(Conv2d(last_channel, mlp[-1], if_bn=False, activation_fn=None))
360 self.mlp_conv = nn.Sequential(*self.mlp_conv)
361
362 def forward(self, xyz, points, idx=None):
363 """
364 Args:
365 xyz: Tensor, (B, 3, N)
366 points: Tensor, (B, f, N)
367
368 Returns:
369 new_xyz: Tensor, (B, 3, npoint)
370 new_points: Tensor, (B, mlp[-1], npoint)
371 """
372 if self.group_all:
373 new_xyz, new_points, idx, grouped_xyz = sample_and_group_all(xyz, points, self.use_xyz)
374 else:
375 new_xyz, new_points, idx, grouped_xyz = sample_and_group_knn(xyz, points, self.npoint, self.nsample, self.use_xyz, idx=idx)
376
377 new_points = self.mlp_conv(new_points)
378 new_points = torch.max(new_points, 3)[0]
379
380 if self.if_idx:
381 return new_xyz, new_points, idx
382 else:
383 return new_xyz, new_points
384
385
386def fps_subsample(pcd, n_points=2048):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected