mps-pointops 0.3.0__py3-none-any.whl

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.
@@ -0,0 +1,7 @@
1
+ """Point cloud ops for PyTorch on Apple Silicon (MPS)."""
2
+
3
+ from . import flat, reference
4
+ from .ops import ball_query, furthest_point_sample, knn
5
+
6
+ __all__ = ["ball_query", "flat", "furthest_point_sample", "knn", "reference"]
7
+ __version__ = "0.3.0"
@@ -0,0 +1,240 @@
1
+ # SPDX-License-Identifier: MIT
2
+ """Deterministic radius-neighbor search on Apple GPUs.
3
+
4
+ The Metal source is original to this project. The public contract follows the
5
+ mathematical definition of radius search and the documented first-K convention
6
+ used by PyTorch3D; no third-party implementation is copied here.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ from functools import lru_cache
12
+ from importlib.resources import files
13
+ import math
14
+ import struct
15
+ from typing import NamedTuple
16
+
17
+ import torch
18
+
19
+
20
+ _MIN_SUPPORTED_RADIUS_F32 = 2.0**-112
21
+
22
+
23
+ class BallQueryResult(NamedTuple):
24
+ """Squared distances and indices, both shaped ``(B, Q, K)``."""
25
+
26
+ distances: torch.Tensor
27
+ indices: torch.Tensor
28
+
29
+
30
+ def _require_compile_shader() -> None:
31
+ if not hasattr(torch.mps, "compile_shader"):
32
+ raise RuntimeError(
33
+ f"mps_pointops needs PyTorch 2.7 or later (torch.mps.compile_shader); found {torch.__version__}"
34
+ )
35
+
36
+
37
+ @lru_cache(maxsize=1)
38
+ def _shader_library():
39
+ _require_compile_shader()
40
+ source = files(__package__).joinpath("kernels", "ball_query.metal").read_text()
41
+ return torch.mps.compile_shader(source)
42
+
43
+
44
+ def _lengths_or_full(
45
+ lengths: torch.Tensor | None, *, batch: int, count: int, device: torch.device, name: str
46
+ ) -> torch.Tensor:
47
+ if lengths is None:
48
+ return torch.full((batch,), count, device=device, dtype=torch.int64)
49
+ if not isinstance(lengths, torch.Tensor):
50
+ raise TypeError(f"{name} must be a torch.Tensor or None")
51
+ if lengths.shape != (batch,) or lengths.dtype != torch.int64 or lengths.device != device:
52
+ raise ValueError(f"{name} must have shape ({batch},), int64 dtype and device {device}")
53
+ if torch.any((lengths < 0) | (lengths > count)).item():
54
+ raise ValueError(f"{name} values must lie in [0, {count}]")
55
+ return lengths.contiguous()
56
+
57
+
58
+ def _checked_radius_and_k(radius: float, k: int) -> tuple[float, float]:
59
+ """Validate ``radius`` and ``k`` for every ball query path, on any device.
60
+
61
+ Returns ``(radius_f32, radius_squared)``: the radius rounded to float32 and
62
+ fl32(fl32(radius) * fl32(radius)).
63
+ """
64
+ if not isinstance(k, int) or isinstance(k, bool) or k < 0 or k > (1 << 63) - 1:
65
+ raise ValueError("k must be a non-negative integer")
66
+ if not isinstance(radius, (int, float)) or isinstance(radius, bool):
67
+ raise ValueError("radius must be a finite non-negative number")
68
+ try:
69
+ radius = float(radius)
70
+ except OverflowError as exc:
71
+ raise ValueError("radius must be a finite non-negative number") from exc
72
+ if not math.isfinite(radius) or radius < 0:
73
+ raise ValueError("radius must be a finite non-negative number")
74
+
75
+ try:
76
+ radius_f32 = struct.unpack("f", struct.pack("f", radius))[0]
77
+ except OverflowError as exc:
78
+ raise ValueError("radius must fit in float32") from exc
79
+ if not math.isfinite(radius_f32) or (radius > 0 and radius_f32 < _MIN_SUPPORTED_RADIUS_F32):
80
+ raise ValueError("positive radius must round to a normal float32 value at least 2**-112")
81
+
82
+ # This is fl32(fl32(radius) * fl32(radius)), as in PyTorch3D's float
83
+ # radius2 = radius * radius. Two binary32 operands multiply exactly in
84
+ # binary64 (at most 48 significant bits); struct.pack rounds once to f32.
85
+ radius_squared = float(radius_f32) * float(radius_f32)
86
+ try:
87
+ radius_squared = struct.unpack("f", struct.pack("f", radius_squared))[0]
88
+ except OverflowError as exc:
89
+ radius_squared = math.inf
90
+ return radius_f32, radius_squared
91
+
92
+
93
+ class _BallQuery(torch.autograd.Function):
94
+ @staticmethod
95
+ def forward(
96
+ ctx,
97
+ queries: torch.Tensor,
98
+ points: torch.Tensor,
99
+ query_lengths: torch.Tensor,
100
+ point_lengths: torch.Tensor,
101
+ radius_sq: float,
102
+ radius_f32: float,
103
+ k: int,
104
+ ) -> tuple[torch.Tensor, torch.Tensor]:
105
+ batch, query_count, _ = queries.shape
106
+ point_count = points.shape[1]
107
+ indices = torch.empty((batch, query_count, k), device=queries.device, dtype=torch.int64)
108
+ distances = torch.empty((batch, query_count, k), device=queries.device, dtype=torch.float32)
109
+
110
+ if indices.numel():
111
+ kernel = (
112
+ _shader_library().ball_query_f32
113
+ if queries.dtype == torch.float32
114
+ else _shader_library().ball_query_f16
115
+ )
116
+ kernel(
117
+ queries,
118
+ points,
119
+ query_lengths,
120
+ point_lengths,
121
+ indices,
122
+ distances,
123
+ batch,
124
+ query_count,
125
+ point_count,
126
+ k,
127
+ radius_sq,
128
+ radius_f32,
129
+ threads=[batch * query_count, 1, 1],
130
+ group_size=[256, 1, 1],
131
+ )
132
+
133
+ ctx.save_for_backward(queries, points, indices)
134
+ ctx.mark_non_differentiable(indices)
135
+ return distances, indices
136
+
137
+ @staticmethod
138
+ def backward(ctx, grad_distances: torch.Tensor | None, grad_indices: None):
139
+ if grad_distances is None:
140
+ return None, None, None, None, None, None, None
141
+
142
+ queries, points, indices = ctx.saved_tensors
143
+ batch, query_count, k = indices.shape
144
+ if not indices.numel() or points.shape[1] == 0:
145
+ return torch.zeros_like(queries), torch.zeros_like(points), None, None, None, None, None
146
+
147
+ valid = indices >= 0
148
+ safe_indices = indices.clamp_min(0)
149
+ gathered = points.float().gather(
150
+ 1, safe_indices.reshape(batch, -1, 1).expand(-1, -1, 3)
151
+ ).reshape(batch, query_count, k, 3)
152
+ delta = queries.float().unsqueeze(2) - gathered
153
+ # Padded slots may gather a non-finite point 0. Mask before multiplying:
154
+ # IEEE-754 makes 0 * NaN equal NaN, not zero.
155
+ delta = torch.where(valid.unsqueeze(-1), delta, torch.zeros_like(delta))
156
+ valid_grad = torch.where(valid, grad_distances, torch.zeros_like(grad_distances))
157
+ coeff = (2.0 * valid_grad).unsqueeze(-1)
158
+ contributions = coeff * delta
159
+
160
+ grad_queries = (
161
+ contributions.sum(dim=2).to(queries.dtype) if ctx.needs_input_grad[0] else None
162
+ )
163
+ grad_points = None
164
+ if ctx.needs_input_grad[1]:
165
+ grad_points = torch.zeros_like(points, dtype=torch.float32)
166
+ grad_points.scatter_add_(
167
+ 1,
168
+ safe_indices.reshape(batch, -1, 1).expand(-1, -1, 3),
169
+ -contributions.reshape(batch, -1, 3),
170
+ )
171
+ grad_points = grad_points.to(points.dtype)
172
+ return grad_queries, grad_points, None, None, None, None, None
173
+
174
+
175
+ def ball_query(
176
+ queries: torch.Tensor,
177
+ points: torch.Tensor,
178
+ *,
179
+ radius: float,
180
+ k: int,
181
+ query_lengths: torch.Tensor | None = None,
182
+ point_lengths: torch.Tensor | None = None,
183
+ ) -> BallQueryResult:
184
+ """Find the first ``k`` points strictly inside each radius.
185
+
186
+ ``queries`` and ``points`` are float32 or float16 MPS tensors with shapes ``(B,Q,3)``
187
+ and ``(B,P,3)``. The output is ordered by input point index, not distance.
188
+ Missing neighbors use index ``-1`` and squared distance ``0``. Under the
189
+ supported Safe math mode, non-finite coordinates never match. Gradients of
190
+ valid squared distances propagate to both coordinate tensors; selection
191
+ indices are non-differentiable. ``radius`` is a Python number, not a
192
+ differentiable Tensor.
193
+
194
+ Coordinates remain on MPS. Supplying lengths checks their bounds with a
195
+ scalar MPS-to-CPU synchronization. Non-contiguous coordinate tensors are
196
+ copied to contiguous MPS storage before dispatch. The radius is rounded
197
+ to float32 first and must then be zero or at least 2**-112. Smaller
198
+ positive values are rejected because Metal may flush coordinate differences
199
+ that matter to the radius decision. Boundary decisions use float32 arithmetic.
200
+ """
201
+ if not isinstance(queries, torch.Tensor) or not isinstance(points, torch.Tensor):
202
+ raise TypeError("queries and points must be torch.Tensor values")
203
+ if queries.ndim != 3 or points.ndim != 3 or queries.shape[-1] != 3 or points.shape[-1] != 3:
204
+ raise ValueError("queries and points must have shapes (B, Q, 3) and (B, P, 3)")
205
+ if queries.shape[0] != points.shape[0]:
206
+ raise ValueError("queries and points must have the same batch size")
207
+ if queries.device.type != "mps" or points.device != queries.device:
208
+ raise ValueError("queries and points must be on the same MPS device")
209
+ if queries.dtype not in (torch.float32, torch.float16) or points.dtype != queries.dtype:
210
+ raise TypeError("queries and points must have the same float32 or float16 dtype")
211
+ radius_f32, radius_squared = _checked_radius_and_k(radius, k)
212
+
213
+ batch, query_count, _ = queries.shape
214
+ point_count = points.shape[1]
215
+ if batch * query_count > (1 << 32) - 1:
216
+ raise ValueError("B * Q exceeds the Metal 1-D dispatch limit")
217
+ if batch * query_count * 3 > (1 << 63) - 1 or batch * point_count * 3 > (1 << 63) - 1:
218
+ raise ValueError("coordinate element count exceeds int64 indexing")
219
+ if batch * query_count * k > (1 << 63) - 1:
220
+ raise ValueError("output element count exceeds int64 indexing")
221
+
222
+ if not torch.backends.mps.is_available():
223
+ raise RuntimeError("PyTorch MPS is unavailable on this host")
224
+ query_lengths = _lengths_or_full(
225
+ query_lengths, batch=batch, count=query_count, device=queries.device, name="query_lengths"
226
+ )
227
+ point_lengths = _lengths_or_full(
228
+ point_lengths, batch=batch, count=point_count, device=queries.device, name="point_lengths"
229
+ )
230
+
231
+ distances, indices = _BallQuery.apply(
232
+ queries.contiguous(),
233
+ points.contiguous(),
234
+ query_lengths,
235
+ point_lengths,
236
+ radius_squared,
237
+ radius_f32,
238
+ k,
239
+ )
240
+ return BallQueryResult(distances=distances, indices=indices)
@@ -0,0 +1,96 @@
1
+ """Farthest point sampling for flat, variable-length point clouds on MPS."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from functools import lru_cache
6
+ from importlib.resources import files
7
+
8
+ import torch
9
+ from torch import Tensor
10
+
11
+ from ._ball_query_mps import _require_compile_shader
12
+
13
+
14
+ _THREADS = 1024
15
+ _MAX_LOCAL_POINTS = (1 << 32) - 1
16
+
17
+
18
+ @lru_cache(maxsize=1)
19
+ def _shader_library():
20
+ _require_compile_shader()
21
+ width_shader = torch.mps.compile_shader(
22
+ "kernel void width(device long* out, uint w [[threads_per_simdgroup]]) { out[0] = w; }"
23
+ )
24
+ width = torch.empty(1, dtype=torch.long, device="mps")
25
+ width_shader.width(width, threads=1, group_size=1)
26
+ actual_width = int(width.item())
27
+ if actual_width != 32:
28
+ raise RuntimeError(
29
+ f"mps_pointops kernels need 32-wide simdgroups, this GPU uses {actual_width}"
30
+ )
31
+ source = files(__package__).joinpath("kernels", "fps_flat.metal").read_text()
32
+ return torch.mps.compile_shader(source)
33
+
34
+
35
+ def fps_flat(x: Tensor, ptr: Tensor, out_ptr: Tensor, starts: Tensor) -> Tensor:
36
+ """Sample flat ``x`` by segment and return global int64 point indices.
37
+
38
+ ``ptr[b:b+2]`` bounds input cloud ``b``; ``out_ptr[b:b+2]`` bounds its
39
+ output. ``starts[b]`` is a *local* index in that cloud. A cloud with zero
40
+ requested samples may be empty. As in the existing FPS kernel, requests
41
+ longer than a cloud repeat its local index zero after exhaustion.
42
+ """
43
+ if not all(isinstance(t, Tensor) for t in (x, ptr, out_ptr, starts)):
44
+ raise TypeError("x, ptr, out_ptr, and starts must be torch.Tensor values")
45
+ if x.ndim != 2 or x.shape[1] != 3:
46
+ raise ValueError(f"x must have shape (Total_N, 3), got {tuple(x.shape)}")
47
+ if ptr.ndim != 1 or ptr.numel() == 0:
48
+ raise ValueError("ptr must be a nonempty one-dimensional tensor")
49
+ batch = ptr.numel() - 1
50
+ if out_ptr.shape != (batch + 1,) or starts.shape != (batch,):
51
+ raise ValueError("out_ptr must have shape (B+1,) and starts must have shape (B,)")
52
+ if x.device.type != "mps" or any(t.device != x.device for t in (ptr, out_ptr, starts)):
53
+ raise ValueError("x, ptr, out_ptr, and starts must be on the same MPS device")
54
+ if x.dtype != torch.float32:
55
+ raise TypeError(f"x must be float32 on MPS, got {x.dtype}")
56
+ if any(t.dtype != torch.int64 for t in (ptr, out_ptr, starts)):
57
+ raise TypeError("ptr, out_ptr, and starts must be int64")
58
+ if batch > (1 << 32) - 1:
59
+ raise ValueError("B exceeds the Metal grid limit")
60
+
61
+ x = x.contiguous()
62
+ ptr = ptr.contiguous()
63
+ out_ptr = out_ptr.contiguous()
64
+ starts = starts.contiguous()
65
+ lengths = ptr[1:] - ptr[:-1]
66
+ counts = out_ptr[1:] - out_ptr[:-1]
67
+ invalid = (
68
+ (ptr[0] != 0)
69
+ | (ptr[-1] != x.shape[0])
70
+ | (out_ptr[0] != 0)
71
+ | torch.any(lengths < 0)
72
+ | torch.any(lengths > _MAX_LOCAL_POINTS)
73
+ | torch.any(counts < 0)
74
+ | torch.any((counts > 0) & ((lengths == 0) | (starts < 0) | (starts >= lengths)))
75
+ )
76
+ # Output shape is data-dependent. Transfer only its total size and the
77
+ # validation flag; coordinates and segment offsets remain on the GPU.
78
+ output_size, is_invalid = torch.stack((out_ptr[-1], invalid.to(torch.int64))).cpu().tolist()
79
+ if is_invalid:
80
+ raise ValueError("invalid ptr, out_ptr, or starts for the input point clouds")
81
+
82
+ out = torch.empty(output_size, dtype=torch.int64, device=x.device)
83
+ if batch == 0 or output_size == 0:
84
+ return out
85
+ min_d2 = torch.empty(x.shape[0], dtype=torch.float32, device=x.device)
86
+ _shader_library().fps_flat(
87
+ x,
88
+ ptr,
89
+ out_ptr,
90
+ starts,
91
+ min_d2,
92
+ out,
93
+ threads=(_THREADS, batch),
94
+ group_size=(_THREADS, 1),
95
+ )
96
+ return out
@@ -0,0 +1,126 @@
1
+ """Metal searches over sorted, flattened point-cloud batches.
2
+
3
+ These private kernels produce rectangular global-index arrays with ``-1``
4
+ padding. The public flat API removes padding to build torch_cluster-style
5
+ ``[query, reference]`` edge indices.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from functools import cache
11
+ from importlib import resources
12
+ import math
13
+ import struct
14
+
15
+ import torch
16
+ from torch import Tensor
17
+
18
+ from ._ball_query_mps import _checked_radius_and_k, _require_compile_shader
19
+ from .ops import _check_simd_width
20
+
21
+
22
+ _MAX_K = 256
23
+ _QUERIES_PER_GROUP = 8
24
+ _UINT_MAX = (1 << 32) - 1
25
+
26
+
27
+ @cache
28
+ def _library():
29
+ _require_compile_shader()
30
+ source = resources.files(__package__).joinpath("kernels", "flat_search.metal").read_text()
31
+ return torch.mps.compile_shader(source)
32
+
33
+
34
+ def _check_inputs(x: Tensor, y: Tensor, ptr_x: Tensor, ptr_y: Tensor) -> tuple[int, int]:
35
+ if x.ndim != 2 or y.ndim != 2 or x.shape[1] != 3 or y.shape[1] != 3:
36
+ raise ValueError("x and y must have shapes (N, 3) and (M, 3)")
37
+ if x.device.type != "mps" or y.device != x.device:
38
+ raise ValueError("x and y must be on the same MPS device")
39
+ if ptr_x.ndim != 1 or ptr_y.ndim != 1 or ptr_x.shape != ptr_y.shape or not ptr_x.numel():
40
+ raise ValueError("ptr_x and ptr_y must have the same nonempty 1-D shape")
41
+ for name, ptr, count in (("ptr_x", ptr_x, x.shape[0]), ("ptr_y", ptr_y, y.shape[0])):
42
+ if ptr.dtype != torch.int64 or ptr.device != x.device:
43
+ raise ValueError(f"{name} must be an int64 tensor on {x.device}")
44
+ if int(ptr[0].item()) != 0 or int(ptr[-1].item()) != count:
45
+ raise ValueError(f"{name} must start at 0 and end at {count}")
46
+ lengths = ptr[1:] - ptr[:-1]
47
+ if bool(torch.any((lengths < 0) | (lengths > _UINT_MAX)).item()):
48
+ raise ValueError(f"{name} must be nondecreasing with per-batch size <= 2**32-1")
49
+ if y.shape[0] > _UINT_MAX:
50
+ raise ValueError("flat query count exceeds the Metal 1-D dispatch limit")
51
+ return int(ptr_x.numel() - 1), int(y.shape[0])
52
+
53
+
54
+ def _torch_cluster_radius_sq(r: float) -> float:
55
+ """The threshold torch_cluster compares against: fl32(r * r), r * r in double.
56
+
57
+ torch_cluster takes ``r`` as a double and passes ``r * r`` to a float32
58
+ kernel. PyTorch3D-style Ball Query instead squares the float32-rounded
59
+ radius; the two can differ by one float32 ULP (for example at r = 0.21).
60
+ """
61
+ try:
62
+ return struct.unpack("f", struct.pack("f", r * r))[0]
63
+ except OverflowError:
64
+ return math.inf
65
+
66
+
67
+ def knn_indices(x: Tensor, y: Tensor, ptr_x: Tensor, ptr_y: Tensor, k: int) -> Tensor:
68
+ """Return global x indices, sorted by (squared distance, x index)."""
69
+ batch_count, query_count = _check_inputs(x, y, ptr_x, ptr_y)
70
+ if x.dtype != torch.float32 or y.dtype != torch.float32:
71
+ raise TypeError("flat Metal kNN requires float32 x and y")
72
+ if not isinstance(k, int) or isinstance(k, bool) or not 0 <= k <= _MAX_K:
73
+ raise ValueError(f"k must be an integer in [0, {_MAX_K}] for the flat Metal kernel")
74
+ if query_count * k > (1 << 63) - 1:
75
+ raise ValueError("kNN output size exceeds int64 indexing")
76
+ out = torch.empty((query_count, k), dtype=torch.int64, device=x.device)
77
+ if query_count == 0 or k == 0:
78
+ return out
79
+ _check_simd_width()
80
+ group = 32 * _QUERIES_PER_GROUP
81
+ groups = (query_count + _QUERIES_PER_GROUP - 1) // _QUERIES_PER_GROUP
82
+ _library().flat_knn_indices(
83
+ y.contiguous(), x.contiguous(), ptr_y.contiguous(), ptr_x.contiguous(), out,
84
+ query_count, batch_count, k,
85
+ threads=[groups * group, 1, 1], group_size=[group, 1, 1],
86
+ )
87
+ return out
88
+
89
+
90
+ def radius_indices(
91
+ x: Tensor, y: Tensor, ptr_x: Tensor, ptr_y: Tensor, r: float,
92
+ max_num_neighbors: int, ignore_same_index: bool = False,
93
+ ) -> Tensor:
94
+ """Return first valid global x indices per query, with ``-1`` padding.
95
+
96
+ When requested, equal global x/y index numbers are excluded before the
97
+ neighbor cap is applied, matching pyg-lib's radius operator.
98
+ """
99
+ batch_count, query_count = _check_inputs(x, y, ptr_x, ptr_y)
100
+ if x.dtype not in (torch.float32, torch.float16) or y.dtype != x.dtype:
101
+ raise TypeError("flat Metal radius requires matching float32 or float16 x and y")
102
+ radius_f32, _ = _checked_radius_and_k(r, max_num_neighbors)
103
+ radius_sq = _torch_cluster_radius_sq(float(r))
104
+ if query_count * max_num_neighbors > (1 << 63) - 1:
105
+ raise ValueError("radius output size exceeds int64 indexing")
106
+ out = torch.empty((query_count, max_num_neighbors), dtype=torch.int64, device=x.device)
107
+ if query_count == 0 or max_num_neighbors == 0:
108
+ return out
109
+ if radius_f32 == 0.0:
110
+ out.fill_(-1)
111
+ return out
112
+ _check_simd_width()
113
+ kernel = (
114
+ _library().flat_radius_indices_f32
115
+ if x.dtype == torch.float32
116
+ else _library().flat_radius_indices_f16
117
+ )
118
+ group = 32 * _QUERIES_PER_GROUP
119
+ groups = (query_count + _QUERIES_PER_GROUP - 1) // _QUERIES_PER_GROUP
120
+ kernel(
121
+ y.contiguous(), x.contiguous(), ptr_y.contiguous(), ptr_x.contiguous(), out,
122
+ query_count, batch_count, max_num_neighbors, radius_sq, radius_f32,
123
+ int(ignore_same_index),
124
+ threads=[groups * group, 1, 1], group_size=[group, 1, 1],
125
+ )
126
+ return out
mps_pointops/compat.py ADDED
@@ -0,0 +1,166 @@
1
+ """Optional stand-ins for ``pointnet2_ops``, ``knn_cuda`` and ``torch_cluster``.
2
+
3
+ Call ``install()`` before importing code that uses them::
4
+
5
+ import mps_pointops.compat
6
+ mps_pointops.compat.install()
7
+
8
+ from pointnet2_ops import pointnet2_utils # served by mps_pointops
9
+ from knn_cuda import KNN
10
+ from torch_cluster import fps, knn, radius
11
+
12
+ ``install()`` registers ``pointnet2_ops``, ``pointnet2_ops.pointnet2_utils``,
13
+ ``knn_cuda`` and ``torch_cluster`` in ``sys.modules``. A name that is already
14
+ importable, such as the real package on a CUDA machine, is left alone unless
15
+ ``force=True``. Call it before importing the packages that use these names.
16
+
17
+ Covered: ``furthest_point_sample``, ``gather_operation``,
18
+ ``grouping_operation`` and ``ball_query`` from ``pointnet2_utils``, and
19
+ ``KNN`` from ``knn_cuda``. Differences from the CUDA versions:
20
+
21
+ - Near ties can resolve differently. The kernels round each squared
22
+ distance without FMA and break ties by the smaller index; the CUDA
23
+ kernels use their own reduction order.
24
+ - ``ball_query`` uses the Metal kernel for MPS inputs and pads in the
25
+ ``pointnet2_ops`` convention.
26
+ - The ``torch_cluster`` shim exposes ``fps``, ``knn``, ``radius`` and their
27
+ same-set graph wrappers for flat three-dimensional point coordinates.
28
+ Other names needed for PyG 2.7 package import raise ``NotImplementedError``
29
+ when called. It does not register PyG's separate ``torch.ops.pyg`` operators.
30
+ """
31
+
32
+ from __future__ import annotations
33
+
34
+ import importlib.util
35
+ from importlib.machinery import ModuleSpec
36
+ import sys
37
+ import types
38
+
39
+ import torch
40
+ from torch import Tensor
41
+
42
+ from . import flat, ops
43
+
44
+
45
+ # ---------------------------------------------------------------- pointnet2_ops
46
+
47
+
48
+ def furthest_point_sample(xyz: Tensor, npoint: int) -> Tensor:
49
+ """Like ``pointnet2_ops.pointnet2_utils.furthest_point_sample``.
50
+
51
+ (B, N, 3) float32 -> (B, npoint) int32. Starts at index 0 and never picks
52
+ points with x^2 + y^2 + z^2 <= 1e-3, as pointnet2_ops does.
53
+ """
54
+ return ops.furthest_point_sample(xyz, npoint, skip_near_origin=True).int()
55
+
56
+
57
+ def gather_operation(features: Tensor, idx: Tensor) -> Tensor:
58
+ """Like ``pointnet2_utils.gather_operation``: (B, C, N), (B, npoint) -> (B, C, npoint)."""
59
+ B, C, _ = features.shape
60
+ return features.gather(2, idx.long().unsqueeze(1).expand(B, C, -1))
61
+
62
+
63
+ def grouping_operation(features: Tensor, idx: Tensor) -> Tensor:
64
+ """Like ``pointnet2_utils.grouping_operation``: (B, C, N), (B, npoint, nsample) -> (B, C, npoint, nsample)."""
65
+ B, C, _ = features.shape
66
+ _, npoint, nsample = idx.shape
67
+ flat = idx.long().reshape(B, 1, npoint * nsample).expand(B, C, -1)
68
+ return features.gather(2, flat).reshape(B, C, npoint, nsample)
69
+
70
+
71
+ def ball_query(radius: float, nsample: int, xyz: Tensor, new_xyz: Tensor) -> Tensor:
72
+ """Like ``pointnet2_utils.ball_query``: (B, N, 3), (B, npoint, 3) -> (B, npoint, nsample) int32.
73
+
74
+ Takes the first ``nsample`` points in input order with squared distance
75
+ ``< radius**2``. Empty slots repeat the first neighbor, and a query with no
76
+ neighbor gets all zeros, as in pointnet2_ops.
77
+ """
78
+ _, idx = ops.ball_query(new_xyz, xyz, radius, nsample)
79
+ first = idx[..., :1].clamp(min=0)
80
+ return torch.where(idx >= 0, idx, first).int()
81
+
82
+
83
+ # ---------------------------------------------------------------- knn_cuda
84
+
85
+
86
+ class KNN(torch.nn.Module):
87
+ """Like ``knn_cuda.KNN``.
88
+
89
+ With ``transpose_mode=True``, ``ref`` is (B, N, 3) and ``query`` is
90
+ (B, M, 3), and ``dist`` and ``idx`` are (B, M, k). Otherwise ``ref`` is
91
+ (B, 3, N), ``query`` is (B, 3, M), and the outputs are (B, k, M). ``dist``
92
+ is Euclidean. No gradients flow, as in knn_cuda.
93
+ """
94
+
95
+ def __init__(self, k: int, transpose_mode: bool = False):
96
+ super().__init__()
97
+ self.k = k
98
+ self.transpose_mode = transpose_mode
99
+
100
+ def forward(self, ref: Tensor, query: Tensor) -> tuple[Tensor, Tensor]:
101
+ if not self.transpose_mode:
102
+ ref, query = ref.transpose(1, 2), query.transpose(1, 2)
103
+ with torch.no_grad():
104
+ dist, idx = ops.knn(query.float(), ref.float(), self.k)
105
+ if not self.transpose_mode:
106
+ dist, idx = dist.transpose(1, 2), idx.transpose(1, 2)
107
+ return dist, idx
108
+
109
+
110
+ # ---------------------------------------------------------------- install
111
+
112
+
113
+ def _module(name: str, **attrs) -> types.ModuleType:
114
+ module = types.ModuleType(name)
115
+ module.__dict__.update(attrs)
116
+ module.__spec__ = ModuleSpec(name, loader=None, is_package="__path__" in attrs)
117
+ return module
118
+
119
+
120
+ def _unsupported_torch_cluster(name: str):
121
+ def unsupported(*args, **kwargs):
122
+ raise NotImplementedError(
123
+ f"torch_cluster.{name} is outside the mps_pointops point-cloud subset"
124
+ )
125
+
126
+ unsupported.__name__ = name
127
+ return unsupported
128
+
129
+
130
+ def install(force: bool = False) -> list[str]:
131
+ """Register the stand-in modules. Returns the names that were installed."""
132
+ pointnet2_utils = _module(
133
+ "pointnet2_ops.pointnet2_utils",
134
+ furthest_point_sample=furthest_point_sample,
135
+ gather_operation=gather_operation,
136
+ grouping_operation=grouping_operation,
137
+ ball_query=ball_query,
138
+ )
139
+ packages = {
140
+ "pointnet2_ops": {
141
+ "pointnet2_ops": _module("pointnet2_ops", pointnet2_utils=pointnet2_utils, __path__=[]),
142
+ "pointnet2_ops.pointnet2_utils": pointnet2_utils,
143
+ },
144
+ "knn_cuda": {"knn_cuda": _module("knn_cuda", KNN=KNN)},
145
+ "torch_cluster": {
146
+ "torch_cluster": _module(
147
+ "torch_cluster",
148
+ fps=flat.fps,
149
+ knn=flat.knn,
150
+ radius=flat.radius,
151
+ knn_graph=flat.knn_graph,
152
+ radius_graph=flat.radius_graph,
153
+ grid_cluster=_unsupported_torch_cluster("grid_cluster"),
154
+ graclus_cluster=_unsupported_torch_cluster("graclus_cluster"),
155
+ random_walk=_unsupported_torch_cluster("random_walk"),
156
+ nearest=_unsupported_torch_cluster("nearest"),
157
+ )
158
+ },
159
+ }
160
+ installed = []
161
+ for top, modules in packages.items():
162
+ if not force and (top in sys.modules or importlib.util.find_spec(top) is not None):
163
+ continue
164
+ sys.modules.update(modules)
165
+ installed.extend(modules)
166
+ return installed