Skip to content

Commit 5b35078

Browse files
committed
update: bootstrap loop matching
1 parent d417dea commit 5b35078

5 files changed

Lines changed: 170 additions & 29 deletions

File tree

‎pyproject.toml‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,11 @@ dependencies = [
5353
"umap-learn>=0.5.7",
5454
]
5555

56+
[dependency-groups]
57+
dev = [
58+
"ty>=0.0.1a34",
59+
]
60+
5661
[tool.uv]
5762
package = true
5863

@@ -106,8 +111,3 @@ executionEnvironments = [
106111

107112
venvPath = "/home/stanfish/Git/scloop"
108113
venv = ".venv"
109-
110-
[dependency-groups]
111-
dev = [
112-
"ty>=0.0.1a34",
113-
]

‎src/scloop/computing/homology.py‎

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,10 +9,15 @@
99
from sklearn.neighbors import radius_neighbors_graph
1010

1111
from ..data.metadata import ScloopMeta
12-
from ..data.ripser_lib import get_boundary_matrix, ripser # type: ignore[import-not-found]
12+
from ..data.ripser_lib import ( # type: ignore[import-not-found]
13+
get_boundary_matrix,
14+
ripser,
15+
)
1316
from ..data.types import Diameter_t, IndexListDistMatrix
1417
from ..data.utils import encode_triangles_and_edges
15-
from ..utils.linear_algebra_gf2.m4ri_lib import solve_multiple_gf2 # type: ignore[import-not-found]
18+
from ..utils.linear_algebra_gf2.m4ri_lib import (
19+
solve_multiple_gf2, # type: ignore[import-not-found]
20+
)
1621

1722
if TYPE_CHECKING:
1823
from ..data.containers import BoundaryMatrixD1

‎src/scloop/data/__init__.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,2 +1,6 @@
11
from .containers import HomologyData
2-
from .ripser_lib import RipserResults, get_boundary_matrix, ripser # type: ignore[import-not-found]
2+
from .ripser_lib import ( # type: ignore[import-not-found]
3+
RipserResults,
4+
get_boundary_matrix,
5+
ripser,
6+
)

‎src/scloop/data/containers.py‎

Lines changed: 111 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,13 @@
11
# Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan)
22
from abc import ABC, abstractmethod
3+
from concurrent.futures import ThreadPoolExecutor, as_completed
34

45
import numpy as np
56
from anndata import AnnData
67
from pydantic import BaseModel, Field, ValidationInfo, field_validator
78
from pydantic.dataclasses import dataclass
89
from scipy.sparse import csr_matrix, triu
10+
from scipy.spatial.distance import directed_hausdorff
911

1012
from ..computing.homology import (
1113
compute_boundary_matrix_data,
@@ -22,6 +24,7 @@
2224
edge_ids_to_rows,
2325
extract_edges_from_coo,
2426
loop_vertices_to_edge_ids,
27+
nearest_neighbor_per_row,
2528
)
2629

2730

@@ -289,30 +292,74 @@ def _compute_loop_representatives(
289292
else:
290293
bootstrap_data.loop_representatives[idx_bootstrap][loop_idx] = loops # type: ignore[attr-defined]
291294

295+
def _assess_bootstrap_geometric_equivalence(
296+
self,
297+
adata: AnnData,
298+
source_class_idx: int,
299+
target_class_idx: int,
300+
idx_bootstrap: int = 0,
301+
) -> tuple[int, int, float]:
302+
assert self.loop_representatives is not None
303+
assert self.bootstrap_data is not None
304+
assert self.meta.preprocess.embedding_method is not None
305+
306+
if idx_bootstrap >= len(self.bootstrap_data.loop_representatives):
307+
return (source_class_idx, target_class_idx, np.nan)
308+
if source_class_idx >= len(self.loop_representatives):
309+
return (source_class_idx, target_class_idx, np.nan)
310+
311+
boot_loops_all = self.bootstrap_data.loop_representatives[idx_bootstrap]
312+
if target_class_idx >= len(boot_loops_all):
313+
return (source_class_idx, target_class_idx, np.nan)
314+
315+
source_loops = self.loop_representatives[source_class_idx]
316+
target_loops = boot_loops_all[target_class_idx]
317+
318+
if len(source_loops) == 0 or len(target_loops) == 0:
319+
return (source_class_idx, target_class_idx, np.nan)
320+
321+
emb = adata.obsm[f"X_{self.meta.preprocess.embedding_method}"]
322+
distances = []
323+
for source_loop in source_loops:
324+
for target_loop in target_loops:
325+
source_coords = emb[source_loop]
326+
target_coords = emb[target_loop]
327+
try:
328+
dist = max(
329+
directed_hausdorff(source_coords, target_coords)[0],
330+
directed_hausdorff(target_coords, source_coords)[0],
331+
)
332+
distances.append(dist)
333+
except (ValueError, IndexError):
334+
distances.append(np.nan)
335+
336+
mean_distance = np.nanmean(distances) if distances else np.nan
337+
return (source_class_idx, target_class_idx, mean_distance)
338+
292339
def _assess_bootstrap_homology_equivalence(
293340
self,
294341
source_class_idx: int,
295342
target_class_idx: int | None = None,
296343
idx_bootstrap: int = 0,
297344
n_pairs_check: int = 10,
298-
) -> bool:
345+
) -> tuple[int, int, bool]:
299346
assert self.boundary_matrix_d1 is not None
300347
assert self.loop_representatives is not None
301348
assert self.bootstrap_data is not None
302349
if target_class_idx is None:
303350
target_class_idx = source_class_idx
304-
if idx_bootstrap >= len(self.bootstrap_data.loop_representatives): # type: ignore[attr-defined]
305-
return False
351+
if idx_bootstrap >= len(self.bootstrap_data.loop_representatives):
352+
return (source_class_idx, target_class_idx, False)
306353
if source_class_idx >= len(self.loop_representatives):
307-
return False
308-
boot_loops_all = self.bootstrap_data.loop_representatives[idx_bootstrap] # type: ignore[attr-defined]
354+
return (source_class_idx, target_class_idx, False)
355+
boot_loops_all = self.bootstrap_data.loop_representatives[idx_bootstrap]
309356
if target_class_idx >= len(boot_loops_all):
310-
return False
357+
return (source_class_idx, target_class_idx, False)
311358

312359
source_loops = self.loop_representatives[source_class_idx]
313360
target_loops = boot_loops_all[target_class_idx]
314361
if len(source_loops) == 0 or len(target_loops) == 0:
315-
return False
362+
return (source_class_idx, target_class_idx, False)
316363

317364
mask_a = self._loops_to_edge_mask(source_loops)
318365
mask_b = self._loops_to_edge_mask(target_loops)
@@ -323,7 +370,7 @@ def _assess_bootstrap_homology_equivalence(
323370
loop_mask_b=mask_b,
324371
n_pairs_check=n_pairs_check,
325372
)
326-
return any(r == 0 for r in results)
373+
return (source_class_idx, target_class_idx, any(r == 0 for r in results))
327374

328375
def _bootstrap(
329376
self,
@@ -340,6 +387,8 @@ def _bootstrap(
340387
loop_lower_t_pct: float = 5,
341388
loop_upper_t_pct: float = 95,
342389
n_pairs_check_equivalence: int = 4,
390+
n_max_workers: int = 4,
391+
k_neighbors_check_equivalence: int = 3,
343392
verbose: bool = True,
344393
**nei_kwargs,
345394
) -> None:
@@ -374,24 +423,65 @@ def _bootstrap(
374423
loop_upper_t_pct=loop_upper_t_pct,
375424
)
376425
"""
377-
========= geometric matching =========
378-
- find loop neighbors using frechet
426+
============= geometric matching =============
427+
- find loop neighbors using hausdorff/frechet
379428
- reduce computation load
380-
======================================
429+
==============================================
381430
"""
382-
# n_source_
383-
# for i in range(len(self.loop_representatives)):
384-
# for j in range(len(self.bootstrap_data)):
385-
431+
n_original_loop_classes = len(self.loop_representatives)
432+
n_bootstrap_loop_classes = len(
433+
self.bootstrap_data.loop_representatives[idx_bootstrap]
434+
)
435+
436+
if n_original_loop_classes == 0 or n_bootstrap_loop_classes == 0:
437+
continue
438+
439+
pairwise_result_matrix = np.full(
440+
(n_original_loop_classes, n_bootstrap_loop_classes), np.nan
441+
)
442+
443+
with ThreadPoolExecutor(max_workers=n_max_workers) as executor:
444+
tasks = {}
445+
for i in range(n_original_loop_classes):
446+
for j in range(n_bootstrap_loop_classes):
447+
task = executor.submit(
448+
self._assess_bootstrap_geometric_equivalence,
449+
i,
450+
j,
451+
idx_bootstrap,
452+
)
453+
tasks[task] = (i, j)
454+
455+
for task in as_completed(tasks):
456+
src_idx, tgt_idx, distance = task.result()
457+
pairwise_result_matrix[src_idx, tgt_idx] = distance
458+
459+
neighbor_indices, neighbor_distances = nearest_neighbor_per_row(
460+
pairwise_result_matrix, k_neighbors_check_equivalence
461+
)
386462

387463
"""
388464
========= homological matching =========
389465
- gf2 regression
390466
========================================
391467
"""
392-
self._assess_bootstrap_homology_equivalence(
393-
source_class_idx=0,
394-
target_class_idx=0,
395-
idx_bootstrap=idx_bootstrap,
396-
n_pairs_check=n_pairs_check_equivalence,
397-
)
468+
with ThreadPoolExecutor(max_workers=n_max_workers) as executor:
469+
tasks = {}
470+
for si in range(n_original_loop_classes):
471+
for k in range(k_neighbors_check_equivalence):
472+
tj = neighbor_indices[si, k]
473+
if tj >= 0:
474+
task = executor.submit(
475+
self._assess_bootstrap_homology_equivalence,
476+
si,
477+
tj,
478+
idx_bootstrap,
479+
n_pairs_check_equivalence,
480+
)
481+
tasks[task] = (si, tj, neighbor_distances[si, k])
482+
483+
for task in as_completed(tasks):
484+
si, tj, geom_dist = tasks[task]
485+
_, _, is_equivalent = task.result()
486+
if is_equivalent:
487+
pass

‎src/scloop/data/utils.py‎

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -122,3 +122,45 @@ def edge_ids_to_rows(edge_ids: np.ndarray, edge_row_ids: np.ndarray) -> np.ndarr
122122
rows[count] = row
123123
count += 1
124124
return rows[:count]
125+
126+
127+
@jit(nopython=True)
128+
def nearest_neighbor_per_row(
129+
distance_matrix: np.ndarray, k: int
130+
) -> tuple[np.ndarray, np.ndarray]:
131+
n_rows, n_cols = distance_matrix.shape
132+
neighbor_indices = np.empty((n_rows, k), dtype=np.int64)
133+
neighbor_distances = np.empty((n_rows, k), dtype=np.float64)
134+
135+
for si in range(n_rows):
136+
distances = distance_matrix[si, :]
137+
valid_count = 0
138+
valid_indices = np.empty(n_cols, dtype=np.int64)
139+
valid_distances = np.empty(n_cols, dtype=np.float64)
140+
141+
for j in range(n_cols):
142+
if not np.isnan(distances[j]):
143+
valid_indices[valid_count] = j
144+
valid_distances[valid_count] = distances[j]
145+
valid_count += 1
146+
147+
if valid_count == 0:
148+
neighbor_indices[si, :] = -1
149+
neighbor_distances[si, :] = np.nan
150+
continue
151+
152+
valid_indices = valid_indices[:valid_count]
153+
valid_distances = valid_distances[:valid_count]
154+
155+
n_keep = min(valid_count, k)
156+
sorted_idx = np.argsort(valid_distances)[:n_keep]
157+
158+
for idx in range(n_keep):
159+
neighbor_indices[si, idx] = valid_indices[sorted_idx[idx]]
160+
neighbor_distances[si, idx] = valid_distances[sorted_idx[idx]]
161+
162+
for idx in range(n_keep, k):
163+
neighbor_indices[si, idx] = -1
164+
neighbor_distances[si, idx] = np.nan
165+
166+
return neighbor_indices, neighbor_distances

0 commit comments

Comments
 (0)