sublevel-flood 0.1__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.
@@ -0,0 +1,5 @@
1
+ .gitignore
2
+ .idea/
3
+ dist/
4
+ __pycache__/
5
+ check.py
@@ -0,0 +1,68 @@
1
+ Metadata-Version: 2.5
2
+ Name: sublevel-flood
3
+ Version: 0.1
4
+ Summary: A package that computes the sublevel Flood bifiltration, a scalable approach for using 2-parameter persistent homology with large-scale, low-dimensional point sets.
5
+ Project-URL: Homepage, https://github.com/MClemot/SLFlood
6
+ Requires-Python: >=3.10
7
+ Requires-Dist: fpsample
8
+ Requires-Dist: gudhi
9
+ Requires-Dist: multipers>=2.6.0
10
+ Requires-Dist: numpy
11
+ Requires-Dist: torch
12
+ Provides-Extra: sklearn
13
+ Requires-Dist: scikit-learn; extra == 'sklearn'
14
+ Provides-Extra: test
15
+ Requires-Dist: pytest; extra == 'test'
16
+ Requires-Dist: scikit-learn; extra == 'test'
17
+ Description-Content-Type: text/markdown
18
+
19
+ # The sublevel Flood bifiltration: towards scalable 2-parameter persistent homology
20
+
21
+ [![PyPI - Version](https://img.shields.io/pypi/v/sublevel-flood?logo=python)](https://pypi.org/project/sublevel-flood/)
22
+
23
+ This repository contains the code for the manuscript *Sublevel Flood bifiltration: scalable 2-parameter persistent homology*.
24
+ The aim is to provide a way to compute 2-parameter persistent homology of large scale point clouds (up to $10^6$).
25
+ To do so, it extends the [*Flood filtration*](https://proceedings.neurips.cc/paper_files/paper/2025/hash/aba03e6f25decb32bda9c5bf81c58305-Abstract-Conference.html) to the 2-parameter context.
26
+
27
+ The Python package constructs a ```SimplexTreeMulti``` from the [```multipers```](https://github.com/DavidLapous/multipers) library.
28
+
29
+ ## Installation
30
+ `sublevel-flood` can be installed via PyPI.
31
+
32
+ ```pip install sublevel-flood```
33
+
34
+ ## Usage
35
+
36
+ The following example computes the sublevel Flood bifiltration (with 100 landmarks) of a point set in $\mathbb{R}^3$ distributed as a sphere with ambient noise, using CUDA. Then, it computes and plots its multiparameter module approximation with `multipers`.
37
+
38
+ ```python
39
+ import slflood
40
+ from slflood.data import noisy_sphere_data
41
+
42
+ import matplotlib.pyplot as plt
43
+ import multipers
44
+
45
+ n_pts = 10000
46
+ n_lms = 100
47
+ dim = 3
48
+
49
+ pts, lms_idx, fun = noisy_sphere_data(n_pts, n_lms, dim,
50
+ bandwidth=0.2)
51
+
52
+ slf = slflood.slflood_bifiltration(pts.cuda(), lms_idx.cuda(), fun.cuda())
53
+ mma = multipers.module_approximation(slf)
54
+ mma.plot(box=[[0, 0], [1, 1]])
55
+ plt.show()
56
+ ```
57
+
58
+ ![](https://raw.githubusercontent.com/MClemot/SLFlood/main/images/noisy_sphere_mma.png)
59
+
60
+ ## Citation
61
+ ```bibtex
62
+ @inproceedings{clemot2026sublevel,
63
+ title={The sublevel {Flood} bifiltration: towards scalable 2-parameter persistent homology},
64
+ author={Cl{\'e}mot, Matt{\'e}o and Digne, Julie and Tierny, Julien},
65
+ year={2026},
66
+ booktitle={NeurIPS},
67
+ }
68
+ ```
@@ -0,0 +1,50 @@
1
+ # The sublevel Flood bifiltration: towards scalable 2-parameter persistent homology
2
+
3
+ [![PyPI - Version](https://img.shields.io/pypi/v/sublevel-flood?logo=python)](https://pypi.org/project/sublevel-flood/)
4
+
5
+ This repository contains the code for the manuscript *Sublevel Flood bifiltration: scalable 2-parameter persistent homology*.
6
+ The aim is to provide a way to compute 2-parameter persistent homology of large scale point clouds (up to $10^6$).
7
+ To do so, it extends the [*Flood filtration*](https://proceedings.neurips.cc/paper_files/paper/2025/hash/aba03e6f25decb32bda9c5bf81c58305-Abstract-Conference.html) to the 2-parameter context.
8
+
9
+ The Python package constructs a ```SimplexTreeMulti``` from the [```multipers```](https://github.com/DavidLapous/multipers) library.
10
+
11
+ ## Installation
12
+ `sublevel-flood` can be installed via PyPI.
13
+
14
+ ```pip install sublevel-flood```
15
+
16
+ ## Usage
17
+
18
+ The following example computes the sublevel Flood bifiltration (with 100 landmarks) of a point set in $\mathbb{R}^3$ distributed as a sphere with ambient noise, using CUDA. Then, it computes and plots its multiparameter module approximation with `multipers`.
19
+
20
+ ```python
21
+ import slflood
22
+ from slflood.data import noisy_sphere_data
23
+
24
+ import matplotlib.pyplot as plt
25
+ import multipers
26
+
27
+ n_pts = 10000
28
+ n_lms = 100
29
+ dim = 3
30
+
31
+ pts, lms_idx, fun = noisy_sphere_data(n_pts, n_lms, dim,
32
+ bandwidth=0.2)
33
+
34
+ slf = slflood.slflood_bifiltration(pts.cuda(), lms_idx.cuda(), fun.cuda())
35
+ mma = multipers.module_approximation(slf)
36
+ mma.plot(box=[[0, 0], [1, 1]])
37
+ plt.show()
38
+ ```
39
+
40
+ ![](https://raw.githubusercontent.com/MClemot/SLFlood/main/images/noisy_sphere_mma.png)
41
+
42
+ ## Citation
43
+ ```bibtex
44
+ @inproceedings{clemot2026sublevel,
45
+ title={The sublevel {Flood} bifiltration: towards scalable 2-parameter persistent homology},
46
+ author={Cl{\'e}mot, Matt{\'e}o and Digne, Julie and Tierny, Julien},
47
+ year={2026},
48
+ booktitle={NeurIPS},
49
+ }
50
+ ```
@@ -0,0 +1,27 @@
1
+ [build-system]
2
+ requires = ["hatchling"]
3
+ build-backend = "hatchling.build"
4
+
5
+ [project]
6
+ name = "sublevel-flood"
7
+ version = "0.1"
8
+ description = "A package that computes the sublevel Flood bifiltration, a scalable approach for using 2-parameter persistent homology with large-scale, low-dimensional point sets."
9
+ readme = "README.md"
10
+ requires-python = ">=3.10"
11
+ dependencies = [
12
+ "fpsample",
13
+ "gudhi",
14
+ "multipers >= 2.6.0",
15
+ "numpy",
16
+ "torch",
17
+ ]
18
+
19
+ [project.urls]
20
+ Homepage = "https://github.com/MClemot/SLFlood"
21
+
22
+ [project.optional-dependencies]
23
+ sklearn = ["scikit-learn"]
24
+ test = ["pytest", "scikit-learn"]
25
+
26
+ [tool.hatch.build.targets.wheel]
27
+ packages = ["slflood"]
@@ -0,0 +1,6 @@
1
+ from .landmarks import generate_landmarks_FPS
2
+ from .slflood import slflood_bifiltration
3
+
4
+ __version__ = "0.1"
5
+
6
+ __all__ = ["slflood_bifiltration", "generate_landmarks_FPS"]
@@ -0,0 +1,3 @@
1
+ from .generate import noisy_sphere_data
2
+
3
+ __all__ = ["noisy_sphere_data"]
@@ -0,0 +1,41 @@
1
+ from multipers.filtrations.density import KDE
2
+ import torch
3
+ from ..landmarks import generate_landmarks_FPS
4
+
5
+ device = 'cpu'
6
+
7
+ def generate_ambient_noise(n_pts, dim, radius=1., seed=None):
8
+ if seed is not None:
9
+ torch.manual_seed(seed)
10
+ directions = torch.normal(0, 1, size=(n_pts, dim), device=device)
11
+ directions /= torch.linalg.norm(directions, dim=-1, keepdim=True)
12
+ radii = torch.rand((n_pts, 1), device=device) ** (1. / dim)
13
+ return directions * radii * radius
14
+
15
+ def generate_sphere(n, dim, device, seed=None):
16
+ if seed is not None:
17
+ torch.manual_seed(seed)
18
+ X = torch.normal(0, 1, size=(n,dim), device=device)
19
+ X = X / torch.norm(X, dim=-1, keepdim=True)
20
+ X += torch.normal(0, 0.01, size=(n,dim), device=device)
21
+ return X
22
+
23
+ def generate_noisy_sphere(n_pts, dim, seed=None):
24
+ pts = generate_sphere(n_pts // 2, dim, device, seed=seed)
25
+ pts = torch.cat([pts, generate_ambient_noise(n_pts // 2, dim, seed=seed) * 1.5])
26
+ return pts
27
+
28
+ def density(pts, bandwidth):
29
+ f = -KDE(bandwidth=bandwidth, return_log=False).fit(pts).score_samples(pts)
30
+ f -= f.min()
31
+ f /= f.max()
32
+ return f
33
+
34
+ def sub_pts(pts, n_lms, start_idx=0):
35
+ return generate_landmarks_FPS(pts, n_lms, start_idx=start_idx, return_idxs=True)
36
+
37
+ def noisy_sphere_data(n_pts, n_lms, dim, bandwidth=.1, seed=None):
38
+ pts = generate_noisy_sphere(n_pts, dim=dim, seed=seed)
39
+ fun = density(pts, bandwidth)
40
+ lms_idx = sub_pts(pts, n_lms)
41
+ return pts, lms_idx, fun
@@ -0,0 +1,63 @@
1
+ import fpsample
2
+ import numpy as np
3
+ import torch
4
+ from typing import Union
5
+
6
+ @torch.no_grad()
7
+ def generate_landmarks_FPS(
8
+ points: torch.Tensor,
9
+ n_lms: int,
10
+ fps_h: Union[None, int] = None,
11
+ start_idx: Union[int, None] = None,
12
+ return_idxs: bool = False,
13
+ ) -> torch.Tensor:
14
+ """
15
+ Selects landmarks using Farthest-Point Sampling (bucket FPS).
16
+
17
+ This method implements a variant of Farthest-Point Sampling from
18
+ [here](https://dl.acm.org/doi/abs/10.1109/TCAD.2023.3274922).
19
+
20
+ Adapted from https://github.com/plus-rkwitt/flooder.
21
+
22
+ Args:
23
+ points (torch.Tensor):
24
+ A (P, d) tensor representing a point cloud. The tensor may reside on any
25
+ device (CPU or GPU) and be of any floating-point dtype.
26
+ n_lms (int):
27
+ The number of landmarks to sample (must be <= P and > 0).
28
+ fps_h (Union[None, int], optional):
29
+ h parameter (depth of kdtree) that is used for farthest point sampling to
30
+ select the landmarks. If None, then h is selected based on the size of the
31
+ point cloud. Defaults to None.
32
+ start_idx (int | None, optional):
33
+ If provided, the sampling starts from this index in the point cloud. If not,
34
+ the start index will be randomly picked from the point cloud.
35
+ return_idxs (bool, optional):
36
+ If true, returns the index of the selected landmarks. Defaults to False.
37
+
38
+ Returns:
39
+ torch.Tensor:
40
+ A (n_l, d) tensor containing a subset of the input `points`, representing the sampled landmarks,
41
+ or a (n_l) tensor of integers containing the indices of the sampled landmarks.
42
+ """
43
+ if n_lms <= 0:
44
+ raise RuntimeError(
45
+ f"Number of landmarks ({n_lms}) must be positive"
46
+ )
47
+ n_pts = len(points)
48
+ n_lms = min(n_lms, n_pts)
49
+ if fps_h is None:
50
+ if n_pts > 200_000:
51
+ fps_h = 9
52
+ elif n_pts > 80_000:
53
+ fps_h = 7
54
+ else:
55
+ fps_h = 5
56
+
57
+ index_set = torch.tensor(
58
+ fpsample.bucket_fps_kdline_sampling(points.cpu(), n_lms, h=fps_h, start_idx=start_idx).astype(np.int64),
59
+ device=points.device,
60
+ )
61
+ if return_idxs:
62
+ return index_set
63
+ return points[index_set]
@@ -0,0 +1,88 @@
1
+ import numpy as np
2
+ from itertools import combinations
3
+
4
+ def get_circumspheres(S):
5
+ """
6
+ Computes the circumspheres of a batch of set of points
7
+
8
+ Parameters
9
+ ----------
10
+ S : (B, M, N) ndarray, where 1 <= M <= N + 1
11
+ The input points
12
+
13
+ Returns
14
+ -------
15
+ C, r : ((2) ndarray, float)
16
+ The center and the squared radius of the circumsphere
17
+ """
18
+
19
+ U = S[:,1:] - S[:,0,None]
20
+ B = np.linalg.norm(U, axis=2, keepdims=True)
21
+ U /= B
22
+ B /= 2
23
+ G = np.matmul(U, U.transpose(0,2,1))
24
+ try:
25
+ x = np.linalg.solve(G, B)[:,:,0]
26
+ except np.linalg.LinAlgError:
27
+ G += np.random.normal(size=G.shape, loc=0, scale=1e-10)
28
+ x = np.linalg.solve(G, B)[:,:,0]
29
+ C = np.einsum('bi,bij->bj', x, U)
30
+ r = np.linalg.norm(C, axis=1)
31
+ C += S[:,0]
32
+ return C, r
33
+
34
+
35
+ def get_middles(p):
36
+ """
37
+ Computes the middles of a batch of pair of points
38
+ """
39
+ center = 0.5 * (p[:, 0] + p[:, 1])
40
+ radius = 0.5 * np.linalg.norm(p[:, 1] - p[:, 0], axis=-1)
41
+ return center, radius
42
+
43
+
44
+ def contains_all(center, radius, pts, tol=1e-7):
45
+ """
46
+ Check if ball contains all 4 tetrahedron vertices.
47
+ """
48
+ d = np.linalg.norm(pts - center[:, None, :], axis=-1)
49
+ return np.all(d <= radius[:, None] + tol, axis=-1)
50
+
51
+
52
+ def minimum_enclosing_balls(X:np.ndarray):
53
+ """
54
+ Compute the minimum enclosing balls of a batch of sets of points
55
+
56
+ Parameters
57
+ ----------
58
+ X : (B, N, D)
59
+ input
60
+
61
+ Returns
62
+ -------
63
+ centers,radius : ((2) ndarray, float)
64
+ centers and radii of the minimum enclosing balls
65
+ """
66
+ B = X.shape[0]
67
+ N = X.shape[1]
68
+ D = X.shape[2]
69
+
70
+ best_r = np.full(B, np.inf)
71
+ best_c = np.zeros((B, D))
72
+
73
+ for k in range(N,1,-1):
74
+ for idx in combinations(range(N), k):
75
+ subset = X[:, idx]
76
+
77
+ if k == 2:
78
+ c, r = get_middles(subset)
79
+ else:
80
+ c, r = get_circumspheres(subset)
81
+
82
+ valid = contains_all(c, r, X)
83
+ update = valid & (r < best_r)
84
+
85
+ best_r[update] = r[update]
86
+ best_c[update] = c[update]
87
+
88
+ return best_c, best_r
@@ -0,0 +1,3 @@
1
+ from .transformer import SublevelFloodBifiltration
2
+
3
+ __all__ = ["SublevelFloodBifiltration"]
@@ -0,0 +1,104 @@
1
+ import slflood
2
+
3
+ from multipers.filtrations.density import KDE, DTM
4
+ import numpy as np
5
+ try:
6
+ from sklearn.base import BaseEstimator, TransformerMixin
7
+ except ImportError as e:
8
+ raise ImportError(
9
+ "slflood.sklearn requires scikit-learn, which can be installed with `pip install sublevel-flood[sklearn]`"
10
+ ) from e
11
+ import torch
12
+
13
+ class SublevelFloodBifiltration(BaseEstimator, TransformerMixin):
14
+ """
15
+ Scikit-learn transformer computing the sublevel Flood bifiltration of point clouds.
16
+ For each input point cloud, a co-density function is estimated via kernel density or distance-to-measure.
17
+
18
+ Parameters
19
+ ----------
20
+ n_lms : int
21
+ Number of landmarks to sample (via farthest point sampling).
22
+
23
+ kde_bandwidth : float, optional
24
+ Bandwidth of the Gaussian kernel used to compute the density functions.
25
+ Mutually exclusive with ``dtm_mass``.
26
+
27
+ dtm_mass : float, optional
28
+ Mass parameter in (0, 1] used to compute the Distance-To-Measure (DTM) functions.
29
+ Mutually exclusive with ``kde_bandwidth``.
30
+
31
+ log_density : bool, default=False
32
+ If True and ``kde_bandwidth`` is set, uses the log-density instead of the density.
33
+
34
+ normalize : bool, default=True
35
+ Whether to normalize the density functions.
36
+
37
+ points_per_edge : int, default=20
38
+ Resolution of the sampling grid used on each simplex.
39
+
40
+ exact : bool, default=False
41
+ Whether to compute the exact (unmasked) or approximate (masked, faster) sublevel Flood bifiltration.
42
+
43
+ device : {'auto', 'cpu', 'cuda'}, default='auto'
44
+ Device used for computation. When 'auto', detects whether CUDA is available.
45
+
46
+ use_triton : bool, default=True
47
+ Whether to use Triton kernels when available.
48
+
49
+ n_jobs : int, default=-1
50
+ Currently unused.
51
+ """
52
+
53
+ def __init__(self, n_lms, kde_bandwidth=None, dtm_mass=None, log_density=False, normalize=True, points_per_edge=20, exact=False, device='auto', use_triton=True, n_jobs=-1):
54
+ self.n_lms = n_lms
55
+ self.kde_bandwidth = kde_bandwidth
56
+ self.dtm_mass = dtm_mass
57
+ self.log_density = log_density
58
+ self.normalize = normalize
59
+ self.points_per_edge = points_per_edge
60
+ self.exact = exact
61
+ self.device = device
62
+ self.use_triton = use_triton
63
+ self.n_jobs = n_jobs
64
+
65
+ def __get_device(self):
66
+ if self.device not in ['cpu', 'cuda', 'auto']:
67
+ raise ValueError(f"device must be 'cpu', 'cuda' or 'auto', got {self.device!r}")
68
+ if self.device == 'auto':
69
+ return 'cuda' if torch.cuda.is_available() else 'cpu'
70
+ if self.device == 'cuda' and not torch.cuda.is_available():
71
+ raise RuntimeError("device='cuda' was requested but CUDA is not available")
72
+ return self.device
73
+
74
+ def fit(self, X, y=None):
75
+ return self
76
+
77
+ def __density(self, X):
78
+ functions = []
79
+ for pts in X:
80
+ if self.kde_bandwidth is not None:
81
+ function = -KDE(bandwidth=self.kde_bandwidth, kernel="gaussian", return_log=self.log_density).fit(pts).score_samples(pts)
82
+ elif self.dtm_mass is not None:
83
+ function = DTM(masses=[self.dtm_mass]).fit(pts).score_samples(pts)[0]
84
+ else:
85
+ raise ValueError("exactly one parameter among kde_bandwidth and dtm_mass must be given")
86
+ if self.normalize:
87
+ function -= function.min()
88
+ function /= function.max()
89
+ functions.append(function)
90
+ return functions
91
+
92
+ def __transform(self, x, f, device):
93
+ x = torch.as_tensor(x, device=device)
94
+ f = torch.as_tensor(f, device=device)
95
+ return slflood.slflood_bifiltration(x, self.n_lms, f,
96
+ points_per_edge=self.points_per_edge,
97
+ exact=self.exact,
98
+ use_triton=self.use_triton)
99
+
100
+ def transform(self, X):
101
+ device = self.__get_device()
102
+ F = self.__density(X)
103
+
104
+ return [self.__transform(x, f, device) for (x,f) in zip(X, F)]
@@ -0,0 +1,258 @@
1
+ """Implementation of sublevel flood bifiltration."""
2
+
3
+ from typing import Union
4
+
5
+ import gudhi
6
+ import multipers as mp
7
+ import numbers
8
+ import numpy as np
9
+ import torch
10
+
11
+ from .landmarks import generate_landmarks_FPS
12
+ from .meb import minimum_enclosing_balls
13
+ from .utils import generate_grid, mask_from_adjacency, pad_bigrades, reset, elapsed
14
+
15
+ try:
16
+ from .triton_maxcummin import maxcummin
17
+ from .triton_mask_topk import compact_mask_indices
18
+ HAS_TRITON = True
19
+ except ImportError:
20
+ HAS_TRITON = False
21
+
22
+ @torch.no_grad()
23
+ def near_validity(points: torch.Tensor,
24
+ simplex_centers: torch.Tensor,
25
+ simplex_radii: torch.Tensor) -> torch.Tensor:
26
+ centers_to_points = torch.cdist(simplex_centers, points)
27
+ return centers_to_points < simplex_radii.unsqueeze(1)
28
+
29
+
30
+ @torch.no_grad()
31
+ def vertex_bigrades(points: torch.Tensor,
32
+ landmarks: torch.Tensor,
33
+ function: torch.Tensor,
34
+ cells: torch.Tensor,
35
+ mask: torch.Tensor):
36
+ landmarks_to_points = torch.cdist(landmarks, points)
37
+ cummin = torch.cummin(landmarks_to_points, dim=1)[0]
38
+ critical = torch.ones_like(cummin, dtype=torch.int8)
39
+ critical[:, 1:] = cummin[:, 1:] != cummin[:, :-1]
40
+
41
+ critical_count = critical.sum(1)
42
+ critical_count_max = critical_count.max().item()
43
+ critical_indices = torch.argsort(critical, dim=1, descending=True, stable=True)[:,:critical_count_max]
44
+ bigrades_diameters = torch.gather(cummin, dim=1, index=critical_indices)
45
+ bigrades_functions = function[critical_indices]
46
+ bigrades = torch.stack((bigrades_diameters, bigrades_functions), dim=-1).cpu().to(torch.float32)
47
+
48
+ scc0 = [bigrades[k,:critical_count[k]] for k in range(len(landmarks))]
49
+
50
+ cell_mask = critical[cells].amax(dim=1)
51
+ mask = torch.max(mask, cell_mask)
52
+
53
+ return scc0, mask
54
+
55
+
56
+ @torch.no_grad()
57
+ def slflood_bifiltration(
58
+ points: torch.Tensor,
59
+ landmarks: Union[int, torch.Tensor],
60
+ function: torch.Tensor,
61
+ max_dimension: int = -1,
62
+ points_per_edge: int = 20,
63
+ near_radius: float = np.sqrt(2),
64
+ exact: bool = False,
65
+ batch_size: int = 32,
66
+ use_triton: bool = False,
67
+ fps_h: Union[None, int] = None,
68
+ start_idx: Union[int, None] = 0,
69
+ verbose = False,
70
+ ):
71
+ """
72
+ Constructs the sublevel Flood bifiltration from a set of points with scalar values and a set of landmarks.
73
+
74
+ Args:
75
+ points (torch.Tensor):
76
+ A (N, d) tensor containing the input point set.
77
+ landmarks (Union[int, torch.Tensor]):
78
+ Either an integer indicating the number of landmarks to randomly sample from `points` using FPS,
79
+ or a tensor of integers of shape (N_l) giving the indices of the entries of `points` to take as landmarks,
80
+ or a tensor of floats of shape (N_l, d) specifying an explicit set of landmarks.
81
+ function (torch.Tensor):
82
+ A (N) tensor containing the scalar values on the points.
83
+ max_dimension (int, optional):
84
+ The top dimension of the simplices to construct.
85
+ Defaults to -1 resulting in the dimension of the ambient space.
86
+ points_per_edge (int, optional):
87
+ Specifies resolution on simplices used for computing filtration values. Defaults to 20.
88
+ near_radius (float, optional):
89
+ Specifies the dilation factor of the minimum enclosing balls to obtain near points. Defaults to sqrt(2).
90
+ exact (bool, optional):
91
+ Whether to compute the exact (unmasked) or approximate (masked) sublevel Flood bifiltration.
92
+ batch_size (int, optional):
93
+ Size of simplex batches. Defaults to 32.
94
+ use_triton (bool, optional):
95
+ Whether to use Triton kernels if available.
96
+ fps_h (Union[None, int], optional):
97
+ h parameter (depth of kdtree) that is used for farthest point sampling to
98
+ select the landmarks. If None, then h is selected based on the size of the
99
+ point cloud. Defaults to None.
100
+ start_idx (int | None, optional):
101
+ If provided, FPS starts from this index in the point cloud. If not,
102
+ the start index will be randomly picked from the point cloud. Defaults to 0.
103
+
104
+ Returns:
105
+ mp.SimplexTreeMulti
106
+ Returns a k-critical mp.SimplexTreeMulti containing the sublevel Flood bifiltration.
107
+ """
108
+
109
+ reset()
110
+
111
+ if max_dimension == -1:
112
+ max_dimension = points.shape[1]
113
+
114
+ if isinstance(landmarks, numbers.Integral):
115
+ if landmarks >= points.shape[0]:
116
+ landmarks = points
117
+ else:
118
+ landmarks = generate_landmarks_FPS(points, landmarks, fps_h, start_idx=start_idx)
119
+ elif landmarks.device != points.device:
120
+ raise RuntimeError(f"landmarks.device ({landmarks.device}) != points.device ({points.device})")
121
+
122
+ device = points.device
123
+ use_triton = use_triton and HAS_TRITON and device.type == 'cuda'
124
+ if use_triton:
125
+ points = points.float()
126
+ if torch.is_floating_point(landmarks):
127
+ landmarks = landmarks.float()
128
+ dtype = points.dtype
129
+
130
+ arg = torch.argsort(function)
131
+ ranks = torch.empty_like(arg)
132
+ ranks[arg] = torch.arange(points.size(0), device=device)
133
+ points = points[arg].contiguous()
134
+ function = function[arg].contiguous()
135
+
136
+ if not torch.is_floating_point(landmarks):
137
+ landmarks = points[ranks[landmarks]]
138
+
139
+ if verbose:
140
+ elapsed("SORT")
141
+
142
+ stree = gudhi.delaunay_complex.DelaunayComplex(landmarks.cpu()).create_simplex_tree(filtration=None)
143
+ cells_list = []
144
+ cells_idx = dict()
145
+ simplices_list = [[] for _ in range(max_dimension-1)]
146
+ simplices_cocells = [[] for _ in range(max_dimension-1)]
147
+ for simplex, _ in stree.get_simplices():
148
+ dim = len(simplex)-1
149
+ if dim == max_dimension:
150
+ cells_idx[tuple(simplex)] = len(cells_list)
151
+ cells_list.append(simplex)
152
+ elif 0 < dim < max_dimension:
153
+ simplices_list[dim-1].append(simplex)
154
+ cocells = []
155
+ for cell,_ in stree.get_cofaces(simplex, max_dimension - dim):
156
+ cocells.append(cells_idx[tuple(cell)])
157
+ simplices_cocells[dim-1].append(cocells)
158
+
159
+ ret_stree = mp.SimplexTreeMulti(num_parameters=2, kcritical=True)
160
+
161
+ cells = torch.tensor(cells_list, device=device)
162
+ simplices_by_dim = [torch.tensor(l, device=device) for l in simplices_list]
163
+
164
+ if verbose:
165
+ elapsed("DELAUNAY")
166
+
167
+ simplex_vertices = landmarks[cells]
168
+ simplex_vertices_cpu = simplex_vertices.cpu().numpy().astype(np.float64, copy=False)
169
+ Lc, Lr = minimum_enclosing_balls(simplex_vertices_cpu)
170
+ simplex_radii = near_radius * torch.tensor(Lr, device=device, dtype=dtype)
171
+ simplex_centers = torch.tensor(Lc, device=device, dtype=dtype)
172
+
173
+ if verbose:
174
+ elapsed("BALLS")
175
+
176
+ mask = near_validity(points, simplex_centers, simplex_radii)
177
+ scc0, mask = vertex_bigrades(points, landmarks, function, cells, mask)
178
+ if exact:
179
+ mask = torch.ones_like(mask)
180
+ if verbose:
181
+ elapsed("MASK")
182
+
183
+ ret_stree.insert_batch(np.arange(len(scc0))[None,:], pad_bigrades(scc0))
184
+ if verbose:
185
+ elapsed("ASSIGN_0")
186
+
187
+ submasks = [mask_from_adjacency(cocells, mask) for cocells in simplices_cocells]
188
+ if verbose:
189
+ elapsed("SUBMASKS")
190
+
191
+ simplices_by_dim.append(cells)
192
+ submasks.append(mask)
193
+
194
+ for simplices, mask in zip(simplices_by_dim, submasks):
195
+ num_simplices, dim = simplices.shape
196
+
197
+ weights = generate_grid(points_per_edge, dim-1, device, dtype)
198
+ simplex_vertices = landmarks[simplices]
199
+ points_on_simplex = weights.unsqueeze(0) @ simplex_vertices
200
+
201
+ if not exact:
202
+ mask_counts = mask.sum(1, keepdim=False, dtype=torch.int32)
203
+ indices = torch.sort(mask_counts)[1]
204
+ mask = mask[indices].contiguous()
205
+ simplices = simplices[indices].contiguous()
206
+ points_on_simplex = points_on_simplex[indices].contiguous()
207
+
208
+ mask_counts_t = mask_counts[indices]
209
+ mask_counts = mask_counts[indices].cpu().numpy()
210
+
211
+ bigrades = []
212
+ for k in range(-(-num_simplices // batch_size)):
213
+ s1, s2 = k*batch_size, min((k+1)*batch_size, num_simplices)
214
+
215
+ # b_simplices = simplices[s1:s2]
216
+ # b_simplices_vertices = landmarks[b_simplices]
217
+ # b_sample = weights.unsqueeze(0) @ b_simplices_vertices
218
+ b_sample = points_on_simplex[s1:s2]
219
+
220
+ if not exact:
221
+ b_mask = mask[s1:s2]
222
+ b_mask_count_max = int(mask_counts[s2 - 1])
223
+ if use_triton:
224
+ b_mask_indices = compact_mask_indices(b_mask, mask_counts_t[s1:s2], b_mask_count_max)
225
+ else:
226
+ b_mask_indices = torch.argsort(b_mask, dim=1, descending=True, stable=True)[:, :b_mask_count_max]
227
+ b_mask_points = points[b_mask_indices]
228
+ b_mask_function = function[b_mask_indices]
229
+ else:
230
+ b_mask_points = points.repeat(s2-s1, 1, 1)
231
+ b_mask_function = function.repeat(s2-s1, 1)
232
+
233
+ if use_triton:
234
+ max_diameters = maxcummin(b_mask_points, b_sample)
235
+ else:
236
+ inter = torch.cdist(b_mask_points, b_sample)
237
+ cummin = inter.cummin(dim=1)[0]
238
+ max_diameters = torch.amax(cummin, dim=2)
239
+
240
+ critical = torch.ones_like(max_diameters, dtype=torch.int8)
241
+ critical[:, 1:] = max_diameters[:, 1:] != max_diameters[:, :-1]
242
+ critical_count = critical.sum(dim=1)
243
+ critical_indices = torch.argsort(critical, dim=1, descending=True, stable=True)
244
+ bigrade_diameters = torch.gather(max_diameters, dim=1, index=critical_indices)
245
+ bigrade_functions = torch.gather(b_mask_function, dim=1, index=critical_indices)
246
+
247
+ b_bigrades = torch.stack((bigrade_diameters, bigrade_functions), dim=-1).cpu().to(torch.float32)
248
+ critical_count = critical_count.cpu().numpy()
249
+
250
+ for i in range(len(b_bigrades)):
251
+ bigrades.append(b_bigrades[i, :critical_count[i]].numpy())
252
+
253
+ ret_stree.insert_batch(simplices.cpu().numpy().T, pad_bigrades(bigrades))
254
+
255
+ if verbose:
256
+ elapsed("LOOP")
257
+
258
+ return ret_stree
@@ -0,0 +1,110 @@
1
+ """
2
+ Triton replacement for:
3
+
4
+ b_mask_indices = torch.argsort(b_mask, dim=1, descending=True, stable=True)[:, :K]
5
+ or
6
+ b_mask_indices = torch.topk(b_mask, k=K, dim=1, sorted=False)[1]
7
+
8
+ Algorithm: classic prefix-sum-based stream compaction, done as a single
9
+ sequential pass per row over blocks of the N dimension, carrying a running
10
+ (true_count_so_far, implicit false_count_so_far) across blocks:
11
+
12
+ For row-relative position j, let C[j] = number of True entries strictly
13
+ before j (exclusive prefix sum of the mask up to j).
14
+ - if mask[j] is True: its destination column is C[j]
15
+ - if mask[j] is False: its destination column is true_count + (j - C[j])
16
+ where true_count is the row's TOTAL true count (known ahead of time -
17
+ reuse the mask_counts you already compute before the batch loop,
18
+ don't recompute here).
19
+
20
+ Only destinations < K are ever written (True entries always land below
21
+ K since true_count <= K by construction of K = max over the batch), so
22
+ most False-entry writes for well-covered rows never happen at all.
23
+ """
24
+
25
+ import torch
26
+ import triton
27
+ import triton.language as tl
28
+
29
+
30
+ @triton.autotune(
31
+ configs=[
32
+ triton.Config({"BLOCK_N": 512}, num_warps=4),
33
+ triton.Config({"BLOCK_N": 1024}, num_warps=4),
34
+ triton.Config({"BLOCK_N": 1024}, num_warps=8),
35
+ triton.Config({"BLOCK_N": 2048}, num_warps=8),
36
+ ],
37
+ key=["N"],
38
+ )
39
+ @triton.jit(do_not_specialize=["K"])
40
+ def _compact_kernel(
41
+ mask_ptr, counts_ptr, out_ptr,
42
+ stride_mb, stride_mn,
43
+ stride_cb,
44
+ stride_ob, stride_ok,
45
+ N, K,
46
+ BLOCK_N: tl.constexpr,
47
+ ):
48
+ row = tl.program_id(0)
49
+ true_count = tl.load(counts_ptr + row * stride_cb)
50
+
51
+ carry_true = 0
52
+ n_blocks = tl.cdiv(N, BLOCK_N)
53
+
54
+ for blk in range(0, n_blocks):
55
+ n_offsets = blk * BLOCK_N + tl.arange(0, BLOCK_N)
56
+ n_valid = n_offsets < N
57
+
58
+ m = tl.load(
59
+ mask_ptr + row * stride_mb + n_offsets * stride_mn,
60
+ mask=n_valid, other=0,
61
+ ).to(tl.int32)
62
+
63
+ incl = tl.cumsum(m, axis=0) # inclusive prefix sum within this block
64
+ excl = incl - m # exclusive prefix sum within this block
65
+
66
+ global_true_rank = carry_true + excl # C[j], global
67
+ global_false_rank = n_offsets - global_true_rank # j - C[j], global
68
+
69
+ dest = tl.where(m == 1, global_true_rank, true_count + global_false_rank)
70
+ store_mask = n_valid & (dest < K)
71
+
72
+ tl.store(out_ptr + row * stride_ob + dest * stride_ok, n_offsets, mask=store_mask)
73
+
74
+ block_true_count = tl.sum(tl.where(n_valid, m, 0), axis=0)
75
+ carry_true += block_true_count
76
+
77
+
78
+ def compact_mask_indices(mask: torch.Tensor, mask_counts: torch.Tensor, k: int) -> torch.Tensor:
79
+ """
80
+ Drop-in replacement for:
81
+ b_mask_indices = torch.argsort(mask, dim=1, descending=True, stable=True)[:, :K]
82
+
83
+ mask: (rows, N) bool, CUDA
84
+ mask_counts: (rows,) int - true count per row (mask.sum(1)).
85
+ Pass the values you've already computed for batch sorting
86
+ rather than recomputing them here - they're the same
87
+ numbers.
88
+ k: output width (b_mask_count_max for this batch)
89
+
90
+ Returns (rows, K) int64 tensor of indices, matching stable
91
+ descending-argsort semantics exactly: True positions in ascending
92
+ original order first, then enough False positions (also ascending) to
93
+ fill out to K columns.
94
+ """
95
+ assert mask.is_cuda
96
+ rows, N = mask.shape
97
+
98
+ mask_i32 = mask.to(torch.int32).contiguous()
99
+ counts_i32 = mask_counts.to(torch.int32).contiguous()
100
+ out = torch.empty((rows, k), device=mask.device, dtype=torch.int32)
101
+
102
+ grid = (rows,)
103
+ _compact_kernel[grid](
104
+ mask_i32, counts_i32, out,
105
+ mask_i32.stride(0), mask_i32.stride(1),
106
+ counts_i32.stride(0),
107
+ out.stride(0), out.stride(1),
108
+ N, k,
109
+ )
110
+ return out.to(torch.int64)
@@ -0,0 +1,141 @@
1
+ """
2
+ Triton replacement for:
3
+
4
+ inter = torch.cdist(b_mask_points, b_sample) # (B, M, G) - materialized
5
+ cummin = inter.cummin(dim=1)[0] # (B, M, G) - materialized
6
+ max_diameters = torch.amax(cummin, dim=2) # (B, M) - only this survives
7
+
8
+ Algorithm (online scan, same structural pattern as flash-attention's
9
+ online softmax):
10
+
11
+ For each batch row b and grid point g, maintain a running minimum
12
+ distance as witness points are added in order (m = 0, 1, 2, ...).
13
+ After adding point m, the "covering radius" at that prefix length is
14
+ max_g(running_min[g]). Only that (B, M) result is ever needed downstream
15
+ so the (B, M, G) intermediate never needs to exist in memory.
16
+
17
+ Grid points are split into blocks of BLOCK_G for parallelism (so a run
18
+ isn't limited to B independent programs, which would badly underuse a
19
+ GPU with far more SMs than typical batch sizes). Each (batch, g-block)
20
+ program computes its own partial max per m-step; a final cheap
21
+ torch.amax over the small g-block axis combines them.
22
+ """
23
+
24
+ import torch
25
+ import triton
26
+ import triton.language as tl
27
+
28
+
29
+ @triton.autotune(
30
+ configs=[
31
+ triton.Config({"BLOCK_G": 256}, num_warps=4),
32
+ triton.Config({"BLOCK_G": 512}, num_warps=4),
33
+ triton.Config({"BLOCK_G": 512}, num_warps=8),
34
+ triton.Config({"BLOCK_G": 1024}, num_warps=8),
35
+ triton.Config({"BLOCK_G": 2048}, num_warps=8),
36
+ ],
37
+ key=["G", "D"],
38
+ )
39
+ @triton.jit(do_not_specialize=["M", "G"])
40
+ def _flood_scan_kernel(
41
+ points_ptr, sample_ptr, out_ptr,
42
+ stride_pb, stride_pm, stride_pd,
43
+ stride_sb, stride_sg, stride_sd,
44
+ stride_ob, stride_om, stride_ogb,
45
+ M, G,
46
+ D: tl.constexpr,
47
+ D_PADDED: tl.constexpr,
48
+ BLOCK_G: tl.constexpr,
49
+ ):
50
+ pid_b = tl.program_id(0)
51
+ pid_g = tl.program_id(1)
52
+
53
+ g_offsets = pid_g * BLOCK_G + tl.arange(0, BLOCK_G)
54
+ g_mask = g_offsets < G
55
+
56
+ d_offsets = tl.arange(0, D_PADDED)
57
+ d_mask = d_offsets < D
58
+
59
+ # Load this program's slice of sample points once: (BLOCK_G, D_PADDED)
60
+ sample_ptrs = (
61
+ sample_ptr
62
+ + pid_b * stride_sb
63
+ + g_offsets[:, None] * stride_sg
64
+ + d_offsets[None, :] * stride_sd
65
+ )
66
+ load_mask = g_mask[:, None] & d_mask[None, :]
67
+ sample_block = tl.load(sample_ptrs, mask=load_mask, other=0.0)
68
+
69
+ running_min = tl.full([BLOCK_G], float("inf"), dtype=tl.float32)
70
+ out_base = out_ptr + pid_b * stride_ob + pid_g * stride_ogb
71
+
72
+ # Sequential scan over M - inherent to the cumulative-min structure.
73
+ # M is a runtime value (not tl.constexpr), so this compiles to an
74
+ # actual device-side loop rather than being unrolled at compile time
75
+ # (same pattern as the K-reduction loop in Triton's matmul tutorial).
76
+ for i in range(0, M):
77
+ point_ptrs = (
78
+ points_ptr
79
+ + pid_b * stride_pb
80
+ + i * stride_pm
81
+ + d_offsets * stride_pd
82
+ )
83
+ point_i = tl.load(point_ptrs, mask=d_mask, other=0.0) # (D_PADDED,)
84
+
85
+ diff = sample_block - point_i[None, :]
86
+ sqdist = tl.sum(diff * diff, axis=1) # (BLOCK_G,) - padded dims contribute 0
87
+ dist = tl.sqrt(sqdist)
88
+
89
+ running_min = tl.minimum(running_min, dist)
90
+
91
+ masked_min = tl.where(g_mask, running_min, float("-inf"))
92
+ local_max = tl.max(masked_min, axis=0)
93
+
94
+ tl.store(out_base + i * stride_om, local_max)
95
+
96
+
97
+ def maxcummin(b_mask_points: torch.Tensor, b_sample: torch.Tensor) -> torch.Tensor:
98
+ """
99
+ Drop-in replacement for:
100
+ inter = torch.cdist(b_mask_points, b_sample)
101
+ cummin = inter.cummin(dim=1)[0]
102
+ max_diameters = torch.amax(cummin, dim=2)
103
+
104
+ b_mask_points: (B, M, D) float32, CUDA
105
+ b_sample: (B, G, D) float32, CUDA
106
+ returns: (B, M) float32
107
+ """
108
+ assert b_mask_points.is_cuda and b_sample.is_cuda, "kernel is CUDA-only"
109
+ assert b_mask_points.dtype == torch.float32 and b_sample.dtype == torch.float32
110
+ B, M, D = b_mask_points.shape
111
+ B2, G, D2 = b_sample.shape
112
+ assert B == B2 and D == D2, "batch/dim mismatch between points and sample"
113
+
114
+ b_mask_points = b_mask_points.contiguous()
115
+ b_sample = b_sample.contiguous()
116
+
117
+ D_PADDED = triton.next_power_of_2(D)
118
+
119
+ def grid(meta):
120
+ return (B, triton.cdiv(G, meta["BLOCK_G"]))
121
+
122
+ # Sized for the smallest BLOCK_G in the autotune list so the buffer
123
+ # shape doesn't depend on which config gets picked; unused trailing
124
+ # columns stay -inf and are ignored by the final amax.
125
+ smallest_block_g = 256
126
+ max_g_blocks = triton.cdiv(G, smallest_block_g)
127
+ partial = torch.full(
128
+ (B, M, max_g_blocks), float("-inf"),
129
+ device=b_mask_points.device, dtype=torch.float32,
130
+ )
131
+
132
+ _flood_scan_kernel[grid](
133
+ b_mask_points, b_sample, partial,
134
+ b_mask_points.stride(0), b_mask_points.stride(1), b_mask_points.stride(2),
135
+ b_sample.stride(0), b_sample.stride(1), b_sample.stride(2),
136
+ partial.stride(0), partial.stride(1), partial.stride(2),
137
+ M, G,
138
+ D=D, D_PADDED=D_PADDED,
139
+ )
140
+
141
+ return partial.amax(dim=2)
@@ -0,0 +1,94 @@
1
+ import numpy as np
2
+ import torch
3
+
4
+
5
+ @torch.no_grad()
6
+ def generate_grid(n: int, dim: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
7
+ """Generates a grid of points on the unit simplex based on the number of points per edge.
8
+ Adapted from https://github.com/plus-rkwitt/flooder.
9
+
10
+ Args:
11
+ n (int):
12
+ Number of points per edge.
13
+ dim (int):
14
+ Dimension of the simplex.
15
+ device (torch.device):
16
+ Device to create the tensor on.
17
+ dtype (torch.dtype):
18
+ Dtype of the tensor.
19
+
20
+ Returns:
21
+ grid (torch.Tensor):
22
+ Tensor of shape (C, dim + 1), containing the grid points (coordinate weights).
23
+ """
24
+
25
+ combs = torch.combinations(torch.arange(n + dim - 1, device=torch.device('cpu')), r=dim)
26
+ padded = torch.cat(
27
+ [
28
+ torch.full((combs.shape[0], 1), -1, device=torch.device('cpu')),
29
+ combs,
30
+ torch.full((combs.shape[0], 1), n + dim - 1, device=torch.device('cpu')),
31
+ ],
32
+ dim=1,
33
+ ) # shape [C, dim + 2]
34
+ grid = torch.diff(padded, dim=1) - 1 # shape [C, dim + 1]
35
+ grid_float = torch.empty_like(grid, dtype=dtype)
36
+ torch.divide(grid, n - 1 , out=grid_float)
37
+ return grid_float.to(device)
38
+
39
+
40
+ @torch.no_grad()
41
+ def mask_from_adjacency(adjacency, mask):
42
+ """
43
+ adjacency: list of length n_k, adjacency[i] = list of cell indices adjacent to k-simplex i
44
+ mask: torch.Tensor of shape (n_cells, *rest), binary
45
+
46
+ Returns: torch.Tensor of shape (n_k, *rest)
47
+ """
48
+ rest_shape = mask.shape[1:]
49
+ n_k = len(adjacency)
50
+
51
+ lengths = torch.tensor([len(a) for a in adjacency], dtype=torch.long, device=mask.device)
52
+ dest_idx = torch.repeat_interleave(torch.arange(n_k, device=mask.device), lengths)
53
+ src_idx = torch.tensor([c for adj in adjacency for c in adj], dtype=torch.long, device=mask.device)
54
+
55
+ values = mask[src_idx] # (n_edges, *rest)
56
+
57
+ idx_scatter = dest_idx.view(-1, *([1] * len(rest_shape))).expand_as(values)
58
+ out = torch.zeros((n_k,) + rest_shape, dtype=mask.dtype, device=mask.device)
59
+ out.scatter_reduce_(0, idx_scatter, values, reduce="amax", include_self=True)
60
+
61
+ return out
62
+
63
+
64
+ def pad_bigrades(bigrades_list):
65
+ """
66
+ bigrades_list: list of N arrays, each of shape (k_i, 2), k_i possibly different.
67
+ Returns: array of shape (N, max_k, 2), padded with (inf, inf).
68
+ """
69
+ N = len(bigrades_list)
70
+ max_k = max(arr.shape[0] for arr in bigrades_list)
71
+
72
+ padded = np.full((N, max_k, 2), np.inf, dtype=np.float64)
73
+ for i, arr in enumerate(bigrades_list):
74
+ k = arr.shape[0]
75
+ padded[i, :k, :] = arr
76
+
77
+ return padded
78
+
79
+
80
+ import time
81
+
82
+ t = time.time()
83
+
84
+ def elapsed(s=None):
85
+ global t
86
+ if s is not None:
87
+ print("[timing]", s, int(1000*(time.time() - t)))
88
+ else:
89
+ print("[timing]", int(1000*(time.time() - t)))
90
+ t = time.time()
91
+
92
+ def reset():
93
+ global t
94
+ t = time.time()
@@ -0,0 +1,17 @@
1
+ import pytest
2
+ pytest.importorskip("sklearn")
3
+
4
+ from slflood.sklearn import SublevelFloodBifiltration
5
+
6
+ from multipers.ml.mma import FilteredComplex2MMA
7
+ import torch
8
+ from sklearn.pipeline import Pipeline
9
+
10
+ def test_sklearn():
11
+ X = torch.rand((10,10000,3))
12
+
13
+ pipeline = Pipeline([("slf", SublevelFloodBifiltration(100, 0.2, use_triton=True)),
14
+ # ("mma", FilteredComplex2MMA(n_jobs=-1))
15
+ ])
16
+
17
+ Y = pipeline.fit_transform(X)
@@ -0,0 +1,49 @@
1
+ import slflood
2
+ import slflood.data
3
+
4
+ import multipers as mp
5
+ import pytest
6
+ import torch
7
+
8
+ # import matplotlib.pyplot as plt
9
+
10
+ def run(n_pts, n_lms, dim, device):
11
+ pts, lms_idx, function = slflood.data.noisy_sphere_data(n_pts, n_lms, dim, 0.2, seed=0)
12
+ if device in ['cuda', 'triton']:
13
+ if not torch.cuda.is_available():
14
+ pytest.skip('CUDA not available')
15
+ else:
16
+ pts, lms_idx, function = pts.cuda(), lms_idx.cuda(), function.cuda()
17
+ ff = slflood.slflood_bifiltration(pts, lms_idx, function,
18
+ use_triton=device=='triton')
19
+ mma_ff = mp.module_approximation(ff)
20
+ # mma_ff.plot(box=[[0, 0], [1, 1]])
21
+ # plt.show()
22
+
23
+
24
+ n_pts = 5000
25
+ n_lms = 100
26
+
27
+
28
+ def test_dim2_cpu():
29
+ run(n_pts, n_lms, 2, 'cpu')
30
+
31
+
32
+ def test_dim2_cuda():
33
+ run(n_pts, n_lms, 2, 'cuda')
34
+
35
+
36
+ def test_dim2_triton():
37
+ run(n_pts, n_lms, 2, 'triton')
38
+
39
+
40
+ def test_dim3_cpu():
41
+ run(n_pts, n_lms, 3, 'cpu')
42
+
43
+
44
+ def test_dim3_cuda():
45
+ run(n_pts, n_lms, 3, 'cuda')
46
+
47
+
48
+ def test_dim3_triton():
49
+ run(n_pts, n_lms, 3, 'triton')