g3ms-pcp 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (35) hide show
  1. g3ms_pcp-0.1.0/.gitignore +10 -0
  2. g3ms_pcp-0.1.0/.python-version +1 -0
  3. g3ms_pcp-0.1.0/PKG-INFO +219 -0
  4. g3ms_pcp-0.1.0/README.md +191 -0
  5. g3ms_pcp-0.1.0/pyproject.toml +52 -0
  6. g3ms_pcp-0.1.0/src/pcp/__init__.py +13 -0
  7. g3ms_pcp-0.1.0/src/pcp/core/PointCloud.py +28 -0
  8. g3ms_pcp-0.1.0/src/pcp/core/__init__.py +3 -0
  9. g3ms_pcp-0.1.0/src/pcp/distance/__init__.py +3 -0
  10. g3ms_pcp-0.1.0/src/pcp/distance/chamfer.py +14 -0
  11. g3ms_pcp-0.1.0/src/pcp/grouping/__init__.py +5 -0
  12. g3ms_pcp-0.1.0/src/pcp/grouping/_query.py +58 -0
  13. g3ms_pcp-0.1.0/src/pcp/grouping/ball_query.py +44 -0
  14. g3ms_pcp-0.1.0/src/pcp/grouping/knn.py +32 -0
  15. g3ms_pcp-0.1.0/src/pcp/interpolation/IDW.py +43 -0
  16. g3ms_pcp-0.1.0/src/pcp/interpolation/__init__.py +3 -0
  17. g3ms_pcp-0.1.0/src/pcp/io/__init__.py +3 -0
  18. g3ms_pcp-0.1.0/src/pcp/io/csv.py +17 -0
  19. g3ms_pcp-0.1.0/src/pcp/nn/PointNet/PointNet.py +77 -0
  20. g3ms_pcp-0.1.0/src/pcp/nn/PointNet/PointNetCls.py +46 -0
  21. g3ms_pcp-0.1.0/src/pcp/nn/PointNet/PointNetSeg.py +44 -0
  22. g3ms_pcp-0.1.0/src/pcp/nn/PointNet/TNet.py +52 -0
  23. g3ms_pcp-0.1.0/src/pcp/nn/PointNet/__init__.py +5 -0
  24. g3ms_pcp-0.1.0/src/pcp/nn/PointNet2/FeaturePropagation.py +71 -0
  25. g3ms_pcp-0.1.0/src/pcp/nn/PointNet2/PointNet2.py +73 -0
  26. g3ms_pcp-0.1.0/src/pcp/nn/PointNet2/PointNet2Cls.py +54 -0
  27. g3ms_pcp-0.1.0/src/pcp/nn/PointNet2/PointNet2Seg.py +83 -0
  28. g3ms_pcp-0.1.0/src/pcp/nn/PointNet2/SetAbstraction.py +119 -0
  29. g3ms_pcp-0.1.0/src/pcp/nn/PointNet2/__init__.py +16 -0
  30. g3ms_pcp-0.1.0/src/pcp/nn/__init__.py +3 -0
  31. g3ms_pcp-0.1.0/src/pcp/py.typed +0 -0
  32. g3ms_pcp-0.1.0/src/pcp/sampling/__init__.py +4 -0
  33. g3ms_pcp-0.1.0/src/pcp/sampling/fps.py +42 -0
  34. g3ms_pcp-0.1.0/src/pcp/sampling/random.py +9 -0
  35. g3ms_pcp-0.1.0/uv.lock +1115 -0
@@ -0,0 +1,10 @@
1
+ # Python-generated files
2
+ __pycache__/
3
+ *.py[oc]
4
+ build/
5
+ dist/
6
+ wheels/
7
+ *.egg-info
8
+
9
+ # Virtual environments
10
+ .venv
@@ -0,0 +1 @@
1
+ 3.10.12
@@ -0,0 +1,219 @@
1
+ Metadata-Version: 2.5
2
+ Name: g3ms-pcp
3
+ Version: 0.1.0
4
+ Summary: Small point-cloud utilities built with PyTorch
5
+ Project-URL: Homepage, https://github.com/G3MS-Lab/pcp
6
+ Project-URL: Repository, https://github.com/G3MS-Lab/pcp
7
+ Project-URL: Issues, https://github.com/G3MS-Lab/pcp/issues
8
+ Author-email: Kittipong Tapyou <kittipong.tpy@gmail.com>
9
+ License: MIT
10
+ Classifier: Development Status :: 3 - Alpha
11
+ Classifier: Intended Audience :: Science/Research
12
+ Classifier: License :: OSI Approved :: MIT License
13
+ Classifier: Programming Language :: Python :: 3
14
+ Classifier: Programming Language :: Python :: 3.10
15
+ Classifier: Programming Language :: Python :: 3.11
16
+ Classifier: Programming Language :: Python :: 3.12
17
+ Classifier: Programming Language :: Python :: 3.13
18
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
19
+ Requires-Python: >=3.10
20
+ Requires-Dist: numpy>=2.0.0
21
+ Requires-Dist: polars>=1.0.0
22
+ Requires-Dist: torch>=2.0.0
23
+ Provides-Extra: dev
24
+ Requires-Dist: nbformat>=5.11.1; extra == 'dev'
25
+ Requires-Dist: pytest>=8.0.0; extra == 'dev'
26
+ Requires-Dist: ruff>=0.5.0; extra == 'dev'
27
+ Description-Content-Type: text/markdown
28
+
29
+ # pcp
30
+
31
+ Small point-cloud utilities built with PyTorch.
32
+
33
+ The package currently includes:
34
+
35
+ - a `PointCloud` tensor wrapper
36
+ - CSV loading
37
+ - farthest-point and random sampling
38
+ - k-nearest-neighbor search
39
+ - Chamfer distance
40
+ - inverse-distance weighted interpolation
41
+ - PointNet classification and segmentation models
42
+
43
+ ## Installation
44
+
45
+ Install the project in editable mode:
46
+
47
+ ```bash
48
+ pip install -e .
49
+ ```
50
+
51
+ Using `uv`:
52
+
53
+ ```bash
54
+ uv pip install -e .
55
+ ```
56
+
57
+ Python 3.10 or newer is required.
58
+
59
+ ## Basic usage
60
+
61
+ ```python
62
+ import torch
63
+
64
+ from pcp import PointCloud
65
+ from pcp.distance import chamfer
66
+ from pcp.grouping import knn
67
+ from pcp.sampling import fps, random
68
+
69
+ cloud = PointCloud(torch.rand(1024, 3)) # (N, C)
70
+
71
+ sampled = fps(cloud, k=256)
72
+ random_sampled = random(cloud, k=256)
73
+
74
+ center = torch.tensor([0.5, 0.5, 0.5])
75
+ neighbors = knn(cloud, center, k=16)
76
+
77
+ distance = chamfer(sampled, random_sampled)
78
+ ```
79
+
80
+ Interpolate scalar or vector values at new coordinates with IDW:
81
+
82
+ ```python
83
+ from pcp.interpolation import IDW
84
+
85
+ known_points = torch.tensor([[0.0, 0.0], [1.0, 0.0], [0.0, 1.0]])
86
+ known_values = torch.tensor([10.0, 20.0, 30.0])
87
+ query_points = torch.tensor([[0.25, 0.25]])
88
+
89
+ interpolated = IDW(
90
+ known_points,
91
+ known_values,
92
+ query_points,
93
+ alpha=2.0,
94
+ n_neighbors=3,
95
+ )
96
+ ```
97
+
98
+ Set `n_neighbors=None` to use every known point. IDW also supports batched
99
+ coordinates `(B, N, C)` and vector values `(B, N, F)`.
100
+
101
+ Load a point cloud from CSV:
102
+
103
+ ```python
104
+ from pcp.io import read_csv
105
+
106
+ cloud = read_csv("points.csv")
107
+ ```
108
+
109
+ ## PointNet
110
+
111
+ PointNet models accept batched tensors in `(B, N, C)` format:
112
+
113
+ - `B`: batch size
114
+ - `N`: number of points
115
+ - `C`: number of input channels
116
+
117
+ ```python
118
+ import torch
119
+
120
+ from pcp.nn.PointNet import PointNetCls, PointNetSeg
121
+
122
+ points = torch.randn(4, 1024, 3) # (B, N, C)
123
+
124
+ classifier = PointNetCls(
125
+ in_channels=3,
126
+ num_classes=10,
127
+ use_tnet=True,
128
+ )
129
+ class_logits = classifier(points) # (4, 10)
130
+
131
+ segmenter = PointNetSeg(
132
+ in_channels=3,
133
+ num_classes=4,
134
+ use_tnet=True,
135
+ )
136
+ point_logits = segmenter(points) # (4, 4, 1024)
137
+ ```
138
+
139
+ An unbatched `PointCloud` with shape `(N, C)` can be prepared for PointNet with:
140
+
141
+ ```python
142
+ batched_points = cloud.points.unsqueeze(0) # (1, N, C)
143
+ ```
144
+
145
+ Segmentation logits use `(B, num_classes, N)`, which can be passed directly to PyTorch's `CrossEntropyLoss` with targets shaped `(B, N)`.
146
+
147
+ ## Source layout
148
+
149
+ ```text
150
+ src/pcp/
151
+ ├── core/ # PointCloud
152
+ ├── distance/ # Chamfer distance
153
+ ├── grouping/ # k-nearest neighbors
154
+ ├── interpolation/ # inverse-distance weighting
155
+ ├── io/ # CSV loading
156
+ ├── nn/ # PointNet models
157
+ └── sampling/ # FPS and random sampling
158
+ ```
159
+
160
+ ## RepKPU point-cloud upsampling
161
+
162
+ `RepKPU` follows the official architecture: a three-stage PointTransformer
163
+ encoder, deformable RepKPoints Extraction Modules (REM), a KP-Queries
164
+ Generation Module (KGM), and cross-attention that turns queries into
165
+ displacements. The full decoder uses three REM blocks; `simple=True` uses one
166
+ REM block and the simple attention/skip path. The forward result is a tuple
167
+ `(upsampled_points, regularization_loss)`.
168
+
169
+ This library implementation uses tiled PyTorch neighborhood lookup rather
170
+ than the reference repository's compiled `pointops` CUDA kernels. Its distance
171
+ work remains quadratic in point count, though temporary distance memory is
172
+ bounded by `B * chunk_size * N`; the portable path is intended for patch-sized
173
+ inputs. The encoder and decoder reuse `pcp.grouping.knn` and
174
+ `pcp.grouping.ball_query`; both utilities retain their original return values
175
+ and add optional neighbor-index/validity-mask returns. It does not require
176
+ `pointops`, `einops`, or compiled Chamfer3D.
177
+
178
+ Inputs are batched coordinates with shape `(B, N, 3)`. The model returns
179
+ coordinates of shape `(B, N * up_rate, 3)`, with each input point followed by
180
+ its generated points. For example:
181
+
182
+ ```python
183
+ import torch
184
+
185
+ from pcp.nn.RepKPU import RepKPU, RepKPUConfig
186
+
187
+ config = RepKPUConfig(
188
+ up_rate=4,
189
+ encoder_dim=32,
190
+ out_dim=64,
191
+ k=16,
192
+ in_dim=64,
193
+ kp_dim=64,
194
+ num_kernel_points=15,
195
+ trans_dim=128,
196
+ head_num=4,
197
+ trans_num=3,
198
+ simple=False,
199
+ )
200
+ model = RepKPU(config)
201
+ points = torch.rand(2, 256, 3)
202
+ upsampled, reg_loss = model(points) # (2, 1024, 3), scalar loss
203
+ ```
204
+
205
+ Configuration fields mirror the official PU1K options: encoder width and
206
+ neighbor count, kernel radii and limits, KP dimensions, upsampling rate,
207
+ attention width/depth, and the simple/full decoder switch. The scalar loss is
208
+ the sum of the official deformable-kernel fitting and repulsion terms; include
209
+ it with the task loss during training. Input and output coordinates are
210
+ `(B, N, 3)` and `(B, N * up_rate, 3)`.
211
+
212
+ The module structure and forward mechanisms follow the official source, but
213
+ this is not a bit-for-bit port: neighborhood search is tiled PyTorch instead
214
+ of `pointops`, and the fixed kernel layouts use a deterministic centered
215
+ Fibonacci sphere instead of `kernel_utils.load_kernels`. Consequently model
216
+ weights and reported results are not checkpoint-compatible with official
217
+ RepKPU. The reference training/evaluation scripts also require their datasets
218
+ and custom Chamfer3D extension. The official code is MIT licensed; its notice
219
+ is included with the RepKPU package code.
@@ -0,0 +1,191 @@
1
+ # pcp
2
+
3
+ Small point-cloud utilities built with PyTorch.
4
+
5
+ The package currently includes:
6
+
7
+ - a `PointCloud` tensor wrapper
8
+ - CSV loading
9
+ - farthest-point and random sampling
10
+ - k-nearest-neighbor search
11
+ - Chamfer distance
12
+ - inverse-distance weighted interpolation
13
+ - PointNet classification and segmentation models
14
+
15
+ ## Installation
16
+
17
+ Install the project in editable mode:
18
+
19
+ ```bash
20
+ pip install -e .
21
+ ```
22
+
23
+ Using `uv`:
24
+
25
+ ```bash
26
+ uv pip install -e .
27
+ ```
28
+
29
+ Python 3.10 or newer is required.
30
+
31
+ ## Basic usage
32
+
33
+ ```python
34
+ import torch
35
+
36
+ from pcp import PointCloud
37
+ from pcp.distance import chamfer
38
+ from pcp.grouping import knn
39
+ from pcp.sampling import fps, random
40
+
41
+ cloud = PointCloud(torch.rand(1024, 3)) # (N, C)
42
+
43
+ sampled = fps(cloud, k=256)
44
+ random_sampled = random(cloud, k=256)
45
+
46
+ center = torch.tensor([0.5, 0.5, 0.5])
47
+ neighbors = knn(cloud, center, k=16)
48
+
49
+ distance = chamfer(sampled, random_sampled)
50
+ ```
51
+
52
+ Interpolate scalar or vector values at new coordinates with IDW:
53
+
54
+ ```python
55
+ from pcp.interpolation import IDW
56
+
57
+ known_points = torch.tensor([[0.0, 0.0], [1.0, 0.0], [0.0, 1.0]])
58
+ known_values = torch.tensor([10.0, 20.0, 30.0])
59
+ query_points = torch.tensor([[0.25, 0.25]])
60
+
61
+ interpolated = IDW(
62
+ known_points,
63
+ known_values,
64
+ query_points,
65
+ alpha=2.0,
66
+ n_neighbors=3,
67
+ )
68
+ ```
69
+
70
+ Set `n_neighbors=None` to use every known point. IDW also supports batched
71
+ coordinates `(B, N, C)` and vector values `(B, N, F)`.
72
+
73
+ Load a point cloud from CSV:
74
+
75
+ ```python
76
+ from pcp.io import read_csv
77
+
78
+ cloud = read_csv("points.csv")
79
+ ```
80
+
81
+ ## PointNet
82
+
83
+ PointNet models accept batched tensors in `(B, N, C)` format:
84
+
85
+ - `B`: batch size
86
+ - `N`: number of points
87
+ - `C`: number of input channels
88
+
89
+ ```python
90
+ import torch
91
+
92
+ from pcp.nn.PointNet import PointNetCls, PointNetSeg
93
+
94
+ points = torch.randn(4, 1024, 3) # (B, N, C)
95
+
96
+ classifier = PointNetCls(
97
+ in_channels=3,
98
+ num_classes=10,
99
+ use_tnet=True,
100
+ )
101
+ class_logits = classifier(points) # (4, 10)
102
+
103
+ segmenter = PointNetSeg(
104
+ in_channels=3,
105
+ num_classes=4,
106
+ use_tnet=True,
107
+ )
108
+ point_logits = segmenter(points) # (4, 4, 1024)
109
+ ```
110
+
111
+ An unbatched `PointCloud` with shape `(N, C)` can be prepared for PointNet with:
112
+
113
+ ```python
114
+ batched_points = cloud.points.unsqueeze(0) # (1, N, C)
115
+ ```
116
+
117
+ Segmentation logits use `(B, num_classes, N)`, which can be passed directly to PyTorch's `CrossEntropyLoss` with targets shaped `(B, N)`.
118
+
119
+ ## Source layout
120
+
121
+ ```text
122
+ src/pcp/
123
+ ├── core/ # PointCloud
124
+ ├── distance/ # Chamfer distance
125
+ ├── grouping/ # k-nearest neighbors
126
+ ├── interpolation/ # inverse-distance weighting
127
+ ├── io/ # CSV loading
128
+ ├── nn/ # PointNet models
129
+ └── sampling/ # FPS and random sampling
130
+ ```
131
+
132
+ ## RepKPU point-cloud upsampling
133
+
134
+ `RepKPU` follows the official architecture: a three-stage PointTransformer
135
+ encoder, deformable RepKPoints Extraction Modules (REM), a KP-Queries
136
+ Generation Module (KGM), and cross-attention that turns queries into
137
+ displacements. The full decoder uses three REM blocks; `simple=True` uses one
138
+ REM block and the simple attention/skip path. The forward result is a tuple
139
+ `(upsampled_points, regularization_loss)`.
140
+
141
+ This library implementation uses tiled PyTorch neighborhood lookup rather
142
+ than the reference repository's compiled `pointops` CUDA kernels. Its distance
143
+ work remains quadratic in point count, though temporary distance memory is
144
+ bounded by `B * chunk_size * N`; the portable path is intended for patch-sized
145
+ inputs. The encoder and decoder reuse `pcp.grouping.knn` and
146
+ `pcp.grouping.ball_query`; both utilities retain their original return values
147
+ and add optional neighbor-index/validity-mask returns. It does not require
148
+ `pointops`, `einops`, or compiled Chamfer3D.
149
+
150
+ Inputs are batched coordinates with shape `(B, N, 3)`. The model returns
151
+ coordinates of shape `(B, N * up_rate, 3)`, with each input point followed by
152
+ its generated points. For example:
153
+
154
+ ```python
155
+ import torch
156
+
157
+ from pcp.nn.RepKPU import RepKPU, RepKPUConfig
158
+
159
+ config = RepKPUConfig(
160
+ up_rate=4,
161
+ encoder_dim=32,
162
+ out_dim=64,
163
+ k=16,
164
+ in_dim=64,
165
+ kp_dim=64,
166
+ num_kernel_points=15,
167
+ trans_dim=128,
168
+ head_num=4,
169
+ trans_num=3,
170
+ simple=False,
171
+ )
172
+ model = RepKPU(config)
173
+ points = torch.rand(2, 256, 3)
174
+ upsampled, reg_loss = model(points) # (2, 1024, 3), scalar loss
175
+ ```
176
+
177
+ Configuration fields mirror the official PU1K options: encoder width and
178
+ neighbor count, kernel radii and limits, KP dimensions, upsampling rate,
179
+ attention width/depth, and the simple/full decoder switch. The scalar loss is
180
+ the sum of the official deformable-kernel fitting and repulsion terms; include
181
+ it with the task loss during training. Input and output coordinates are
182
+ `(B, N, 3)` and `(B, N * up_rate, 3)`.
183
+
184
+ The module structure and forward mechanisms follow the official source, but
185
+ this is not a bit-for-bit port: neighborhood search is tiled PyTorch instead
186
+ of `pointops`, and the fixed kernel layouts use a deterministic centered
187
+ Fibonacci sphere instead of `kernel_utils.load_kernels`. Consequently model
188
+ weights and reported results are not checkpoint-compatible with official
189
+ RepKPU. The reference training/evaluation scripts also require their datasets
190
+ and custom Chamfer3D extension. The official code is MIT licensed; its notice
191
+ is included with the RepKPU package code.
@@ -0,0 +1,52 @@
1
+ [project]
2
+ name = "g3ms-pcp"
3
+ version = "0.1.0"
4
+ description = "Small point-cloud utilities built with PyTorch"
5
+ readme = "README.md"
6
+ license = { text = "MIT" }
7
+ authors = [
8
+ { name = "Kittipong Tapyou", email = "kittipong.tpy@gmail.com" }
9
+ ]
10
+ requires-python = ">=3.10"
11
+ dependencies = [
12
+ "numpy>=2.0.0",
13
+ "polars>=1.0.0",
14
+ "torch>=2.0.0",
15
+ ]
16
+ classifiers = [
17
+ "Development Status :: 3 - Alpha",
18
+ "Intended Audience :: Science/Research",
19
+ "License :: OSI Approved :: MIT License",
20
+ "Programming Language :: Python :: 3",
21
+ "Programming Language :: Python :: 3.10",
22
+ "Programming Language :: Python :: 3.11",
23
+ "Programming Language :: Python :: 3.12",
24
+ "Programming Language :: Python :: 3.13",
25
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
26
+ ]
27
+
28
+ [project.urls]
29
+ Homepage = "https://github.com/G3MS-Lab/pcp"
30
+ Repository = "https://github.com/G3MS-Lab/pcp"
31
+ Issues = "https://github.com/G3MS-Lab/pcp/issues"
32
+
33
+ [project.optional-dependencies]
34
+ dev = [
35
+ "nbformat>=5.11.1",
36
+ "pytest>=8.0.0",
37
+ "ruff>=0.5.0",
38
+ ]
39
+
40
+ [build-system]
41
+ requires = ["hatchling"]
42
+ build-backend = "hatchling.build"
43
+
44
+ [tool.hatch.build.targets.wheel]
45
+ packages = ["src/pcp"]
46
+
47
+ [tool.hatch.build.targets.sdist]
48
+ exclude = [
49
+ "/.venv",
50
+ "/data",
51
+ "/notebooks",
52
+ ]
@@ -0,0 +1,13 @@
1
+ from pcp import core, distance, grouping, interpolation, io, sampling, nn
2
+ from pcp.core.PointCloud import PointCloud
3
+
4
+ __all__ = [
5
+ "core",
6
+ "distance",
7
+ "grouping",
8
+ "interpolation",
9
+ "io",
10
+ "sampling",
11
+ "nn",
12
+ "PointCloud",
13
+ ]
@@ -0,0 +1,28 @@
1
+ from dataclasses import dataclass
2
+ from typing import Any
3
+ import torch
4
+
5
+ @dataclass
6
+ class PointCloud:
7
+ points: torch.Tensor
8
+
9
+ def __init__(self, points: torch.Tensor):
10
+ self.points = points
11
+
12
+ def __getitem__(self, index: Any):
13
+ if isinstance(index, tuple):
14
+ return self.points[index]
15
+
16
+ new_pos = self.points[index]
17
+ if new_pos.ndim == 2:
18
+ new_pos = new_pos.unsqueeze(0)
19
+
20
+ return new_pos
21
+
22
+ @property
23
+ def shape(self):
24
+ return self.points.shape
25
+
26
+ def to_numpy(self):
27
+ return self.points.numpy()
28
+
@@ -0,0 +1,3 @@
1
+ from pcp.core.PointCloud import PointCloud
2
+
3
+ __all__ = ["PointCloud"]
@@ -0,0 +1,3 @@
1
+ from pcp.distance.chamfer import chamfer
2
+
3
+ __all__ = ["chamfer"]
@@ -0,0 +1,14 @@
1
+ import torch
2
+
3
+ from pcp.core.PointCloud import PointCloud
4
+
5
+
6
+ def chamfer(
7
+ points_a: PointCloud | torch.Tensor,
8
+ points_b: PointCloud | torch.Tensor,
9
+ ) -> torch.Tensor:
10
+ a = points_a.points if isinstance(points_a, PointCloud) else points_a
11
+ b = points_b.points if isinstance(points_b, PointCloud) else points_b
12
+
13
+ distances = torch.cdist(a, b)
14
+ return distances.min(dim=-1).values.mean() + distances.min(dim=-2).values.mean()
@@ -0,0 +1,5 @@
1
+ from pcp.grouping.ball_query import ball_query
2
+ from pcp.grouping.knn import knn
3
+ from pcp.grouping._query import index_features, index_points
4
+
5
+ __all__ = ["ball_query", "index_features", "index_points", "knn"]
@@ -0,0 +1,58 @@
1
+ """Shared tiled nearest-neighbor primitives for point grouping utilities."""
2
+
3
+ import torch
4
+
5
+
6
+ def nearest_indices(
7
+ points: torch.Tensor,
8
+ centroids: torch.Tensor,
9
+ k: int,
10
+ chunk_size: int = 128,
11
+ ) -> tuple[torch.Tensor, torch.Tensor]:
12
+ """Return k nearest source indices and distances for batched coordinates."""
13
+ k = min(k, points.shape[1])
14
+ index_parts = []
15
+ distance_parts = []
16
+ for start in range(0, centroids.shape[1], chunk_size):
17
+ query = centroids[:, start : start + chunk_size]
18
+ distances = torch.cdist(query, points)
19
+ distances, indices = distances.topk(k, dim=-1, largest=False)
20
+ index_parts.append(indices)
21
+ distance_parts.append(distances)
22
+ return torch.cat(index_parts, dim=1), torch.cat(distance_parts, dim=1)
23
+
24
+
25
+ def index_points(
26
+ points: torch.Tensor,
27
+ indices: torch.Tensor,
28
+ fill: float = 0.0,
29
+ ) -> torch.Tensor:
30
+ """Gather ``(B,N,C)`` points by ``(B,S[,K])`` indices; ``N`` means padding."""
31
+ batch, count, channels = points.shape
32
+ valid = indices < count
33
+ safe = indices.clamp(max=count - 1)
34
+ gathered = points.gather(
35
+ 1,
36
+ safe.reshape(batch, -1).unsqueeze(-1).expand(-1, -1, channels),
37
+ ).reshape(*indices.shape, channels)
38
+ return gathered.masked_fill(~valid.unsqueeze(-1), fill)
39
+
40
+
41
+ def index_features(
42
+ features: torch.Tensor,
43
+ indices: torch.Tensor,
44
+ fill: float = 0.0,
45
+ ) -> torch.Tensor:
46
+ """Gather ``(B,C,N)`` features into ``(B,C,S,K)`` with sentinel padding."""
47
+ batch, channels, count = features.shape
48
+ valid = indices < count
49
+ safe = indices.clamp(max=count - 1)
50
+ gathered = (
51
+ features.transpose(1, 2)
52
+ .gather(
53
+ 1,
54
+ safe.reshape(batch, -1).unsqueeze(-1).expand(-1, -1, channels),
55
+ )
56
+ .reshape(*indices.shape, channels)
57
+ )
58
+ return gathered.permute(0, 3, 1, 2).masked_fill(~valid.unsqueeze(1), fill)
@@ -0,0 +1,44 @@
1
+ import torch
2
+
3
+ from pcp.core.PointCloud import PointCloud
4
+ from pcp.grouping._query import index_points, nearest_indices
5
+
6
+
7
+ def ball_query(
8
+ radius: float,
9
+ n_samples: int,
10
+ points: PointCloud | torch.Tensor,
11
+ centroids: PointCloud | torch.Tensor,
12
+ return_mask: bool = False,
13
+ ):
14
+ points = points.points if isinstance(points, PointCloud) else points
15
+ centroids = centroids.points if isinstance(centroids, PointCloud) else centroids
16
+
17
+ unbatched = points.ndim == 2 and centroids.ndim == 2
18
+ if points.ndim == 2:
19
+ points = points.unsqueeze(0)
20
+ if centroids.ndim == 2:
21
+ centroids = centroids.unsqueeze(0)
22
+
23
+ _, num_points, _ = points.shape
24
+ k = min(n_samples, num_points)
25
+ group_idx, distances = nearest_indices(points, centroids, k)
26
+ valid = distances <= radius
27
+
28
+ first_idx = group_idx[:, :, :1]
29
+ first_idx = first_idx.repeat(1, 1, k)
30
+ group_idx = torch.where(valid, group_idx, first_idx)
31
+
32
+ if k < n_samples:
33
+ padding = group_idx[:, :, :1].repeat(1, 1, n_samples - k)
34
+ group_idx = torch.cat([group_idx, padding], dim=-1)
35
+ valid = torch.cat([valid, torch.zeros_like(padding, dtype=torch.bool)], dim=-1)
36
+
37
+ group_points = index_points(points, group_idx)
38
+
39
+ if unbatched:
40
+ group_idx = group_idx.squeeze(0)
41
+ group_points = group_points.squeeze(0)
42
+ valid = valid.squeeze(0)
43
+ result = (group_idx, group_points)
44
+ return (*result, valid) if return_mask else result
@@ -0,0 +1,32 @@
1
+ import torch
2
+
3
+ from pcp.core.PointCloud import PointCloud
4
+ from pcp.grouping._query import index_points, nearest_indices
5
+
6
+
7
+ def knn(
8
+ points: PointCloud | torch.Tensor,
9
+ centroids: PointCloud | torch.Tensor,
10
+ k: int,
11
+ return_indices: bool = False,
12
+ ):
13
+ points = points.points if isinstance(points, PointCloud) else points
14
+ centroids = centroids.points if isinstance(centroids, PointCloud) else centroids
15
+
16
+ unbatched = points.ndim == 2 and centroids.ndim <= 2
17
+ if centroids.ndim == 1:
18
+ centroids = centroids.unsqueeze(0)
19
+
20
+ if points.ndim == 2:
21
+ points = points.unsqueeze(0)
22
+ if centroids.ndim == 2:
23
+ centroids = centroids.unsqueeze(0)
24
+ indices, _ = nearest_indices(points, centroids, k)
25
+ result = index_points(points, indices)
26
+
27
+ if unbatched:
28
+ result = result.squeeze(0)
29
+ indices = indices.squeeze(0)
30
+
31
+ result = PointCloud(result)
32
+ return (result, indices) if return_indices else result