Numerical contracts¶
The mathematical selection rule, the stored float32 result, and parity with a CUDA/CPU library are separate claims. Ties and radius boundaries are part of the API contract; binary equivalence across all hardware is not implied.
Selection and output¶
For query q_i, reference x_j, and coordinate dimension D, the real
squared distance is
s(i,j) = sum[d=0..D-1] (q[i,d] - x[j,d])^2.
Dense knn selects the smallest k distances, breaking exact ties by the
smaller reference index, and returns Euclidean distances sqrt(s).
Dense Ball Query retains the first K matching reference indices in their
original input order and returns squared distances. A valid zero distance
is possible; use index >= 0, not distance != 0, to recognize padding.
The dense Ball Query threshold is
r32 = fl32(r)
R2_dense = fl32(r32 * r32)
accept(i,j) iff s32(i,j) < R2_dense.
fl32 is a float32 rounding operation. The flat torch_cluster-style radius
path instead uses R2_flat = fl32(r * r) after the Python double-precision
product. These can differ by one float32 ULP, so cross-API boundary comparisons
must use the intended threshold. The dense kernel also takes a normalized
path for very small radii to reduce subnormal flush errors. Its exact
threshold, FTZ limits, FMA source order, and edge-case tests are in the
full Ball Query numerical contract.
FPS starts at the requested index and repeatedly chooses the point whose
distance to its nearest selected center is greatest. Exact ties choose the
smallest input index. flat.fps computes the sample count from ratio using
the torch_cluster dtype and rounding rule; for deterministic parity set
random_start=False. See the README equations.
Differentiation¶
Neighbor indices and radius membership are discrete and are not differentiated.
With a selected Ball Query pair, upstream distance gradient G, and
delta=q-x, the first-order contributions are 2G*delta to the query and
-2G*delta to the reference point. Multiple references to one point are
accumulated; floating-point summation order may vary. The
Chamfer contract
defines both nearest-neighbor directions, padding masks, weights, point and
batch reduction divisors, L1 subgradients at coincident coordinates, and
normal-loss scope. A tested gradient formula does not make the nearest-index
decision itself differentiable.
Metal math modes¶
| Mode | How to start the process | Interpretation |
|---|---|---|
| Safe | PYTORCH_MPS_FAST_MATH=0 |
Primary numerical contract and regression baseline. |
| Fast | PYTORCH_MPS_FAST_MATH=1 |
Separate observed test mode. Nonfinite handling and boundary bits are not promoted to Safe-mode guarantees. |
Set PYTORCH_ENABLE_MPS_FALLBACK=0 during parity runs so a missing Metal path
cannot be hidden by CPU fallback. Call torch.mps.synchronize() around timed
regions; an API return alone does not mean the GPU finished. Compare Safe and
Fast in different Python processes, since shader compilation settings are
cached. Subnormal flush and compiler contraction can affect the lowest bits
even when the source makes FMA order explicit. The
Metal probe and methods
record the observed behavior and limits.