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.
- sublevel_flood-0.1/.gitignore +5 -0
- sublevel_flood-0.1/PKG-INFO +68 -0
- sublevel_flood-0.1/README.md +50 -0
- sublevel_flood-0.1/images/noisy_sphere_mma.png +0 -0
- sublevel_flood-0.1/pyproject.toml +27 -0
- sublevel_flood-0.1/slflood/__init__.py +6 -0
- sublevel_flood-0.1/slflood/data/__init__.py +3 -0
- sublevel_flood-0.1/slflood/data/generate.py +41 -0
- sublevel_flood-0.1/slflood/landmarks.py +63 -0
- sublevel_flood-0.1/slflood/meb.py +88 -0
- sublevel_flood-0.1/slflood/sklearn/__init__.py +3 -0
- sublevel_flood-0.1/slflood/sklearn/transformer.py +104 -0
- sublevel_flood-0.1/slflood/slflood.py +258 -0
- sublevel_flood-0.1/slflood/triton_mask_topk.py +110 -0
- sublevel_flood-0.1/slflood/triton_maxcummin.py +141 -0
- sublevel_flood-0.1/slflood/utils.py +94 -0
- sublevel_flood-0.1/tests/test_sklearn.py +17 -0
- sublevel_flood-0.1/tests/test_slflood.py +49 -0
|
@@ -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
|
+
[](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
|
+

|
|
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
|
+
[](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
|
+

|
|
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
|
+
```
|
|
Binary file
|
|
@@ -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,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,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')
|