Skip to content

Commit 84f6621

Browse files
committed
update: bootstrap track store
1 parent c6eb99d commit 84f6621

5 files changed

Lines changed: 90 additions & 13 deletions

File tree

‎.github/workflows/test-import.yml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@ jobs:
2727
run: |
2828
export PATH="$HOME/.cargo/bin:$PATH"
2929
make build
30-
make sync-fresh
30+
make fresh-sync
3131
3232
- name: Test import
3333
run: |

‎Makefile‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@ build: build-m4ri
2121
CPLUS_INCLUDE_PATH=$(PROJECT_ROOT)/src/scloop/data:$(DM_PREFIX) \
2222
uv build
2323

24-
sync-fresh: build-m4ri
24+
fresh-sync: build-m4ri
2525
CPLUS_INCLUDE_PATH=$(PROJECT_ROOT)/src/scloop/data:$(DM_PREFIX)/discrete-frechet-distance uv sync
2626

2727
sync: clean build-m4ri

‎src/scloop/data/analysis_containers.py‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,45 @@
11
# Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan)
22
from __future__ import annotations
33

4+
from typing import Optional
5+
46
import numpy as np
57
from pydantic import ConfigDict, Field
68
from pydantic.dataclasses import dataclass
79

810

11+
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
12+
class LoopMatch:
13+
idx_bootstrap: int
14+
birth_bootstrap: float
15+
death_bootstrap: float
16+
target_class_idx: int
17+
geometric_distance: Optional[float] = None
18+
neighbor_rank: Optional[int] = None
19+
extra: dict = Field(default_factory=dict)
20+
21+
22+
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
23+
class LoopTrack:
24+
source_class_idx: int
25+
birth_root: float
26+
death_root: float
27+
matches: list[LoopMatch] = Field(default_factory=list)
28+
29+
def presence_prob(self, n_bootstraps: int) -> float:
30+
if n_bootstraps == 0:
31+
return 0.0
32+
hit_boots = {m.idx_bootstrap for m in self.matches}
33+
return len(hit_boots) / n_bootstraps
34+
35+
936
@dataclass
1037
class BootstrapAnalysis:
1138
num_bootstraps: int = 0
1239
persistence_diagrams: list[list] = Field(default_factory=list)
1340
cocycles: list[list] = Field(default_factory=list)
1441
loop_representatives: list[list[list[list[int]]]] = Field(default_factory=list)
42+
loop_tracks: dict[int, LoopTrack] = Field(default_factory=dict)
1543

1644

1745
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))

‎src/scloop/data/containers.py‎

Lines changed: 59 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -14,10 +14,15 @@
1414
compute_loop_homological_equivalence,
1515
compute_persistence_diagram_and_cocycles,
1616
)
17-
from .analysis_containers import BootstrapAnalysis, HodgeAnalysis
17+
from .analysis_containers import (
18+
BootstrapAnalysis,
19+
HodgeAnalysis,
20+
LoopMatch,
21+
LoopTrack,
22+
)
1823
from .loop_reconstruction import reconstruct_n_loop_representatives
1924
from .metadata import BootstrapMeta, ScloopMeta
20-
from .types import Diameter_t, Index_t, IndexListDownSample, Size_t
25+
from .types import Diameter_t, Index_t, IndexListDownSample, LoopDistMethod, Size_t
2126
from .utils import (
2227
decode_edges,
2328
decode_triangles,
@@ -239,7 +244,6 @@ def _compute_loop_representatives(
239244
return [], []
240245
top_k = min(top_k, loop_births.size)
241246

242-
# get top k homology classes
243247
indices_top_k = np.argpartition(loop_deaths - loop_births, -top_k)[-top_k:]
244248

245249
dm_upper = triu(pairwise_distance_matrix, k=1).tocoo()
@@ -298,9 +302,11 @@ def _assess_bootstrap_geometric_equivalence(
298302
source_class_idx: int,
299303
target_class_idx: int,
300304
idx_bootstrap: int = 0,
305+
method: LoopDistMethod = "hausdorff",
301306
) -> tuple[int, int, float]:
302307
assert self.loop_representatives is not None
303308
assert self.bootstrap_data is not None
309+
assert self.meta.preprocess is not None
304310
assert self.meta.preprocess.embedding_method is not None
305311

306312
if idx_bootstrap >= len(self.bootstrap_data.loop_representatives):
@@ -334,7 +340,7 @@ def _assess_bootstrap_geometric_equivalence(
334340
distances.append(np.nan)
335341

336342
mean_distance = np.nanmean(distances) if distances else np.nan
337-
return (source_class_idx, target_class_idx, mean_distance)
343+
return (source_class_idx, target_class_idx, float(mean_distance))
338344

339345
def _assess_bootstrap_homology_equivalence(
340346
self,
@@ -372,6 +378,26 @@ def _assess_bootstrap_homology_equivalence(
372378
)
373379
return (source_class_idx, target_class_idx, any(r == 0 for r in results))
374380

381+
def _ensure_loop_tracks(self) -> None:
382+
if self.bootstrap_data is None:
383+
return
384+
if self.bootstrap_data.loop_tracks:
385+
return
386+
if self.persistence_diagram is None or self.loop_representatives is None:
387+
return
388+
loop_births = np.array(self.persistence_diagram[1][0], dtype=np.float32)
389+
loop_deaths = np.array(self.persistence_diagram[1][1], dtype=np.float32)
390+
top_k = len(self.loop_representatives)
391+
if top_k == 0:
392+
return
393+
indices_top_k = np.argpartition(loop_deaths - loop_births, -top_k)[-top_k:]
394+
for track_idx, loop_idx in enumerate(indices_top_k):
395+
birth = float(loop_births[loop_idx])
396+
death = float(loop_deaths[loop_idx])
397+
self.bootstrap_data.loop_tracks[track_idx] = LoopTrack(
398+
source_class_idx=track_idx, birth_root=birth, death_root=death
399+
)
400+
375401
def _bootstrap(
376402
self,
377403
adata: AnnData,
@@ -392,7 +418,7 @@ def _bootstrap(
392418
verbose: bool = True,
393419
**nei_kwargs,
394420
) -> None:
395-
self.bootstrap_data = BootstrapAnalysis(num_bootstraps=n_bootstrap)
421+
self.bootstrap_data = BootstrapAnalysis()
396422
if self.meta.bootstrap is None:
397423
self.meta.bootstrap = BootstrapMeta(indices_resample=[])
398424
else:
@@ -446,6 +472,7 @@ def _bootstrap(
446472
for j in range(n_bootstrap_loop_classes):
447473
task = executor.submit(
448474
self._assess_bootstrap_geometric_equivalence,
475+
adata,
449476
i,
450477
j,
451478
idx_bootstrap,
@@ -459,7 +486,6 @@ def _bootstrap(
459486
neighbor_indices, neighbor_distances = nearest_neighbor_per_row(
460487
pairwise_result_matrix, k_neighbors_check_equivalence
461488
)
462-
463489
"""
464490
========= homological matching =========
465491
- gf2 regression
@@ -478,10 +504,32 @@ def _bootstrap(
478504
idx_bootstrap,
479505
n_pairs_check_equivalence,
480506
)
481-
tasks[task] = (si, tj, neighbor_distances[si, k])
507+
tasks[task] = (si, tj, neighbor_distances[si, k], k)
482508

483509
for task in as_completed(tasks):
484-
si, tj, geom_dist = tasks[task]
485-
_, _, is_equivalent = task.result()
486-
if is_equivalent:
487-
pass
510+
si, tj, geo_dist, neighbor_rank = tasks[task]
511+
_, _, is_homologically_equivalent = task.result()
512+
if self.bootstrap_data is not None and is_homologically_equivalent:
513+
self._ensure_loop_tracks()
514+
track = self.bootstrap_data.loop_tracks.get(si)
515+
birth_boot = float(
516+
self.bootstrap_data.persistence_diagrams[idx_bootstrap][1][
517+
0
518+
][tj]
519+
)
520+
death_boot = float(
521+
self.bootstrap_data.persistence_diagrams[idx_bootstrap][1][
522+
1
523+
][tj]
524+
)
525+
track.matches.append(
526+
LoopMatch(
527+
idx_bootstrap=idx_bootstrap,
528+
birth_bootstrap=birth_boot,
529+
death_bootstrap=death_boot,
530+
target_class_idx=tj,
531+
geometric_distance=float(geo_dist),
532+
neighbor_rank=neighbor_rank,
533+
)
534+
)
535+
self.bootstrap_data.num_bootstraps += 1

‎src/scloop/data/types.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
FeatureSelectionMethod = Literal["hvg", "delve", "none"]
77
EmbeddingMethod = Literal["pca", "diffmap", "scvi"]
88
EmbeddingNeighbors = Literal["pca", "scvi"]
9+
LoopDistMethod = Literal["hausdorff", "frechet"]
910

1011
Index_t = Annotated[int, Field(ge=0)]
1112
Size_t = Annotated[int, Field(ge=0)]

0 commit comments

Comments
 (0)