diff --git a/.gitignore b/.gitignore index 0b5fab2..0134128 100644 --- a/.gitignore +++ b/.gitignore @@ -7,6 +7,11 @@ wheels/ *.egg-info *.so +*.a +*.la +src/scloop/utils/linear_algebra_gf2/include/ +src/scloop/utils/linear_algebra_gf2/pkgconfig/ + # Virtual environments .venv @@ -16,10 +21,21 @@ wheels/ .claude/ .specify/ +.ccb/ +.codex +.cursor/ +graphify-out/ todo/ *.ipynb +*.ipynb.bak* .vscode/ *.h5ad +*.csv +*.pdf +*.pkl +*.tar -# Refactor planning -REFACTOR_PLAN.md +/*.md +!/README.md +/*.png +/tmp/ diff --git a/examples/.gitignore b/examples/.gitignore index e33609d..557b79d 100644 --- a/examples/.gitignore +++ b/examples/.gitignore @@ -1 +1,6 @@ *.png +*.md +/*_SCVI_*/ +/checkpoints/ +/examples/ +/*data* diff --git a/lefthook.yml b/lefthook.yml new file mode 100644 index 0000000..fe5d99b --- /dev/null +++ b/lefthook.yml @@ -0,0 +1,24 @@ +pre-commit: + parallel: true + jobs: + - name: ruff + glob: "*.{py,pyi}" + stage_fixed: true + run: uv run --no-sync ruff check --select I --fix {staged_files} && uv run --no-sync ruff format {staged_files} + - name: pyproject-fmt + glob: "pyproject.toml" + stage_fixed: true + run: pyproject-fmt {staged_files} || true + +pre-bump: + fail_on_changes: always + jobs: + - name: ruff check + glob: "*.{py,pyi}" + run: uv run --no-sync ruff check --fix {all_files} + - name: ruff format + glob: "*.{py,pyi}" + run: uv run --no-sync ruff format {all_files} + - name: pyproject-fmt + glob: "pyproject.toml" + run: pyproject-fmt {all_files} || true diff --git a/mise.toml b/mise.toml index 601fb36..f10f261 100644 --- a/mise.toml +++ b/mise.toml @@ -1,5 +1,6 @@ [tools] uv = "latest" +lefthook = "latest" [env] CPLUS_INCLUDE_PATH = "{{config_root}}/src/scloop/data" @@ -54,6 +55,7 @@ run = [ [tasks.bump] run = [ + "lefthook run pre-bump", "uv version --bump patch", "git add uv.lock pyproject.toml", 'git commit -m "new tag"', diff --git a/src/scloop/analyzing/bootstrap.py b/src/scloop/analyzing/bootstrap.py index a1f4cd8..fc39119 100644 --- a/src/scloop/analyzing/bootstrap.py +++ b/src/scloop/analyzing/bootstrap.py @@ -201,13 +201,14 @@ def run_single_bootstrap( extra_diameter_homology_equivalence: float = DEFAULT_EXTRA_DIAM_EQUIVALENCE, filter_column_homology_equivalence: bool = True, column_trim_method: ColumnTrimMethod = DEFAULT_COLUMN_TRIM_METHOD, - death_scale: float | None = None, n_neighbors_column_trim: int = DEFAULT_N_NEIGHBORS_COLUMN_TRIM, full_pairwise_distance_matrix: csr_matrix | None = None, full_vertex_ids: list[int] | None = None, reconstruct_on_full_data: bool = False, candidate_method: Literal["geometric", "image"] = "geometric", require_homological_equivalence: bool = True, + compute_homotopy_coherence: bool = True, + source_cocycle_bases: list[tuple[float, np.ndarray] | None] | None = None, **kwargs, ) -> BootstrapResult: if candidate_method not in {"geometric", "image"}: @@ -296,6 +297,7 @@ def run_single_bootstrap( do_clean_cocycle_region=True, foreign_chord_mult=foreign_chord_mult, max_perimeter_mult=max_perimeter_mult, + validate_representatives=False, ) else: bootstrap_loop_classes = compute_loop_representatives( @@ -323,6 +325,7 @@ def run_single_bootstrap( bootstrap=True, foreign_chord_mult=foreign_chord_mult, max_perimeter_mult=max_perimeter_mult, + validate_representatives=True, ) if image_homology_result is not None: @@ -432,6 +435,8 @@ def run_single_bootstrap( or target_loop.representatives is None ): continue + source_loops = source_loop.filter_valid(source_loop.representatives) + target_loops = target_loop.filter_valid(target_loop.representatives) max_column_diameter = None if filter_column_homology_equivalence: if extra_diameter_homology_equivalence < 0: @@ -443,9 +448,33 @@ def run_single_bootstrap( max(source_loop.death, target_loop.death) + float(extra_diameter_homology_equivalence) * max_lifetime ) + cocycle_basis_masks = None + if not compute_homotopy_coherence and source_cocycle_bases is not None: + cocycle_basis = source_cocycle_bases[source_idx] + if cocycle_basis is not None: + max_column_diameter, cocycle_basis_masks = cocycle_basis + column_scores = ( + source_loop.column_proximity_scores( + original_boundary_matrix_d1, + embedding, + n_neighbors_column_trim, + ) + if column_trim_method == "loop_proximity" + and cocycle_basis_masks is None + else None + ) + source_fillings = ( + source_loop.filter_valid( + source_loop.fillings( + original_boundary_matrix_d1, column_trim_method, column_scores + ) + ) + if compute_homotopy_coherence + else None + ) match.topological_equivalence = check_homological_equivalence( - source_loops=source_loop.representatives, - target_loops=target_loop.representatives, + source_loops=source_loops, + target_loops=target_loops, boundary_matrix_d1=original_boundary_matrix_d1, n_pairs_check=n_pairs_check_equivalence, with_relaxation=with_relaxation_equivalence, @@ -453,18 +482,12 @@ def run_single_bootstrap( max_n_edges_relaxation=max_n_edges_relaxation_equivalence, max_column_diameter=max_column_diameter, cocycle_edge_mask=cocycle_edge_masks[source_idx], + compute_homotopy_coherence=compute_homotopy_coherence, + cocycle_basis_masks=cocycle_basis_masks, + source_fillings=source_fillings, column_trim_method=column_trim_method, embedding=embedding, - death_scale=death_scale, - column_scores=( - source_loop.column_proximity_scores( - original_boundary_matrix_d1, - embedding, - n_neighbors_column_trim, - ) - if column_trim_method == "loop_proximity" - else None - ), + column_scores=column_scores, ) match.boundary_checked = True @@ -499,10 +522,6 @@ def run_bootstrap_pipeline( **kwargs, ) -> list[BootstrapResult]: results: list[BootstrapResult] = [] - _original_deaths = [lc.death for lc in original_loop_classes if lc is not None] - global_death_scale = ( - max(_original_deaths) if _original_deaths else meta.bootstrap.threshold_homology - ) full_pairwise_distance_matrix: csr_matrix | None = None full_vertex_ids: list[int] | None = None @@ -516,19 +535,31 @@ def run_bootstrap_pipeline( ) ) - if kwargs.get("column_trim_method", DEFAULT_COLUMN_TRIM_METHOD) == "loop_proximity": + column_trim_method = kwargs.get("column_trim_method", DEFAULT_COLUMN_TRIM_METHOD) + warm_fillings = kwargs.get("require_homological_equivalence", True) and kwargs.get( + "compute_homotopy_coherence", True + ) + warm_embedding = None + if column_trim_method == "loop_proximity" and warm_fillings: assert meta.preprocess is not None assert meta.preprocess.embedding_method is not None warm_embedding = np.array(adata.obsm[f"X_{meta.preprocess.embedding_method}"]) - for loop_class in original_loop_classes: - if loop_class is not None and loop_class.representatives is not None: - loop_class.column_proximity_scores( - original_boundary_matrix_d1, - warm_embedding, - kwargs.get( - "n_neighbors_column_trim", DEFAULT_N_NEIGHBORS_COLUMN_TRIM - ), - ) + for loop_class in original_loop_classes: + if loop_class is None or loop_class.representatives is None: + continue + column_scores = ( + loop_class.column_proximity_scores( + original_boundary_matrix_d1, + warm_embedding, + kwargs.get("n_neighbors_column_trim", DEFAULT_N_NEIGHBORS_COLUMN_TRIM), + ) + if warm_embedding is not None + else None + ) + if warm_fillings: + loop_class.fillings( + original_boundary_matrix_d1, column_trim_method, column_scores + ) ExecutorClass = ThreadPoolExecutor @@ -542,7 +573,6 @@ def run_bootstrap_pipeline( meta=meta, original_loop_classes=original_loop_classes, original_boundary_matrix_d1=original_boundary_matrix_d1, - death_scale=global_death_scale, full_pairwise_distance_matrix=full_pairwise_distance_matrix, full_vertex_ids=full_vertex_ids, reconstruct_on_full_data=reconstruct_on_full_data, diff --git a/src/scloop/computing/coherence.py b/src/scloop/computing/coherence.py index 19303bf..dfebe98 100644 --- a/src/scloop/computing/coherence.py +++ b/src/scloop/computing/coherence.py @@ -1,214 +1,55 @@ # Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan) from __future__ import annotations -from typing import Mapping, Sequence +from collections.abc import Mapping, Sequence import numpy as np -from ..data.types import HomotopyCoherenceMethod -from ..data.utils import edge_idx_decode -from ..utils.distance_metrics import compute_loop_frechet - - -def global_h1_death_scale(persistence_diagram: list | None) -> float | None: - if not persistence_diagram or len(persistence_diagram) < 2: - return None - deaths = np.asarray(persistence_diagram[1][1], dtype=float) - finite = deaths[np.isfinite(deaths)] - return float(finite.max()) if finite.size else None - - -def _flip_triangle( - cycle: frozenset[int], triangle: frozenset[int], num_vertices: int -) -> frozenset[int] | None: - """Apply one elementary triangle deformation to a simple cycle.""" - - shared = cycle & triangle - if len(shared) == 2: - return cycle ^ triangle - if len(shared) != 1: - return None - - shared_vertices = set(edge_idx_decode(next(iter(shared)), num_vertices)) - triangle_vertices = { - vertex for edge in triangle for vertex in edge_idx_decode(edge, num_vertices) - } - new_vertices = triangle_vertices - shared_vertices - if len(new_vertices) != 1: - return None - cycle_vertices = { - vertex for edge in cycle for vertex in edge_idx_decode(edge, num_vertices) - } - if next(iter(new_vertices)) in cycle_vertices: - return None - return cycle ^ triangle - - -def _cycle_vertices(cycle: frozenset[int], num_vertices: int) -> list[int]: - adjacency: dict[int, list[int]] = {} - for edge in cycle: - tail, head = edge_idx_decode(edge, num_vertices) - adjacency.setdefault(tail, []).append(head) - adjacency.setdefault(head, []).append(tail) - start = next(iter(adjacency)) - order = [start] - previous, current = -1, start - while len(order) <= len(adjacency): - neighbours = adjacency[current] - if len(neighbours) != 2: - break - first, second = neighbours - following = first if first != previous else second - if following == start: - break - order.append(following) - previous, current = current, following - return order - - -def _cycle_coords( - cycle: frozenset[int], num_vertices: int, embedding: np.ndarray -) -> np.ndarray: - vertices = _cycle_vertices(cycle, num_vertices) - return np.ascontiguousarray(embedding[vertices], dtype=np.float64) - - -def _frechet_to_target( - cycle: frozenset[int], - target_coords: np.ndarray, - num_vertices: int, - embedding: np.ndarray, -) -> float: - cycle_coords = _cycle_coords(cycle, num_vertices, embedding) - return float( - min( - compute_loop_frechet(cycle_coords, target_coords), - compute_loop_frechet(target_coords, cycle_coords), - ) +from ..data.boundary import BoundaryMatrixD1 +from ..data.constants import DEFAULT_COLUMN_TRIM_METHOD +from ..data.types import ColumnTrimMethod +from .homology import compute_loop_homological_equivalence + + +def compute_loop_fillings( + loop_mask: np.ndarray, + boundary_matrix_d1: BoundaryMatrixD1, + death: float, + column_trim_method: ColumnTrimMethod = DEFAULT_COLUMN_TRIM_METHOD, + column_scores: np.ndarray | None = None, +) -> list[tuple[int, ...] | None]: + n_loops = loop_mask.shape[0] + result = compute_loop_homological_equivalence( + boundary_matrix_d1=boundary_matrix_d1, + loop_mask_a=loop_mask, + loop_mask_b=np.zeros((1, loop_mask.shape[1]), dtype=bool), + n_pairs_check=n_loops, + with_relaxation=False, + max_column_diameter=death, + column_trim_method=column_trim_method, + column_scores=column_scores, + compute_filling=True, ) + fillings: list[tuple[int, ...] | None] = [None] * n_loops + for (loop_index, _), deformation in zip( + result.loop_pairs_matched, result.mapping_deformation_matched + ): + fillings[loop_index] = deformation["triangle_ids"] + return fillings -def _greedy_reversal( - source: frozenset[int], - target: frozenset[int], - triangles: tuple[frozenset[int], ...], - num_vertices: int, - embedding: np.ndarray, - death_scale: float, +def compute_coherence( + filling: Sequence[int] | None, + deformation: Sequence[int], + triangle_areas: Mapping[int, float], ) -> float | None: - target_coords = _cycle_coords(target, num_vertices, embedding) - source_coords = _cycle_coords(source, num_vertices, embedding) - cycle = source - reversal = 0.0 - total_step_abs = 0.0 - max_step_abs = 0.0 - n_steps = 0 - remaining = set(range(len(triangles))) - - while remaining: - best_key = None - best_step = None - for triangle_index in remaining: - triangle = triangles[triangle_index] - next_cycle = _flip_triangle(cycle, triangle, num_vertices) - if next_cycle is None: - continue - - current_remaining_cost = _frechet_to_target( - cycle, target_coords, num_vertices, embedding - ) - next_remaining_cost = _frechet_to_target( - next_cycle, target_coords, num_vertices, embedding - ) - current_departed_cost = _frechet_to_target( - cycle, source_coords, num_vertices, embedding - ) - next_departed_cost = _frechet_to_target( - next_cycle, source_coords, num_vertices, embedding - ) - step_reversal = max( - next_remaining_cost - current_remaining_cost, - 0.0, - ) - step_abs = abs(next_remaining_cost - current_remaining_cost) + abs( - next_departed_cost - current_departed_cost - ) - key = (next_remaining_cost, triangle_index) - if best_key is None or key < best_key: - best_key = key - best_step = ( - triangle_index, - next_cycle, - step_reversal, - step_abs, - ) - - if best_step is None: - return None - - triangle_index, cycle, step_reversal, step_abs = best_step - remaining.remove(triangle_index) - reversal += step_reversal - total_step_abs += step_abs - max_step_abs = max(max_step_abs, step_abs) - n_steps += 1 - - if cycle != target: + if filling is None: return None - # multiply by 2 due to two way frechet - return 1 - min(max_step_abs / (2 * death_scale), 1) - - -def path_finding_coherence( - source_edges: Sequence[int], - target_edges: Sequence[int], - triangles: Sequence[Sequence[int]], - num_vertices: int, - embedding: np.ndarray, - death_scale: float, -) -> float | None: - - source = frozenset(source_edges) - target = frozenset(target_edges) - triangle_sets = tuple(frozenset(triangle) for triangle in triangles) - forward = _greedy_reversal( - source, - target, - triangle_sets, - num_vertices, - embedding, - death_scale, - ) - - reverse = _greedy_reversal( - target, - source, - triangle_sets, - num_vertices, - embedding, - death_scale, - ) - - candidates = [cost for cost in (forward, reverse) if cost is not None] - if not candidates: + source = set(filling) + homotopy = set(deformation) + area_union = sum(triangle_areas[t] for t in source | homotopy) + if area_union <= 0: return None - return min(candidates) - - -def compute_coherence( - source_edges: Sequence[int], - target_edges: Sequence[int], - triangles: Sequence[Sequence[int]], - num_vertices: int, - embedding: np.ndarray, - death_scale: float, - method: HomotopyCoherenceMethod = "path_finding", -) -> float | None: - return path_finding_coherence( - source_edges, - target_edges, - triangles, - num_vertices, - embedding, - death_scale, - ) + # F ∩ (F + H) = F \ H and F ∪ (F + H) = F ∪ H over GF2 + # Based on Jaccard index + return sum(triangle_areas[t] for t in source - homotopy) / area_union diff --git a/src/scloop/computing/homology.py b/src/scloop/computing/homology.py index e704d33..bda3398 100644 --- a/src/scloop/computing/homology.py +++ b/src/scloop/computing/homology.py @@ -646,6 +646,8 @@ def compute_loop_homological_equivalence( cocycle_edge_mask: np.ndarray | None = None, column_trim_method: ColumnTrimMethod = DEFAULT_COLUMN_TRIM_METHOD, column_scores: np.ndarray | None = None, + compute_filling: bool = True, + cocycle_basis_masks: np.ndarray | None = None, ) -> LoopClassEquivalence: """ Parameters @@ -662,6 +664,11 @@ def compute_loop_homological_equivalence( ``diameter`` keeps the largest-diameter triangles (legacy). ``loop_proximity`` keeps the columns nearest to the loops, ranked by ``column_scores`` (one score per boundary-matrix column, lower is nearer). + compute_filling: bool + Solves d2 x = a + b over GF2 + cocycle_basis_masks: np.ndarray | None + Boolean mask of shape (n_classes, n_edges), one row per H1 class alive in the + complex up to ``max_column_diameter`` """ assert loop_mask_a.shape[1] == boundary_matrix_d1.shape[0] assert loop_mask_b.shape[1] == boundary_matrix_d1.shape[0] @@ -694,6 +701,20 @@ def compute_loop_homological_equivalence( ] result.n_loop_pairs_checked = n_pairs_check + if not compute_filling: + if cocycle_basis_masks is None: + raise ValueError("compute_filling=False requires cocycle_basis_masks") + sums = loop_sums[:n_pairs_check] + parities = ( + sums.astype(np.int64) @ np.asarray(cocycle_basis_masks, dtype=np.int64).T + ) % 2 + is_boundary = ~parities.any(axis=1) + if max_column_diameter is not None: + row_diams = np.asarray(boundary_matrix_d1.row_simplex_diams, dtype=float) + is_boundary &= ~(sums & (row_diams > max_column_diameter)).any(axis=1) + result.loop_pairs_matched = [pairs_kept[i] for i in np.flatnonzero(is_boundary)] + return result + one_ridx_A = np.asarray(boundary_matrix_d1.data[0], dtype=int) one_cidx_A = np.asarray(boundary_matrix_d1.data[1], dtype=int) nrow_A = boundary_matrix_d1.shape[0] diff --git a/src/scloop/computing/loops.py b/src/scloop/computing/loops.py index c147eb0..90aceb9 100644 --- a/src/scloop/computing/loops.py +++ b/src/scloop/computing/loops.py @@ -10,17 +10,22 @@ from numba import jit from pydantic import PositiveFloat from scipy.sparse import csr_matrix, triu +from sklearn.neighbors import NearestNeighbors from ..data.base_components import LoopClass from ..data.boundary import BoundaryMatrixD1 from ..data.constants import ( DEFAULT_FOREIGN_CHORD_MULT, + DEFAULT_K_LOCAL_SCALE, DEFAULT_K_YEN, DEFAULT_LIFE_PCT, + DEFAULT_MAX_INSERT_PER_EDGE, DEFAULT_MAX_PERIMETER_MULT, DEFAULT_N_COCYCLES_USED, DEFAULT_N_FORCE_DEVIATE, DEFAULT_N_REPS_PER_LOOP, + DEFAULT_SPLIT_EDGE_LENGTH_MULT, + DEFAULT_SPLIT_POINT_DISTANCE_MULT, NUMERIC_EPSILON, ) from ..data.types import Count_t, Percent_t @@ -148,6 +153,7 @@ def compute_loop_representatives( do_clean_cocycle_region: bool = False, foreign_chord_mult: float = DEFAULT_FOREIGN_CHORD_MULT, max_perimeter_mult: float = DEFAULT_MAX_PERIMETER_MULT, + validate_representatives: bool = True, ) -> list[LoopClass | None]: assert pairwise_distance_matrix.shape is not None @@ -178,8 +184,18 @@ def compute_loop_representatives( results: list[LoopClass | None] = [None] * len(indices_top_k) build_foreign = foreign_chord_mult > 1.0 + check_reps = validate_representatives and persistence_pair_simplices is not None + death_keys: list[tuple[float, tuple[int, ...]]] = [] + if check_reps: + assert persistence_pair_simplices is not None + death_keys = [ + (float(loop_deaths[c]), tuple(-v for v in sorted(simplex, reverse=True))) + if len(simplex) > 0 + else (math.inf, ()) + for c, simplex in enumerate(persistence_pair_simplices[1]) + ] cocycle_edges_per_class: dict[int, set[tuple[int, int]]] = {} - if build_foreign: + if build_foreign or check_reps: for cls_idx in range(len(cocycles)): cocycle_edges_per_class[cls_idx] = { (min(a, b), max(a, b)) @@ -276,6 +292,19 @@ def compute_loop_representatives( max_perimeter_mult=max_perimeter_mult, ) + representatives_valid = None + if check_reps: + own_key = death_keys[loop_idx] + alive = [ + c + for c in range(len(death_keys)) + if loop_births[c] <= own_key[0] and death_keys[c] >= own_key + ] + representatives_valid = [ + _is_death_cycle(loop, int(loop_idx), alive, cocycle_edges_per_class) + for loop in loops_local + ] + loops = [[vertex_ids[v] for v in loop] for loop in loops_local] loops_coords = loops_to_coords(embedding=embedding, loops_vertices=loops) @@ -288,12 +317,29 @@ def compute_loop_representatives( death_simplex=death_simplex, cocycles=cocycles[loop_idx], representatives=loops, + representatives_valid=representatives_valid, coordinates_vertices_representatives=loops_coords, ) return results +def _is_death_cycle( + loop: Sequence[int], + own_class: int, + alive_classes: list[int], + cocycle_edges: dict[int, set[tuple[int, int]]], +) -> bool: + """[loop] = [∂τ] just before the death triangle τ enters, via the alive cocycle basis.""" + edges: set[tuple[int, int]] = set() + for u, v in zip(loop, [*loop[1:], loop[0]]): + if u != v: + edges ^= {(min(u, v), max(u, v))} + return all( + len(edges & cocycle_edges[c]) % 2 == int(c == own_class) for c in alive_classes + ) + + def reconstruct_n_loop_representatives( cocycles_dim1: List, edges: np.ndarray, @@ -379,7 +425,7 @@ def reconstruct_n_loop_representatives( if foreign_cocycle_edges and foreign_chord_mult > 1.0: own_chord_keys = {(min(e), max(e)) for e in all_cocycle_edges} foreign_keys = { - (min(int(u), int(v)), max(int(u), int(v))) for u, v in foreign_cocycle_edges + (min(u, v), max(u, v)) for u, v in foreign_cocycle_edges } - own_chord_keys for idx, e in enumerate(edge_list): if e in foreign_keys: @@ -472,7 +518,7 @@ def _select_diverse_loops( max_perimeter_mult: float = DEFAULT_MAX_PERIMETER_MULT, ) -> Tuple[List[List[int]], List[float]]: pairs = sorted( - [(float(d), list(c)) for d, c in zip(distances, cycles) if math.isfinite(d)], + [(d, list(c)) for d, c in zip(distances, cycles) if math.isfinite(d)], key=lambda x: x[0], ) if not pairs: @@ -496,10 +542,200 @@ def _select_diverse_loops( idxs = [] for i in range(n_return): pct = (lower_pct + step * i) / 100 - idx = min(int(math.floor(n_total * pct)), n_total - 1) + idx = min(math.floor(n_total * pct), n_total - 1) idxs.append(idx) selected = [pairs[i] for i in idxs] dists = [p[0] for p in selected] loops = [p[1] for p in selected] return loops, dists + + +def _distances_to_edge(points: np.ndarray, a: np.ndarray, b: np.ndarray) -> np.ndarray: + ab = b - a + denom = float(ab @ ab) + if denom == 0.0: + return np.linalg.norm(points - a, axis=1) + s = np.clip((points - a) @ ab / denom, 0.0, 1.0) + return np.linalg.norm(points - (a + s[:, None] * ab), axis=1) + + +def _select_split_point( + a: int, + b: int, + u: int, + v: int, + embedding: np.ndarray, + used_mask: np.ndarray, + max_diameter: float, + limit: float | None, + pool: np.ndarray | None, + pool_embedding: np.ndarray, +) -> int | None: + points = pool_embedding + d_a = np.linalg.norm(points - embedding[a], axis=1) + mask = d_a <= max_diameter + mask &= ~used_mask if pool is None else ~used_mask[pool] + idx = np.flatnonzero(mask) + if idx.size == 0: + return None + candidate_points = points[idx] + keep = np.linalg.norm(candidate_points - embedding[b], axis=1) <= max_diameter + if not keep.any(): + return None + candidates = (idx if pool is None else pool[idx])[keep] + candidate_points = candidate_points[keep] + + if limit is not None: + keep = _distances_to_edge(candidate_points, embedding[u], embedding[v]) <= limit + if not keep.any(): + return None + candidates = candidates[keep] + candidate_points = candidate_points[keep] + + midpoint = 0.5 * (embedding[a] + embedding[b]) + return int( + candidates[np.argmin(np.linalg.norm(candidate_points - midpoint, axis=1))] + ) + + +def _densify_loops( + vertices: list[int], + embedding: np.ndarray, + local_scale: np.ndarray, + max_diameter: float, + max_insert_per_edge: int, + split_edge_length_mult: float | None = None, + split_point_distance_mult: float | None = None, +) -> list[int]: + if len(vertices) < 2: + return list(vertices) + + refined: list[int] = [vertices[0]] + used_mask = np.zeros(embedding.shape[0], dtype=bool) + used_mask[vertices] = True + + if split_edge_length_mult is None: + split_edge_length_mult = 1.0 + + ball_cutoff = 0.25 * float( + np.linalg.norm(embedding.max(axis=0) - embedding.min(axis=0)) + ) + + def _target_length(a: int, b: int) -> float: + return float(split_edge_length_mult * 0.5 * (local_scale[a] + local_scale[b])) + + for u, v in zip(vertices[:-1], vertices[1:]): + poly = [u, v] + lengths = [float(np.linalg.norm(embedding[u] - embedding[v]))] + targets = [_target_length(u, v)] + limit = ( + split_point_distance_mult * 0.5 * (local_scale[u] + local_scale[v]) + if split_point_distance_mult is not None + else None + ) + pool: np.ndarray | None = None + pool_embedding: np.ndarray = embedding + limit_select = limit + pool_ready = limit is None + stalled_pairs: set[tuple[int, int]] = set() + n_inserted = 0 + while n_inserted < max_insert_per_edge: + edges_to_split = [] + for i in range(len(poly) - 1): + a, b = poly[i], poly[i + 1] + if (a, b) in stalled_pairs: + continue + if lengths[i] > targets[i]: + edges_to_split.append((targets[i] / lengths[i], i, a, b)) + if not edges_to_split: + break + edges_to_split.sort(key=lambda x: x[0]) + if not pool_ready: + assert limit is not None + half_len = 0.5 * float(np.linalg.norm(embedding[u] - embedding[v])) + if limit + half_len < ball_cutoff: + mid_uv = 0.5 * (embedding[u] + embedding[v]) + near_mask = np.linalg.norm(embedding - mid_uv, axis=1) <= ( + limit + half_len + ) * (1.0 + 1e-12) + if 2 * int(near_mask.sum()) <= embedding.shape[0]: + near = np.flatnonzero(near_mask) + pool = near[ + _distances_to_edge( + embedding[near], embedding[u], embedding[v] + ) + <= limit + ] + pool_embedding = embedding[pool] + limit_select = None + pool_ready = True + + inserted_this_round = False + for _, i, a, b in edges_to_split: + p = _select_split_point( + a=a, + b=b, + u=u, + v=v, + embedding=embedding, + used_mask=used_mask, + max_diameter=max_diameter, + limit=limit_select, + pool=pool, + pool_embedding=pool_embedding, + ) + if p is None: + stalled_pairs.add((a, b)) + continue + poly.insert(i + 1, p) + lengths[i] = float(np.linalg.norm(embedding[a] - embedding[p])) + lengths.insert( + i + 1, float(np.linalg.norm(embedding[p] - embedding[b])) + ) + targets[i] = _target_length(a, p) + targets.insert(i + 1, _target_length(p, b)) + used_mask[p] = True + n_inserted += 1 + inserted_this_round = True + break + if not inserted_this_round: + break + refined.extend(poly[1:]) + + return refined + + +def refine_loop_representatives( + loop_classes: list[LoopClass], + embedding: np.ndarray, + local_scale: np.ndarray | None = None, + max_insert_per_edge: int = DEFAULT_MAX_INSERT_PER_EDGE, + split_edge_length_mult: float | None = DEFAULT_SPLIT_EDGE_LENGTH_MULT, + split_point_distance_mult: float | None = DEFAULT_SPLIT_POINT_DISTANCE_MULT, + life_pct: float = 0.0, + k_local_scale: int = DEFAULT_K_LOCAL_SCALE, +) -> None: + if local_scale is None: + nn = NearestNeighbors(n_neighbors=k_local_scale + 1).fit(embedding) + knn_distances, _ = nn.kneighbors(embedding) + local_scale = np.asarray(knn_distances[:, 1:].mean(axis=1), dtype=np.float64) + for loop_class in loop_classes: + if loop_class is None or not loop_class.representatives: + continue + max_diameter = loop_class.birth + life_pct * ( + loop_class.death - loop_class.birth + ) + refined_all = [] + for rep in loop_class.representatives: + refined = _densify_loops( + vertices=list(rep), + embedding=embedding, + local_scale=local_scale, + max_diameter=max_diameter, + max_insert_per_edge=max_insert_per_edge, + split_edge_length_mult=split_edge_length_mult, + split_point_distance_mult=split_point_distance_mult, + ) + refined_all.append(refined) + loop_class.representatives_refined = refined_all diff --git a/src/scloop/computing/matching.py b/src/scloop/computing/matching.py index 094fa11..2d69698 100644 --- a/src/scloop/computing/matching.py +++ b/src/scloop/computing/matching.py @@ -15,16 +15,15 @@ from ..data.boundary import BoundaryMatrixD1 from ..data.constants import ( DEFAULT_COLUMN_TRIM_METHOD, - DEFAULT_N_NEIGHBORS_COLUMN_TRIM, DEFAULT_MAX_N_EDGES_RELAXATION_EQUIVALENCE, DEFAULT_N_HUBS_RELAXATION_EQUIVALENCE, + DEFAULT_N_NEIGHBORS_COLUMN_TRIM, DEFAULT_N_PAIRS_CHECK, DEFAULT_WITH_RELAXATION_EQUIVALENCE, ) from ..data.types import ( ColumnTrimMethod, Count_t, - HomotopyCoherenceMethod, LoopDistMethod, LoopEdges, PositiveFloat, @@ -127,6 +126,28 @@ def cocycle_to_edge_mask( return mask +def cocycle_basis_to_edge_masks( + cocycles: list, + persistence_diagram: list, + persistence_pair_simplices: list, + diameter: float, + boundary_matrix_d1: BoundaryMatrixD1, + vertex_ids: list[int], +) -> np.ndarray: + births, deaths = persistence_diagram + _, death_simplices = persistence_pair_simplices + masks = [] + for cocycle, birth, death, death_simplex in zip( + cocycles, births, deaths, death_simplices + ): + if birth > diameter or (len(death_simplex) > 0 and death <= diameter): + continue + mask = cocycle_to_edge_mask(cocycle, boundary_matrix_d1, vertex_ids) + if mask is not None: + masks.append(mask) + return np.array(masks, dtype=bool).reshape(len(masks), boundary_matrix_d1.shape[0]) + + def compute_geometric_distance( source_coords_list: list[list[list[float]]], target_coords_list: list[list[list[float]]], @@ -153,29 +174,23 @@ def check_homological_equivalence( max_column_diameter: PositiveFloat | None = None, cocycle_edge_mask: np.ndarray | None = None, compute_homotopy_coherence: bool = True, - homotopy_coherence_method: HomotopyCoherenceMethod = "path_finding", + cocycle_basis_masks: np.ndarray | None = None, + source_fillings: list[tuple[int, ...] | None] | None = None, column_trim_method: ColumnTrimMethod = DEFAULT_COLUMN_TRIM_METHOD, column_scores: np.ndarray | None = None, embedding: np.ndarray | None = None, - death_scale: float = 1.0, n_neighbors_column_trim: int = DEFAULT_N_NEIGHBORS_COLUMN_TRIM, ) -> LoopClassEquivalence: if len(source_loops) == 0 or len(target_loops) == 0: return LoopClassEquivalence() - loop_edges_a = loops_to_edge_mask( - loops=source_loops, - boundary_matrix_d1=boundary_matrix_d1, - return_valid_indices=compute_homotopy_coherence, + mask_a = loops_to_edge_mask( + loops=source_loops, boundary_matrix_d1=boundary_matrix_d1 ) - loop_edges_b = loops_to_edge_mask( - loops=target_loops, - boundary_matrix_d1=boundary_matrix_d1, - return_valid_indices=compute_homotopy_coherence, + mask_b = loops_to_edge_mask( + loops=target_loops, boundary_matrix_d1=boundary_matrix_d1 ) - - mask_a = loop_edges_a.mask if isinstance(loop_edges_a, LoopEdges) else loop_edges_a - mask_b = loop_edges_b.mask if isinstance(loop_edges_b, LoopEdges) else loop_edges_b + assert isinstance(mask_a, np.ndarray) and isinstance(mask_b, np.ndarray) if ( column_scores is None @@ -207,78 +222,33 @@ def check_homological_equivalence( cocycle_edge_mask=cocycle_edge_mask, column_trim_method=column_trim_method, column_scores=column_scores, + compute_filling=compute_homotopy_coherence or cocycle_basis_masks is None, + cocycle_basis_masks=cocycle_basis_masks, ) - if not compute_homotopy_coherence or embedding is None: + if not compute_homotopy_coherence or source_fillings is None: return result - assert isinstance(loop_edges_a, LoopEdges) - assert isinstance(loop_edges_b, LoopEdges) - row_edge_ids = np.asarray(boundary_matrix_d1.row_simplex_ids, dtype=int) - row_indices = np.asarray(boundary_matrix_d1.data[0], dtype=int) - column_indices = np.asarray(boundary_matrix_d1.data[1], dtype=int) - edges_per_column = row_edge_ids[ - row_indices[np.argsort(column_indices, kind="stable")] - ].reshape(-1, 3) - triangle_edges = { - triangle_id: tuple(edges) - for triangle_id, edges in zip( - boundary_matrix_d1.col_simplex_ids, edges_per_column.tolist() + triangle_areas = boundary_matrix_d1.compute_col_areas() + result.homotopy_coherence_matched = [ + compute_coherence( + filling=source_fillings[source_index], + deformation=deformation["triangle_ids"], + triangle_areas=triangle_areas, ) - } - - for (source_index, target_index), deformation in zip( - result.loop_pairs_matched, - result.mapping_deformation_matched, - ): - result.homotopy_coherence_matched.append( - compute_coherence( - source_edges=loop_edges_a.edge_ids_per_rep[source_index], - target_edges=loop_edges_b.edge_ids_per_rep[target_index], - triangles=[ - triangle_edges[triangle_id] - for triangle_id in deformation["triangle_ids"] - ], - num_vertices=boundary_matrix_d1.num_vertices, - embedding=embedding, - death_scale=death_scale, - method=homotopy_coherence_method, - ) + for (source_index, _), deformation in zip( + result.loop_pairs_matched, result.mapping_deformation_matched ) - - for (source_index, target_index), deformation in zip( - result.loop_pairs_matched_relax, - result.mapping_deformation_matched_relax, - ): - source = set(loop_edges_a.edge_ids_per_rep[source_index]) - target = set(loop_edges_b.edge_ids_per_rep[target_index]) - relaxation_edges = set(deformation["relaxation_edge_ids"]) - triangles = [ - triangle_edges[triangle_id] for triangle_id in deformation["triangle_ids"] - ] - scores = [ - compute_coherence( - source_edges=tuple(source ^ relaxation_edges), - target_edges=tuple(target), - triangles=triangles, - num_vertices=boundary_matrix_d1.num_vertices, - embedding=embedding, - death_scale=death_scale, - method=homotopy_coherence_method, - ), - compute_coherence( - source_edges=tuple(source), - target_edges=tuple(target ^ relaxation_edges), - triangles=triangles, - num_vertices=boundary_matrix_d1.num_vertices, - embedding=embedding, - death_scale=death_scale, - method=homotopy_coherence_method, - ), - ] - valid_scores = [score for score in scores if score is not None] - result.homotopy_coherence_matched_relax.append( - float(np.mean(valid_scores)) if valid_scores else None + ] + result.homotopy_coherence_matched_relax = [ + compute_coherence( + filling=source_fillings[source_index], + deformation=deformation["triangle_ids"], + triangle_areas=triangle_areas, + ) + for (source_index, _), deformation in zip( + result.loop_pairs_matched_relax, result.mapping_deformation_matched_relax ) + ] return result diff --git a/src/scloop/data/analysis_containers.py b/src/scloop/data/analysis_containers.py index e4967da..d749f60 100644 --- a/src/scloop/data/analysis_containers.py +++ b/src/scloop/data/analysis_containers.py @@ -563,6 +563,7 @@ def _get_track_embedding( idx_track: Index_t, embedding_alt: np.ndarray | None = None, keep_matches: str = "equivalent", + use_refined: bool = False, ) -> list[np.ndarray]: assert idx_track in self.loop_tracks loops = [] @@ -576,14 +577,22 @@ def _get_track_embedding( if embedding_alt is None: if loop_class.coordinates_vertices_representatives is not None: loops.extend( - loop_class.coordinates_vertices_representatives + loop_class.filter_valid( + loop_class.coordinates_vertices_representatives + ) ) else: - if loop_class.representatives is not None: + reps = loop_class.representatives + if ( + use_refined + and loop_class.representatives_refined is not None + ): + reps = loop_class.representatives_refined + if reps is not None: loops.extend( loops_to_coords( embedding=embedding_alt, - loops_vertices=loop_class.representatives, + loops_vertices=loop_class.filter_valid(reps), ) ) return loops @@ -594,6 +603,7 @@ def _get_loop_embedding( idx_loop_class: Index_t, idx_loop: Index_t | None = None, embedding_alt: np.ndarray | None = None, + use_refined: bool = False, ) -> list[list[list[float]]]: if idx_bootstrap < len(self.selected_loop_classes) and idx_loop_class < len( self.selected_loop_classes[idx_bootstrap] @@ -602,29 +612,30 @@ def _get_loop_embedding( if loop_class is not None: if embedding_alt is None: if loop_class.coordinates_vertices_representatives is not None: + coords = loop_class.filter_valid( + loop_class.coordinates_vertices_representatives + ) if idx_loop is None: - return loop_class.coordinates_vertices_representatives + return coords else: - assert idx_loop < len( - loop_class.coordinates_vertices_representatives - ) - return [ - loop_class.coordinates_vertices_representatives[ - idx_loop - ] - ] + assert idx_loop < len(coords) + return [coords[idx_loop]] else: - if loop_class.representatives is not None: + reps = loop_class.representatives + if use_refined and loop_class.representatives_refined is not None: + reps = loop_class.representatives_refined + if reps is not None: + reps = loop_class.filter_valid(reps) if idx_loop is None: return loops_to_coords( embedding=embedding_alt, - loops_vertices=loop_class.representatives, + loops_vertices=reps, ) else: - assert idx_loop < len(loop_class.representatives) + assert idx_loop < len(reps) return loops_to_coords( embedding=embedding_alt, - loops_vertices=[loop_class.representatives[idx_loop]], + loops_vertices=[reps[idx_loop]], ) return [] @@ -1080,6 +1091,11 @@ def from_super( for coords in super_obj.coordinates_vertices_representatives ] representatives = [list(rep) for rep in super_obj.representatives] + representatives_refined = ( + [list(rep) for rep in super_obj.representatives_refined] + if super_obj.representatives_refined is not None + else None + ) if len(coordinates_vertices) > 0: if ref_area is None: @@ -1089,6 +1105,12 @@ def from_super( if ref_area * signed_area_2d(coordinates_vertices[i]) < 0: coordinates_vertices[i] = coordinates_vertices[i][::-1] representatives[i] = representatives[i][::-1] + if representatives_refined is not None and i < len( + representatives_refined + ): + representatives_refined[i] = representatives_refined[i][ + ::-1 + ] coordinates_edges = [ (emb[0:-1, :] + emb[1:, :]) / 2 for emb in coordinates_vertices @@ -1121,6 +1143,8 @@ def from_super( death_simplex=super_obj.death_simplex, cocycles=super_obj.cocycles, representatives=representatives, + representatives_refined=representatives_refined, + representatives_valid=super_obj.representatives_valid, coordinates_vertices_representatives=[ c.tolist() for c in coordinates_vertices ], diff --git a/src/scloop/data/base_components.py b/src/scloop/data/base_components.py index 275b2e9..ca0fe26 100644 --- a/src/scloop/data/base_components.py +++ b/src/scloop/data/base_components.py @@ -1,15 +1,15 @@ # Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan) from __future__ import annotations -from typing import TYPE_CHECKING +from itertools import compress +from typing import TYPE_CHECKING, Self import numpy as np from pydantic import BaseModel, Field, model_validator from pydantic.dataclasses import dataclass from scipy.spatial.distance import cdist -from typing_extensions import Self -from .types import Diameter_t, Index_t, Percent_t, PositiveFloat +from .types import ColumnTrimMethod, Diameter_t, Index_t, Percent_t, PositiveFloat if TYPE_CHECKING: import h5py @@ -116,9 +116,12 @@ class LoopClass(BaseModel): death_simplex: list[Index_t] = Field(default_factory=list) cocycles: list | None = None representatives: list[list[Index_t]] | None = None + representatives_refined: list[list[Index_t]] | None = None + representatives_valid: list[bool] | None = None coordinates_vertices_representatives: list[list[list[float]]] | None = None _cached_column_scores: np.ndarray | None = None + _cached_fillings: list[tuple[Index_t, ...] | None] | None = None model_config = {"arbitrary_types_allowed": True} @@ -128,6 +131,12 @@ def check_birth_death(self) -> Self: raise ValueError("loop dies before its birth") return self + def filter_valid[T](self, per_rep: list[T]) -> list[T]: + valid = self.representatives_valid + if valid and len(valid) == len(per_rep): + return list(compress(per_rep, valid)) + return per_rep + @property def lifetime(self): return self.death - self.birth @@ -154,6 +163,31 @@ def column_proximity_scores( ) return self._cached_column_scores + def fillings( + self, + boundary_matrix_d1: BoundaryMatrixD1, + column_trim_method: ColumnTrimMethod, + column_scores: np.ndarray | None = None, + ) -> list[tuple[Index_t, ...] | None]: + if self._cached_fillings is not None: + return self._cached_fillings + + if self.representatives is None: + raise ValueError("loop class has no representatives to fill") + from ..computing.coherence import compute_loop_fillings + from ..computing.matching import loops_to_edge_mask + + loop_mask = loops_to_edge_mask(self.representatives, boundary_matrix_d1) + assert isinstance(loop_mask, np.ndarray) + self._cached_fillings = compute_loop_fillings( + loop_mask=loop_mask, + boundary_matrix_d1=boundary_matrix_d1, + death=self.death, + column_trim_method=column_trim_method, + column_scores=column_scores, + ) + return self._cached_fillings + @property def persistence_pair(self) -> PersistencePair: return PersistencePair( @@ -210,6 +244,20 @@ def to_hdf5_group(self, group: h5py.Group, compress: bool = True) -> None: str(i), data=np.array(rep, dtype=np.int64), **kw ) + if self.representatives_refined is not None: + reps_ref_grp = group.create_group("representatives_refined") + for i, rep in enumerate(self.representatives_refined): + reps_ref_grp.create_dataset( + str(i), data=np.array(rep, dtype=np.int64), **kw + ) + + if self.representatives_valid is not None: + group.create_dataset( + "representatives_valid", + data=np.asarray(self.representatives_valid, dtype=bool), + **kw, + ) + if self.coordinates_vertices_representatives is not None: coords_grp = group.create_group("coordinates_vertices_representatives") for i, coords in enumerate(self.coordinates_vertices_representatives): @@ -257,6 +305,21 @@ def _read_base_fields(group: h5py.Group) -> dict: for i in range(len(reps_grp)): representatives.append(np.asarray(reps_grp[str(i)]).tolist()) + representatives_refined = None + if "representatives_refined" in group: + reps_ref_grp: h5py.Group = group["representatives_refined"] # type: ignore[assignment] + representatives_refined = [] + for i in range(len(reps_ref_grp)): + representatives_refined.append( + np.asarray(reps_ref_grp[str(i)]).tolist() + ) + + representatives_valid = ( + np.asarray(group["representatives_valid"], dtype=bool).tolist() + if "representatives_valid" in group + else None + ) + coordinates_vertices_representatives = None if "coordinates_vertices_representatives" in group: coords_grp: h5py.Group = group["coordinates_vertices_representatives"] # type: ignore[assignment] @@ -275,6 +338,8 @@ def _read_base_fields(group: h5py.Group) -> dict: death_simplex=death_simplex, cocycles=cocycles, representatives=representatives, + representatives_refined=representatives_refined, + representatives_valid=representatives_valid, coordinates_vertices_representatives=coordinates_vertices_representatives, ) diff --git a/src/scloop/data/boundary.py b/src/scloop/data/boundary.py index 82410b5..4efc409 100644 --- a/src/scloop/data/boundary.py +++ b/src/scloop/data/boundary.py @@ -104,6 +104,7 @@ def _from_hdf5_group_data(cls, group: h5py.Group) -> dict: class BoundaryMatrixD1(BoundaryMatrix): _cached_edge_set: set[tuple[Index_t, Index_t]] | None = None _cached_col_vertices: np.ndarray | None = None + _cached_col_areas: dict[Index_t, float] | None = None @property def row_simplex_decode(self) -> list[tuple[Index_t, Index_t]]: @@ -137,6 +138,22 @@ def compute_col_vertices(self) -> np.ndarray: self._cached_col_vertices = verts return self._cached_col_vertices + def compute_col_areas(self) -> dict[Index_t, float]: + if self._cached_col_areas is not None: + return self._cached_col_areas + + row_indices = np.asarray(self.data[0], dtype=np.int64) + col_indices = np.asarray(self.data[1], dtype=np.int64) + sides = np.asarray(self.row_simplex_diams, dtype=np.float64)[ + row_indices[np.argsort(col_indices, kind="stable")] + ].reshape(-1, 3) + a, b, c = sides.T + # Heron's formula + heron = (a + b + c) * (-a + b + c) * (a - b + c) * (a + b - c) + areas = 0.25 * np.sqrt(np.clip(heron, 0.0, None)) + self._cached_col_areas = dict(zip(self.col_simplex_ids, areas.tolist())) + return self._cached_col_areas + def to_hdf5_group(self, group: h5py.Group, compress: bool = True) -> None: group.attrs["_type"] = "BoundaryMatrixD1" super().to_hdf5_group(group, compress=compress) diff --git a/src/scloop/data/constants.py b/src/scloop/data/constants.py index 579e0c5..76dc3dc 100644 --- a/src/scloop/data/constants.py +++ b/src/scloop/data/constants.py @@ -5,7 +5,13 @@ from IPython.display import Javascript -from .types import ColumnTrimMethod, LoopDistMethod, PositiveFloat +from .types import ( + CandidateMethod, + ColumnTrimMethod, + LoopDistMethod, + PositiveFloat, + PresenceTestMethod, +) CROSS_MATCH_KEY = "X_scloop_aligned" CROSS_MATCH_RESULT_KEY = "scloop_cross_match" @@ -28,6 +34,8 @@ DEFAULT_CUTOFF_PVAL: float = 0.05 DEFAULT_MAX_ROWS_BOUNDARY_MATRIX: int = 30000 DEFAULT_N_BOOTSTRAP: int = 10 +DEFAULT_CANDIDATE_METHOD: CandidateMethod = "image" +DEFAULT_PRESENCE_METHOD: PresenceTestMethod = "chi2" DEFAULT_AUTO_THRESHOLD_FACTOR: float = 1.75 @@ -40,6 +48,10 @@ # Keep reconstructed reps with perimeter <= α * L_min before diversity # sampling. math.inf disables the gate. Typical range: 1.5–3. DEFAULT_MAX_PERIMETER_MULT: float = float("inf") +DEFAULT_MAX_INSERT_PER_EDGE: int = 20 +DEFAULT_SPLIT_EDGE_LENGTH_MULT: float = 1.25 +DEFAULT_SPLIT_POINT_DISTANCE_MULT: float = 4.0 +DEFAULT_K_LOCAL_SCALE: int = 10 DEFAULT_N_PAIRS_CHECK_EQUIVALENCE: int = 4 # typically one neighbor is sufficient for checking DEFAULT_K_NEIGHBORS_CHECK_EQUIVALENCE: int = 1 diff --git a/src/scloop/data/containers.py b/src/scloop/data/containers.py index 5068044..8316332 100644 --- a/src/scloop/data/containers.py +++ b/src/scloop/data/containers.py @@ -26,10 +26,10 @@ from ..computing.homology import ( compute_persistence_diagram_and_cocycles, ) -from ..computing.loops import compute_loop_representatives -from ..computing.coherence import global_h1_death_scale +from ..computing.loops import compute_loop_representatives, refine_loop_representatives from ..computing.matching import ( check_homological_equivalence, + cocycle_basis_to_edge_masks, cocycle_to_edge_mask, compute_geometric_distance, loops_to_edge_mask, @@ -54,6 +54,7 @@ DEFAULT_K_YEN, DEFAULT_LIFE_PCT, DEFAULT_LOOP_DIST_METHOD, + DEFAULT_MAX_INSERT_PER_EDGE, DEFAULT_MAX_N_EDGES_RELAXATION_EQUIVALENCE, DEFAULT_MAX_PERIMETER_MULT, DEFAULT_MAXITER_EIGENDECOMPOSITION, @@ -68,6 +69,8 @@ DEFAULT_N_PAIRS_CHECK_EQUIVALENCE, DEFAULT_N_REPS_PER_LOOP, DEFAULT_NOISE_SCALE, + DEFAULT_SPLIT_EDGE_LENGTH_MULT, + DEFAULT_SPLIT_POINT_DISTANCE_MULT, DEFAULT_TIMEOUT_EIGENDECOMPOSITION, DEFAULT_WEIGHT_HODGE, DEFAULT_WITH_RELAXATION_EQUIVALENCE, @@ -359,6 +362,34 @@ def _compute_boundary_matrix_d1( **nei_kwargs, ) + def _cocycle_bases_before_death(self) -> list[tuple[float, np.ndarray] | None]: + assert self.boundary_matrix_d1 is not None + assert self.persistence_diagram is not None + assert self.persistence_pair_simplices is not None + assert self.cocycles is not None + bases: list[tuple[float, np.ndarray] | None] = [] + for loop_class in self.selected_loop_classes: + if loop_class is None: + bases.append(None) + continue + scale = float( + np.nextafter(np.float32(loop_class.death), np.float32(-np.inf)) + ) + bases.append( + ( + scale, + cocycle_basis_to_edge_masks( + cocycles=self.cocycles[1], + persistence_diagram=self.persistence_diagram[1], + persistence_pair_simplices=self.persistence_pair_simplices[1], + diameter=scale, + boundary_matrix_d1=self.boundary_matrix_d1, + vertex_ids=self._original_vertex_ids, + ), + ) + ) + return bases + @property def _original_vertex_ids(self): assert self.meta.preprocess is not None @@ -488,6 +519,50 @@ def _compute_hodge_analysis_for_track( kwargs_gene_trends=kwargs_gene_trends, ) + def _refine_loop_representatives( + self, + embedding: np.ndarray, + local_scale: np.ndarray | None = None, + max_insert_per_edge: int = DEFAULT_MAX_INSERT_PER_EDGE, + split_edge_length_mult: float | None = DEFAULT_SPLIT_EDGE_LENGTH_MULT, + split_point_distance_mult: float | None = DEFAULT_SPLIT_POINT_DISTANCE_MULT, + life_pct: float = 0.0, + include_bootstrap: bool = True, + ) -> None: + assert self.selected_loop_classes is not None + loop_classes_refined = [] + seen_ids: set[int] = set() + for i, c in enumerate(self.selected_loop_classes): + if c is None: + continue + loop_classes_refined.append(c) + if include_bootstrap and self.bootstrap_data is not None: + if i in self.bootstrap_data.loop_tracks: + for boot_id, loop_id in self.bootstrap_data.loop_tracks[ + i + ].track_ipairs: + if boot_id < len( + self.bootstrap_data.selected_loop_classes + ) and loop_id < len( + self.bootstrap_data.selected_loop_classes[boot_id] + ): + cm = self.bootstrap_data.selected_loop_classes[boot_id][ + loop_id + ] + if id(cm) in seen_ids: + continue + seen_ids.add(id(cm)) + loop_classes_refined.append(cm) + refine_loop_representatives( + loop_classes=loop_classes_refined, + embedding=embedding, + local_scale=local_scale, + max_insert_per_edge=max_insert_per_edge, + split_edge_length_mult=split_edge_length_mult, + split_point_distance_mult=split_point_distance_mult, + life_pct=life_pct, + ) + def _compute_loop_representatives( self, embedding: np.ndarray, @@ -618,6 +693,7 @@ def _get_loop_embedding( embedding_alt: np.ndarray | None = None, include_bootstrap: bool = True, keep_matches: str = "equivalent", + use_refined: bool = False, ) -> list[list[list[float]]]: """ Use embedding stored in LoopClass by default @@ -631,38 +707,38 @@ def _get_loop_embedding( if loop_class is not None: if embedding_alt is None: if loop_class.coordinates_vertices_representatives is not None: + coords = loop_class.filter_valid( + loop_class.coordinates_vertices_representatives + ) if ( idx_loop is None or include_bootstrap ): # if include bootstrap, then idx_loop means loop index among all loops in a track - loops.extend( - loop_class.coordinates_vertices_representatives - ) + loops.extend(coords) else: - assert idx_loop < len( - loop_class.coordinates_vertices_representatives - ) - loops.append( - loop_class.coordinates_vertices_representatives[ - idx_loop - ] - ) + assert idx_loop < len(coords) + loops.append(coords[idx_loop]) else: - if loop_class.representatives is not None: + reps = loop_class.representatives + if ( + use_refined + and loop_class.representatives_refined is not None + ): + reps = loop_class.representatives_refined + if reps is not None: + reps = loop_class.filter_valid(reps) if idx_loop is None or include_bootstrap: loops.extend( loops_to_coords( embedding=embedding_alt, - loops_vertices=loop_class.representatives, + loops_vertices=reps, ) ) else: - assert idx_loop < len(loop_class.representatives) + assert idx_loop < len(reps) loops.extend( loops_to_coords( embedding=embedding_alt, - loops_vertices=[ - loop_class.representatives[idx_loop] - ], + loops_vertices=[reps[idx_loop]], ) ) @@ -674,6 +750,7 @@ def _get_loop_embedding( idx_track=selector, embedding_alt=embedding_alt, keep_matches=keep_matches, + use_refined=use_refined, ) ) else: @@ -696,6 +773,7 @@ def _get_loop_embedding( idx_loop_class=selector[1], idx_loop=idx_loop, embedding_alt=embedding_alt, + use_refined=use_refined, ) ) return loops @@ -751,6 +829,8 @@ def _assess_bootstrap_homology_equivalence( column_trim_method: ColumnTrimMethod = DEFAULT_COLUMN_TRIM_METHOD, embedding: np.ndarray | None = None, n_neighbors_column_trim: Count_t = DEFAULT_N_NEIGHBORS_COLUMN_TRIM, + compute_homotopy_coherence: bool = True, + cocycle_basis: tuple[float, np.ndarray] | None = None, ) -> tuple[int, int, bool]: assert self.bootstrap_data is not None self._ensure_loop_tracks() @@ -776,9 +856,8 @@ def _assess_bootstrap_homology_equivalence( or target_loop_class.representatives is None ): return (source_class_idx, target_class_idx, False) - - source_loops = source_loop_class.representatives - target_loops = target_loop_class.representatives + source_loops = source_loop_class.filter_valid(source_loop_class.representatives) + target_loops = target_loop_class.filter_valid(target_loop_class.representatives) if len(source_loops) == 0 or len(target_loops) == 0: return (source_class_idx, target_class_idx, False) @@ -805,6 +884,29 @@ def _assess_bootstrap_homology_equivalence( boundary_matrix_d1=self.boundary_matrix_d1, vertex_ids=self._original_vertex_ids, ) + cocycle_basis_masks = None + if not compute_homotopy_coherence and cocycle_basis is not None: + max_column_diameter, cocycle_basis_masks = cocycle_basis + column_scores = ( + source_loop_class.column_proximity_scores( + self.boundary_matrix_d1, + embedding, + n_neighbors_column_trim, + ) + if column_trim_method == "loop_proximity" + and embedding is not None + and cocycle_basis_masks is None + else None + ) + source_fillings = ( + source_loop_class.filter_valid( + source_loop_class.fillings( + self.boundary_matrix_d1, column_trim_method, column_scores + ) + ) + if compute_homotopy_coherence + else None + ) result = check_homological_equivalence( source_loops=source_loops, target_loops=target_loops, @@ -815,18 +917,12 @@ def _assess_bootstrap_homology_equivalence( max_n_edges_relaxation=max_n_edges_relaxation, max_column_diameter=max_column_diameter, cocycle_edge_mask=cocycle_edge_mask, + compute_homotopy_coherence=compute_homotopy_coherence, + cocycle_basis_masks=cocycle_basis_masks, + source_fillings=source_fillings, column_trim_method=column_trim_method, embedding=embedding, - death_scale=global_h1_death_scale(self.persistence_diagram), - column_scores=( - source_loop_class.column_proximity_scores( - self.boundary_matrix_d1, - embedding, - n_neighbors_column_trim, - ) - if column_trim_method == "loop_proximity" and embedding is not None - else None - ), + column_scores=column_scores, ) is_equivalent = result.is_equivalent(relax=with_relaxation) return (source_class_idx, target_class_idx, is_equivalent) @@ -880,6 +976,7 @@ def _bootstrap( method_geometric_equivalence: LoopDistMethod = DEFAULT_LOOP_DIST_METHOD, candidate_method: Literal["geometric", "image"] = "geometric", require_homological_equivalence: bool = True, + compute_homotopy_coherence: bool = True, reconstruct_on_full_data: bool = False, verbose: bool = False, progress_main: Progress | None = None, @@ -896,6 +993,11 @@ def _bootstrap( else: self.meta.bootstrap.indices_resample.clear() + cocycle_bases: list[tuple[float, np.ndarray] | None] = ( + self._cocycle_bases_before_death() + if require_homological_equivalence and not compute_homotopy_coherence + else [None] * len(self.selected_loop_classes) + ) if ( use_parallel or candidate_method == "image" @@ -935,6 +1037,8 @@ def _bootstrap( method_geometric_equivalence=method_geometric_equivalence, candidate_method=candidate_method, require_homological_equivalence=require_homological_equivalence, + compute_homotopy_coherence=compute_homotopy_coherence, + source_cocycle_bases=cocycle_bases, reconstruct_on_full_data=reconstruct_on_full_data, verbose=verbose, progress_main=progress_main, @@ -1085,6 +1189,8 @@ def _bootstrap( column_trim_method=column_trim_method, embedding=embedding, n_neighbors_column_trim=n_neighbors_column_trim, + compute_homotopy_coherence=compute_homotopy_coherence, + cocycle_basis=cocycle_bases[si], ) tasks[task] = (si, tj, neighbor_distances[si, k], k) @@ -1155,6 +1261,8 @@ def _bootstrap_parallel( method_geometric_equivalence: LoopDistMethod = DEFAULT_LOOP_DIST_METHOD, candidate_method: Literal["geometric", "image"] = "geometric", require_homological_equivalence: bool = True, + compute_homotopy_coherence: bool = True, + source_cocycle_bases: list[tuple[float, np.ndarray] | None] | None = None, reconstruct_on_full_data: bool = False, verbose: bool = False, progress_main: Progress | None = None, @@ -1194,6 +1302,8 @@ def _bootstrap_parallel( method_geometric_equivalence=method_geometric_equivalence, candidate_method=candidate_method, require_homological_equivalence=require_homological_equivalence, + compute_homotopy_coherence=compute_homotopy_coherence, + source_cocycle_bases=source_cocycle_bases, n_pairs_check_equivalence=n_pairs_check_equivalence, with_relaxation_equivalence=with_relaxation_equivalence, n_hubs_relaxation_equivalence=n_hubs_relaxation_equivalence, diff --git a/src/scloop/data/types.py b/src/scloop/data/types.py index a0bd78e..59cc6b8 100644 --- a/src/scloop/data/types.py +++ b/src/scloop/data/types.py @@ -9,11 +9,11 @@ EmbeddingNeighbors = Literal["pca", "scvi"] LoopDistMethod = Literal["hausdorff", "frechet"] ColumnTrimMethod = Literal["diameter", "loop_proximity"] -HomotopyCoherenceMethod = Literal["path_finding"] MultipleTestCorrectionMethod = Literal["bonferroni", "benjamini-hochberg"] PresenceTestMethod = Literal["fisher", "chi2", "wilcoxon"] CrossMatchModelTypes = Literal["mlp", "nf"] CrossMatchRoutes = Literal["geometric", "image"] +CandidateMethod = Literal["geometric", "image"] LogLevel = Literal[ "TRACE", "DEBUG", diff --git a/src/scloop/matching/cross_matching.py b/src/scloop/matching/cross_matching.py index f48a955..cb60ec7 100644 --- a/src/scloop/matching/cross_matching.py +++ b/src/scloop/matching/cross_matching.py @@ -87,6 +87,7 @@ def _reconstruct_image_loop_classes( [birth_simplices[idx] for idx in selected], [death_simplices[idx] for idx in selected], ), + validate_representatives=False, **kwargs_reconstruct, ) return { @@ -122,10 +123,10 @@ def _attribute_to_loop_class( return None, None, None assert side.boundary_matrix_d1 is not None column_scores = None - if ( - kwargs_equivalence.get("column_trim_method", DEFAULT_COLUMN_TRIM_METHOD) - == "loop_proximity" - ): + column_trim_method = kwargs_equivalence.get( + "column_trim_method", DEFAULT_COLUMN_TRIM_METHOD + ) + if column_trim_method == "loop_proximity": column_scores = image_loop_class.column_proximity_scores( side.boundary_matrix_d1, side.embedding, @@ -152,12 +153,9 @@ def _attribute_to_loop_class( max_column_diameter=max(image_loop_class.death, loop_class.death) + extra_diameter * max_lifetime, cocycle_edge_mask=cocycle_edge_mask, + compute_homotopy_coherence=False, column_scores=column_scores, embedding=side.embedding, - death_scale=max( - (lc.death for lc in side.loop_classes if lc is not None), - default=None, - ), **kwargs_equivalence, ) if equivalence.is_equivalent(relax=with_relaxation): diff --git a/src/scloop/matching/utils.py b/src/scloop/matching/utils.py deleted file mode 100644 index 3fefe6c..0000000 --- a/src/scloop/matching/utils.py +++ /dev/null @@ -1,860 +0,0 @@ -import os -import re - -import numpy as np - -import subprocess - -def send_cmd_windows(cmd) : # cmd is a string - 'dont forget stdout = good output, stderr = error like outputs e.g. help output for windows' - out = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE).communicate()[0].decode('UTF-8').split('\n') - # choosing first good output - return out - -def send_cmd_linux(cmd) : - out = os.popen(cmd).read().split('\n') - return out - -#### define your system -send_cmd = send_cmd_windows - -############ - -from match.utils_plot import * - - -################ - -##### PERSISTENCE - -def compute_bars_tightreps(inp = None, filename = 'data') : - '''This function computes the barcode and representatives of a: - - inp = point cloud in the form of an array - - filename = file containing the lower diagonal matrix containing the pairwise distances for a finite metric space''' - if inp is None : # a filename input which gives lower distance matrix - if filename.endswith('.lower_distance_matrix') : - ldm_file = filename - else : - ldm_file = filename + '.lower_distance_matrix' - else : - data = inp # a point cloud in 2D or 3D - ldm_file = '{}.lower_distance_matrix'.format(filename) - pairwise = np.sqrt( np.sum( (data[:, None, :] - data[None, :, :])**2, axis=-1) ) - - f = open(ldm_file, "w") # erase possibly pre-existing - for i in range(len(pairwise)) : - f.write(', '.join([ str(x) for x in pairwise[i,:i]])) - f.write('\n') - f.close() - - software = "./ripser-representatives" - options = "" - - command = "{} {} {}".format(software, options, ldm_file) - - out = send_cmd(command) - return out - -def extract_bars_reps(out, only_dim_1 = False, verbose = False) : - '''This function converts the output of compute_bars_tight_reps into list of bars and reps, organised by dimension''' - # find after which line it starts enumerating intervals in dim 0,1,2 - line_PH = {0:len(out), 1:len(out), 2:len(out)} - for i in range(len(out)) : - if out[i].startswith('persistent homology intervals in dim ') : - dim = out[i].rstrip().split(' ')[-1][:-1] # after split '0:' - dim = int(dim) - line_PH[dim] = i - - bars = {0:[],1:[],2:[]} - reps = {0:[],1:[],2:[]} - tight_reps = {0:[],1:[],2:[]} - - # reps 0 [ [[0],[1]], [[1],[2]], etc ] - # reps 1 [ [ [0,1], [1,2], [2,3], [3,0] ], same] - - if not only_dim_1 : - - # 0-dim PH bars and reps - dim = 0 - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - x = re.search(r"\[(\d*.\d*),(\d*.\d*)\)", out[i]) - if not x : raise ValueError("no match found") - if x.group(2) != ' ' : # finite bars first - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - y = re.search(r"\{\[(\d+)\], \[(\d+)\]\}", out[i]) - if not y : - raise ValueError("finite interval detected but not represented by two vertices") - reps[dim] += [ [ [int(y.group(1))], [int(y.group(2))] ] ] - else : - bars[dim] += [ [float(x.group(1)), np.inf] ] - y = re.search(r"\{\[(\d+)\].*\}", out[i]) - reps[dim] += [ [ [int(y.group(1))] ] ] - i += 1 - - # 1-dim PH bars and reps - dim = 1 - - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - x = re.search(r"\[(\d*.\d*),(\d*.\d*)\)", out[i]) - if i == len(out) - 1 : # trivial string '' - break - - # all "finite" bars (no missing second endpoint) - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - - i += 1 # next line for tight reps - y = re.findall(r"\[(\d+),(\d+)\] \(\d*.\d*\)", out[i]) - y = [list(elem) for elem in y] - y = [[int(e[0]), int(e[1])] for e in y] - tight_reps[dim] += [ y ] - - i += 1 # again, next line for reps - y = re.findall(r"\[(\d+),(\d+)\] \(\d*.\d*\)", out[i]) - y = [list(elem) for elem in y] - y = [[int(e[0]), int(e[1])] for e in y] - reps[dim] += [ y ] - - i += 1 - - if verbose : - print(bars) - print(reps) - print(tight_reps) - - return bars, reps, tight_reps - -def extract_bars_reps_indices(out, only_dim_1 = False, verbose = False) : - '''This function converts the output of compute_bars_tight_reps into list of bars, representatives and indices of the persistence pairs, - organised by dimension. REMARK: you need to use the modified version or ripser_tight_representative_cycles''' - # find after which line it starts enumerating intervals in dim 0,1,2 - line_PH = {0:len(out), 1:len(out), 2:len(out)} - for i in range(len(out)) : - if out[i].startswith('persistent homology intervals in dim ') : - dim = out[i].rstrip().split(' ')[-1][:-1] # after split '0:' - dim = int(dim) - line_PH[dim] = i - # print(line_PH) - # allows for repeated line: then takes last one - - bars = {0:[],1:[],2:[]} - reps = {0:[],1:[],2:[]} - tight_reps = {0:[],1:[],2:[]} - indices = {1:[],2:[]} - - # reps 0 [ [[0],[1]], [[1],[2]], etc ] - # reps 1 [ [ [0,1], [1,2], [2,3], [3,0] ], same] - - if not only_dim_1 : - - # 0-dim PH bars and reps - dim = 0 - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - x = re.search(r"\[(\d*.\d*),(\d*.\d*)\)", out[i]) - if not x : raise ValueError("no intervals found") - if x.group(2) != ' ' : # finite bars first - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - y = re.search(r"\{\[(\d+)\], \[(\d+)\]\}", out[i]) - if not y : - raise ValueError("finite interval detected but not represented by two vertices") - reps[dim] += [ [ [int(y.group(1))], [int(y.group(2))] ] ] - else : - bars[dim] += [ [float(x.group(1)), np.inf] ] - y = re.search(r"\{\[(\d+)\].*\}", out[i]) - reps[dim] += [ [ [int(y.group(1))] ] ] - i += 1 - - - # 1-dim PH bars and reps - dim = 1 - - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - x = re.search(r"\[(\d*.\d*),(\d*.\d*)\)", out[i]) - if i == len(out) - 1 : # trivial string '' - break - - # all "finite" bars (no missing second endpoint) - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - - - # indices - z = re.search(r"indices: (\d*)-(\d*)", out[i]) - if not z : raise ValueError("no iindices found --- are you using the modified version of ripser-tight-representative-cycles?") - indices[dim] += [ [int(z.group(1)), int(z.group(2))] ] - - i += 1 # next line for tight reps - y = re.findall(r"\[(\d+),(\d+)\] \(\d*.\d*\)", out[i]) - y = [list(elem) for elem in y] - y = [[int(e[0]), int(e[1])] for e in y] - tight_reps[dim] += [ y ] - - i += 1 # again, next line for reps - y = re.findall(r"\[(\d+),(\d+)\] \(\d*.\d*\)", out[i]) - y = [list(elem) for elem in y] - y = [[int(e[0]), int(e[1])] for e in y] - reps[dim] += [ y ] - - i += 1 - - if verbose : - print(bars) - print(reps) - print(tight_reps) - - return bars, reps, tight_reps, indices - -##### IMAGE-PERSISTENCE - -def compute_image_bars(filename_X = 'X', filename_Z = 'Z', threshold = None) : - '''This function computes the barcode of the image-persistence of X inside of Z. - The input consists on the two lower distance matrices, using the extension explained in the reference paper, and the treshold up - to which their VR complexes coincide.''' - - if filename_X.endswith('.lower_distance_matrix') : - ldm_file_X = filename_X - ldm_file_Z = filename_Z - else : - ldm_file_X = filename_X + '.lower_distance_matrix' - ldm_file_Z = filename_Z + '.lower_distance_matrix' - - software = "./ripser-image" - - if threshold is None : - options = "--dim 1 --subfiltration {}".format(ldm_file_X) - else : - options = "--dim 1 --threshold {} --subfiltration {}".format(threshold, ldm_file_X) - - command = "{} {} {}".format(software, options, ldm_file_Z) - - out = send_cmd(command) - - return out - -def extract_bars(out, only_dim_1 = False, verbose = False) : - ''' This function converts the output of compute_image_bars into list of bars organised by dimension - (simpler version than extract_bars_reps, no reps for image-persistence)''' - - line_PH = {0:len(out), 1:len(out), 2:len(out)} - for i in range(len(out)) : - if out[i].startswith('persisten') : # not the same output message depending on ripser-feature version! - dim = out[i].rstrip().split(' ')[-1][:-1] # after split '0:' - dim = int(dim) - line_PH[dim] = i - - bars = {0:[],1:[],2:[]} - - if not only_dim_1 : - - # 0-dim PH bars - dim = 0 - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - x = re.search(r"\[(\d*.\d*),(\d*.\d*)\)", out[i]) - if not x : raise ValueError("no match found") - if x.group(2) != ' ' : # finite bars first - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - else : - bars[dim] += [ [float(x.group(1)), np.inf] ] - i += 1 - - # 1-dim PH bars - dim = 1 - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - x = re.search(r"\[(\d*.\d*),(\d*.*\d*)\)", out[i]) - if i == len(out) - 1 : # trivial string '' - break - if x : - if x.group(2) != ' ' : - # "finite" bar (no missing second endpoint) - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - else : - bars[dim] += [ [float(x.group(1)), np.inf] ] - - i += 1 - - if verbose : - print(bars) - - return bars - - -def extract_bars_indices(out, only_dim_1 = False, verbose = False) : - ''' This function converts the output of compute_image_bars into list of bars and indices of the persistence pairs organised by dimension - (simpler version than extract_bars_reps_indices, no reps for image-persistence). REMARK: need to use the modified version of ripser-image!''' - - line_PH = {0:len(out), 1:len(out), 2:len(out)} - for i in range(len(out)) : - if out[i].startswith('persisten') : # not the same output message depending on ripser-feature version! - dim = out[i].rstrip().split(' ')[-1][:-1] # after split '0:' - dim = int(dim) - line_PH[dim] = i - - - bars = {0:[],1:[],2:[]} - indices = {1: [], 2: []} - - if not only_dim_1 : - - # 0-dim PH bars - dim = 0 - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - x = re.search(r"\[(\d*.\d*),(\d*.\d*)\)", out[i]) - if not x : raise ValueError("no match found") - if x.group(2) != ' ' : # finite bars first - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - else : - bars[dim] += [ [float(x.group(1)), np.inf] ] - i += 1 - - # 1-dim PH bars - dim = 1 - i = line_PH[dim]+1 - while i < line_PH[dim + 1] : - - #bars - x = re.search(r"\[(\d*.\d*),(\d*.*\d*)\)", out[i]) - if i == len(out) - 1 : # trivial string '' - break - if x : - if x.group(2) != ' ' : - # "finite" bar (no missing second endpoint) - bars[dim] += [ [float(x.group(1)), float(x.group(2))] ] - else : - bars[dim] += [ [float(x.group(1)), np.inf] ] - - #indices - z = re.search(r"indices: (\d*)-(\d*)", out[i]) - indices[dim] += [ [int(z.group(1)), int(z.group(2))] ] - if not z : raise ValueError("no iindices found --- are you using the modified version of ripser-image?") - - i += 1 - - if verbose : - print(bars) - - return bars, indices - -##### MATCHING - -def Jaccard(a,b,c,d) : - # Jaccard index = intersection over union of two intervals [a,b] and [c,d] - # used to measure affinity of two intervals - M1 = max(a,c) - m1 = min(b,d) - if M1 < m1 : - Jac = (m1 - M1) / (max(b,d) - min(a,c)) - else : - Jac = 0 - return Jac - -def compute_affinity(birth_X, death_X, death, birth_Y, death_Y, affinity_method = 'A') : - if affinity_method == 'A' : # Yohai's and Omer's - a_X_Y = Jaccard( birth_X, death_X, birth_Y, death_Y ) - a_X_Z = Jaccard( birth_X, death_X, birth_X, death ) - a_Y_Z = Jaccard( birth_Y, death_Y, birth_Y, death ) - affinity = a_X_Y * a_X_Z * a_Y_Z - - if affinity_method == 'B' : - a_XZ_YZ = Jaccard( birth_X, death, birth_Y, death ) - a_X_Z = Jaccard( birth_X, death_X, birth_X, death ) - a_Y_Z = Jaccard( birth_Y, death_Y, birth_Y, death ) - affinity = a_XZ_YZ * a_X_Z * a_Y_Z - - if affinity_method == 'C' : - a_X_Y = Jaccard( birth_X, death_X, birth_Y, death_Y ) - a_XZ_YZ = Jaccard( birth_X, death, birth_Y, death ) - a_X_Z = Jaccard( birth_X, death_X, birth_X, death ) - a_Y_Z = Jaccard( birth_Y, death_Y, birth_Y, death ) - affinity = a_X_Y * a_XZ_YZ * a_X_Z * a_Y_Z - - if affinity_method == 'D' : - a_X_Y = Jaccard( birth_X, death_X, birth_Y, death_Y ) - a_XZ_YZ = Jaccard( birth_X, death, birth_Y, death ) - affinity = a_X_Y * a_XZ_YZ - - return affinity - - -def argsort(seq, option = 'desc'): - # what permutation to apply to indices in order to sort seq in ascending / descending order - if option == 'asc' : - return sorted(range(len(seq)), reverse = False, key= seq.__getitem__) - elif option == 'desc' : - return sorted(range(len(seq)), reverse = True, key= seq.__getitem__) - -def show_matches(X, Y, matched_X_Y, affinity_X_Y, tight_reps_X, tight_reps_Y, dim = 1, - zoom_factor = 3, show_together = False) : - ''' This function displays the matches between two point-clouds X and Y after performing the matching and obtaining - the lists: matched_X_Y and affinity_X_Y. Also needed the corresponding lists of tight_reps.''' - - Z = np.vstack((X,Y)) - arg = argsort(affinity_X_Y) - affinity_X_Y = np.array(affinity_X_Y)[arg] - matched_X_Y = np.array(matched_X_Y)[arg] - - if len(matched_X_Y) == 1 or not show_together : - for match, aff in zip(matched_X_Y, affinity_X_Y) : - print('new match') - a,b = match - - if X.shape[1] == 2 : - fig, axes = plt.subplots(1,2, figsize = (8,5), sharex = True, sharey = True) - axes[0].scatter(Y[:,0], Y[:,1], alpha = .2) - plot_cycreps(Z, [tight_reps_X[dim][a]], pts_to_show = X, ax = axes[0]) - axes[1].scatter(X[:,0], X[:,1], alpha = .2) - plot_cycreps(Z, [tight_reps_Y[dim][b]], pts_to_show = Y, ax = axes[1]) - for ax in axes: - ax.set_aspect('equal') - axes[0].set_xlabel('X') - axes[1].set_xlabel('Y') - #plt.text(2, 5, 'a match with affinity {}'.format(aff)) - fig.suptitle('a match with affinity {}'.format(aff), y=0.78) - plt.show() - - if X.shape[1] == 3 : - fig = plt.figure(figsize=plt.figaspect(1.)*zoom_factor) # figaspect(0.5)*1.5 - fig.suptitle('a match with affinity {}'.format(aff), y=0.7) - ax = fig.add_subplot(1,2,1, projection='3d') - ax.scatter(Y[:,0],Y[:,1],Y[:,2], alpha = .1) - plot_cycreps(Z, [tight_reps_X[dim][a]], pts_to_show = X, ax = ax) - ax.set_title('X') - - ax = fig.add_subplot(1,2,2, projection='3d') - ax.scatter(X[:,0],X[:,1],X[:,2], alpha = .1) #, c = '#729dcf', s = 50, edgecolors='black') - plot_cycreps(Z, [tight_reps_Y[dim][b]], pts_to_show = Y, ax = ax) - ax.set_title('Y') - plt.show() - - else : # show_together : - - if X.shape[1] == 2 : - fig, axes = plt.subplots(len(matched_X_Y), 2, figsize = (10, 6 * len(matched_X_Y)), - sharex = True, sharey = True) - # will be buggy if len(matched_X_Y) == 1 - i = 0 - for match, aff in zip(matched_X_Y, affinity_X_Y) : - a,b = match - - axes[i,0].scatter(Y[:,0], Y[:,1], alpha = .2) - plot_cycreps(Z, [tight_reps_X[dim][a]], pts_to_show = X, ax = axes[i,0]) - axes[i,1].scatter(X[:,0], X[:,1], alpha = .2) - plot_cycreps(Z, [tight_reps_Y[dim][b]], pts_to_show = Y, ax = axes[i,1]) - - axes[i,0].set_xlabel('X') - axes[i,0].set_title('aff =') - axes[i,1].set_title(aff) - axes[i,1].set_xlabel('Y') - #plt.text(2, 5, 'a match with affinity {}'.format(aff)) - #fig.suptitle('a match with affinity {}'.format(aff), y=0.78) - for ax in axes.ravel(): - ax.set_aspect('equal') - plt.show() - - - if X.shape[1] == 3 : - fig, axes = plt.subplots(len(matched_X_Y), 2, subplot_kw={'projection': '3d'}, - figsize = (10, 4.5 * len(matched_X_Y))) - #fig = plt.figure(figsize=plt.figaspect(1.)*zoom_factor) # figaspect(0.5)*1.5 - #fig.suptitle('all matches', y=0.7) - - i = 0 - for match, aff in zip(matched_X_Y, affinity_X_Y) : - a,b = match - #print('a,b=',a,b) - ax = axes[i,0] - #ax = fig.add_subplot(i,2,1, projection='3d') - ax.scatter(Y[:,0],Y[:,1],Y[:,2], alpha = .1) - plot_cycreps(Z, [tight_reps_X[dim][a]], pts_to_show = X, ax = ax) - ax.set_title('X aff = {}'.format(aff)) - - ax = axes[i,1] - #ax = fig.add_subplot(i,2,2, projection='3d') - ax.scatter(X[:,0],X[:,1],X[:,2], alpha = .1) #, c = '#729dcf', s = 50, edgecolors='black') - plot_cycreps(Z, [tight_reps_Y[dim][b]], pts_to_show = Y, ax = ax) - ax.set_title('Y a = {} b = {}'.format(a,b)) - - i += 1 - - plt.tight_layout() - plt.show() - - -def duplicates_list(a) : # duplicates_list([1,2,3,1,2,2]) = [1, 2, 2] - seen = set() - dupes = [x for x in a if x in seen or seen.add(x)] - return dupes - -def find_occurences_list(a, val) : - indices = [i for i, x in enumerate(a) if x == val] - return indices - - -def find_match(bars_X, bars_X_Z, indices_X, indices_X_Z, bars_Y, bars_Y_Z, indices_Y, indices_Y_Z, dim = 1, affinity_method = 'A', - check_Morse = False, check_ambiguous_deaths = False) : - ''' This funtion find the matches between the barcodes of X and Y providing the barcodes of their image-persistence modules in the union. - Affinity score is automatically set to A but can be changed. Optiona outputs to check if the filtrations provided are Morse and if - there are image-bars sharing death times in the barcodes.''' - - matched_X_Y = [] - affinity_X_Y = [] - - # consider all image-bars - births_X_Z = [a[0] for a in bars_X_Z[dim]] - births_Y_Z = [a[0] for a in bars_Y_Z[dim]] - deaths_X_Z = [a[1] for a in bars_X_Z[dim]] - deaths_Y_Z = [a[1] for a in bars_Y_Z[dim]] - - # consider normal bars - births_X = [a[0] for a in bars_X[dim]] - births_Y = [a[0] for a in bars_Y[dim]] - deaths_X = [a[1] for a in bars_X[dim]] - deaths_Y = [a[1] for a in bars_Y[dim]] - - if check_Morse : - # adding noise to your point clouds does not solve the following exceptions. - # It will create a distance matrix with unique values, so that only adding 1 edge at a time - # but possibly many triangles that kill cycles simultaneously in Rips complexes - if len(duplicates_list(deaths_X_Z)) > 0 : - print('Found duplicate deaths in X_Z') - if len(duplicates_list(deaths_Y_Z)) > 0 : - print('Found duplicate deaths in Y_Z') - if len(duplicates_list(births_X)) > 0 : - print('Found duplicate births in X') # should never happen for unique distance values - if len(duplicates_list(births_Y)) > 0 : - print('Found duplicate births in Y') # should never happen for unique distance values - - - considered_deaths_X_Z = set(deaths_X_Z) - considered_deaths_Y_Z = set(deaths_Y_Z) - - # find common (considered) deaths in image - common_deaths = considered_deaths_X_Z.intersection(considered_deaths_Y_Z) - - - if check_ambiguous_deaths : - if set(duplicates_list(deaths_X_Z)).intersection(set(duplicates_list(deaths_Y_Z))) != set() : - print('Found common duplicate deaths in X_Z and Y_Z!!!') - if set(duplicates_list(deaths_X_Z)).intersection(common_deaths) != set() : - print('Found duplicate death in X_Z common with Y_Z') - if set(duplicates_list(deaths_Y_Z)).intersection(common_deaths) != set() : - print('Found duplicate death in Y_Z common with X_Z') - - # determine ambiguous deaths - ambiguous_deaths_X_Z = set(duplicates_list(deaths_X_Z)).intersection(common_deaths) - ambiguous_deaths_Y_Z = set(duplicates_list(deaths_Y_Z)).intersection(common_deaths) - ambiguous_deaths = ambiguous_deaths_X_Z.union(ambiguous_deaths_Y_Z) - if ambiguous_deaths != set() : - print('We will solve ambiguous deaths matching.') - - # now, find common births with the normal bars - - # First case: non-ambiguous matching - # in this case, deaths in X_Z and Y_Z are unique so matching can be made without ambiguity (even if duplicate deaths in X or in Y) - - for death in common_deaths.difference(ambiguous_deaths) : - oXZ = deaths_X_Z.index(death) - oYZ = deaths_Y_Z.index(death) - birth_X = births_X_Z[oXZ] - birth_Y = births_Y_Z[oYZ] - - # Now we match with the persistence bars of X and Y - Occ_X = find_occurences_list(births_X, birth_X) - Occ_Y = find_occurences_list(births_Y, birth_Y) - - if len(Occ_X) == 1 and len(Occ_Y) == 1: # if there are no ambiguous births - a = births_X.index(birth_X) # unique - b = births_Y.index(birth_Y) # unique - - matched_X_Y += [[a, b]] - affinity = compute_affinity(birth_X, deaths_X[a], death, birth_Y, deaths_Y[b], affinity_method = affinity_method) - affinity_X_Y += [affinity] - else: - # the way we are computing persistent homology, the indices of the - # persistent homology of Y can be compared with the indices in the - # image - persistent homology of Y inside Z - - for k, oX in enumerate(Occ_X): - pos_index_X = indices_X[dim][oX][0] - pos_index_XZ = indices_X_Z[dim][oXZ][0] - if pos_index_X == pos_index_XZ: - a = oX - for l, oY in enumerate(Occ_Y): - pos_index_Y = indices_Y[dim][oY][0] - pos_index_YZ = indices_Y_Z[dim][oYZ][0] # the bars are presented in the same order, independently on how we arrange X an Y - if pos_index_Y == pos_index_YZ: - b = oY - matched_X_Y += [[a, b]] - affinity = compute_affinity(birth_X, deaths_X[a], death, birth_Y, deaths_Y[b], affinity_method = affinity_method) - affinity_X_Y += [affinity] - - # Second case: ambiguous matching - - for death in ambiguous_deaths : - # detect the indices of the bars with ambiguous death times - Occ_XZ = find_occurences_list(deaths_X_Z, death) - Occ_YZ = find_occurences_list(deaths_Y_Z, death) - - for i, oXZ in enumerate(Occ_XZ) : - # extract negative index of the image persistence bar of X - neg_index_XZ = indices_X_Z[dim][oXZ][1] - for j, oYZ in enumerate(Occ_YZ) : - # extract negative index of the image persistence bar of Y - neg_index_YZ = indices_Y_Z[dim][oYZ][1] - if neg_index_XZ == neg_index_YZ: - # match the image bars when the indices coincide - birth_X = births_X_Z[oXZ] # ! not i - birth_Y = births_Y_Z[oYZ] # ! not j - - # Now we match with the persistence bars of X and Y - Occ_X = find_occurences_list(births_X, birth_X) - Occ_Y = find_occurences_list(births_Y, birth_Y) - - if len(Occ_X) == 1 and len(Occ_Y) == 1: # if there are no ambiguous births - a = births_X.index(birth_X) # unique - b = births_Y.index(birth_Y) # unique - matched_X_Y += [[a, b]] - affinity = compute_affinity(birth_X, deaths_X[a], death, birth_Y, deaths_Y[b], affinity_method = affinity_method) - affinity_X_Y += [affinity] - else: - # the way we are computing persistent homology, the indices of the - # persistent homology of Y can be compared with the indices in the - # image - persistent homology of Y inside Z - for k, oX in enumerate(Occ_X): - pos_index_X = indices_X[dim][oX][0] - pos_index_XZ = indices_X_Z[dim][oXZ][0] - if pos_index_X == pos_index_XZ: - a = oX - for l, oY in enumerate(Occ_Y): - pos_index_Y = indices_Y[dim][oY][0] - pos_index_YZ = indices_Y_Z[dim][oYZ][0] # the bars are presented in the same order, independently on how we arrange X an Y - if pos_index_Y == pos_index_YZ: - b = oY - matched_X_Y += [[a, b]] - affinity = compute_affinity(birth_X, deaths_X[a], death, birth_Y, deaths_Y[b], affinity_method = affinity_method) - affinity_X_Y += [affinity] - - return matched_X_Y, affinity_X_Y - - -def create_matrices_image(X, Y, filename_X = 'X', filename_Y = 'Y', filename_Z = 'Z', return_thr = False): - '''Function to create the matrices for the computation of image-persistence so that we can compare the indices of the persistence pairs''' - Z = np.vstack((X,Y)) - nb_X = len(X) - ldm_file_X = '{}.lower_distance_matrix'.format(filename_X) - ldm_file_Y = '{}.lower_distance_matrix'.format(filename_Y) - ldm_file_Z = '{}.lower_distance_matrix'.format(filename_Z) - - pairwise_Z = np.sqrt( np.sum( (Z[:, None, :] - Z[None, :, :])**2, axis=-1) ) - maxi = np.max(pairwise_Z) - pairwise_X = pairwise_Z.copy() - pairwise_X[nb_X:] = 2 * maxi + 1 # add min offset 1 because maxi can be small - pairwise_Y = pairwise_Z.copy() - pairwise_Y[:,:nb_X] = 2 * maxi + 1 #observe here the change wrt the previous line - - # so that later we can apply thresholding in Ripser-image (bug fixed) - # threshold = 2 * maxi # not 2 * maxi - 1 as maxi could be very small - - f = open(ldm_file_X, "w") # erase possibly pre-existing - for i in range(len(pairwise_X)) : - f.write(', '.join([ str(x) for x in pairwise_X[i,:i]])) - f.write('\n') - f.close() - - f = open(ldm_file_Y, "w") # erase possibly pre-existing - for i in range(len(pairwise_Y)) : - f.write(', '.join([ str(x) for x in pairwise_Y[i,:i]])) - f.write('\n') - f.close() - - f = open(ldm_file_Z, "w") # erase possibly pre-existing - for i in range(len(pairwise_Z)) : - f.write(', '.join([ str(x) for x in pairwise_Z[i,:i]])) - f.write('\n') - f.close() - - if return_thr : - threshold = 2 * maxi + 0.5 # to make sure we still include people <= maxi (in case maxi = 0... paranoia lol) - return ldm_file_X, ldm_file_Y, ldm_file_Z, threshold - - return ldm_file_X, ldm_file_Y, ldm_file_Z - - -def matching(X,Y, dim = 1, verbose_figs = False, affinity_method = 'A', check_Morse = False) : - - '''Function that takes as input two pointclouds X and Y and computes the relevant barcodes and the matching.''' - - # Compute the lower distance matrices - ldm_file_X, ldm_file_Y, ldm_file_Z, threshold = \ - create_matrices_image(X, Y, filename_X = 'X', filename_Y = 'Y', return_thr = True) - - # Image persistence - apply thresholding (bug fixed) - out_X_Z = compute_image_bars(filename_X = ldm_file_X, filename_Z = ldm_file_Z, threshold = threshold) - bars_X_Z, indices_X_Z = extract_bars_indices(out_X_Z, only_dim_1 = True) - - out_Y_Z = compute_image_bars(filename_X = ldm_file_Y, filename_Z = ldm_file_Z, threshold = threshold) - bars_Y_Z, indices_Y_Z = extract_bars_indices(out_Y_Z, only_dim_1 = True) - - # Persistent homology - out_X = compute_bars_tightreps(inp = None, filename = ldm_file_X) - # this way we obtain the same bars but the vertices are indexes in the same way as in the image persistence - bars_X, reps_X, tight_reps_X, indices_X = extract_bars_reps_indices(out_X, only_dim_1 = True) - - out_Y = compute_bars_tightreps(inp = None, filename = ldm_file_Y) - bars_Y, reps_Y, tight_reps_Y, indices_Y = extract_bars_reps_indices(out_Y, only_dim_1 = True) - - matched_X_Y, affinity_X_Y = find_match(bars_X, bars_X_Z, indices_X, indices_X_Z, - bars_Y, bars_Y_Z, indices_Y, indices_Y_Z, dim = 1, - affinity_method = affinity_method, check_Morse = check_Morse, - check_ambiguous_deaths = False) - - if verbose_figs : - show_matches(X,Y,matched_X_Y, affinity_X_Y, tight_reps_X, tight_reps_Y, dim = dim, show_together = True) - - return matched_X_Y, affinity_X_Y, (bars_X, reps_X, tight_reps_X), (bars_Y, reps_Y, tight_reps_Y) - - -###### PREVALENCE, CROSS-PREVALENCE - -def multiple_matching(X, list_Y, dim = 1, verbose_figs = False, affinity_method = 'A') : - - out_X = compute_bars_tightreps(X) - bars_X, reps_X, tight_reps_X, indices_X = extract_bars_reps_indices(out_X, only_dim_1 = True) - - list_matched_X_Y = {} - list_affinity_X_Y = {} - - list_bars_reps_Y = [] - - for y, Y in enumerate(list_Y) : - print('Matching X to Y_{} ...'.format(y)) - - # Compute the lower distance matrices - ldm_file_X, ldm_file_Y, ldm_file_Z, threshold = \ - create_matrices_image(X, Y, filename_X = 'X', filename_Y = 'Y', return_thr = True) - - # Image persistence - apply thresholding - out_X_Z = compute_image_bars(filename_X = ldm_file_X, filename_Z = ldm_file_Z, threshold = threshold) - bars_X_Z, indices_X_Z = extract_bars_indices(out_X_Z, only_dim_1 = True) - - out_Y_Z = compute_image_bars(filename_X = ldm_file_Y, filename_Z = ldm_file_Z, threshold = threshold) - bars_Y_Z, indices_Y_Z = extract_bars_indices(out_Y_Z, only_dim_1 = True) - - out_Y = compute_bars_tightreps(Y) - bars_Y, reps_Y, tight_reps_Y, indices_Y = extract_bars_reps_indices(out_Y, only_dim_1 = True) - - list_bars_reps_Y += [ [bars_Y, reps_Y, tight_reps_Y] ] - - matched_X_Y, affinity_X_Y = find_match(bars_X, bars_X_Z, indices_X, indices_X_Z, - bars_Y, bars_Y_Z, indices_Y, indices_Y_Z, dim = 1, - affinity_method = affinity_method, check_Morse = False, - check_ambiguous_deaths = False) - - list_matched_X_Y[y] = matched_X_Y - list_affinity_X_Y[y] = affinity_X_Y - - if verbose_figs : - - show_matches(X, list_Y[y], matched_X_Y[y], affinity_X_Y[y], tight_reps_X, tight_reps_Y, dim = dim, show_together = True) - - list_matched_X_Y = list(list_matched_X_Y.values()) - list_affinity_X_Y = list(list_affinity_X_Y.values()) - bars_reps_X = [bars_X, reps_X, tight_reps_X] - - return list_matched_X_Y, list_affinity_X_Y, bars_reps_X, list_bars_reps_Y - -def cross_matching(list_X, dim = 1, verbose_figs = False, affinity_method = 'A') : - - list_matched_X_Y = {} - list_affinity_X_Y = {} - list_bars_reps_indices_X = {} - list_bars_reps_X = {} - - # compute PH and reps and indices of individual spaces - for i,X in enumerate(list_X) : - out_X = compute_bars_tightreps(X) - bars_X, reps_X, tight_reps_X, indices_X = extract_bars_reps_indices(out_X, only_dim_1 = True) - list_bars_reps_indices_X[i] = [bars_X, reps_X, tight_reps_X, indices_X] - list_bars_reps_X[i] = [bars_X, reps_X, tight_reps_X] - - # match any X_i to any X_j (j > i) - for i,X in enumerate(list_X) : - for j in range(i+1, len(list_X)) : - print('Matching X_{} to X_{} ...'.format(i,j)) - - X = list_X[i] - Y = list_X[j] - - # Compute the lower distance matrices - ldm_file_X, ldm_file_Y, ldm_file_Z, threshold = \ - create_matrices_image(X, Y, filename_X = 'X', filename_Y = 'Y', return_thr = True) - - # Image persistence - apply thresholding (bug fixed) - out_X_Z = compute_image_bars(filename_X = ldm_file_X, filename_Z = ldm_file_Z, threshold = threshold) - bars_X_Z, indices_X_Z = extract_bars_indices(out_X_Z, only_dim_1 = True) - - out_Y_Z = compute_image_bars(filename_X = ldm_file_Y, filename_Z = ldm_file_Z, threshold = threshold) - bars_Y_Z, indices_Y_Z = extract_bars_indices(out_Y_Z, only_dim_1 = True) - - bars_X, reps_X, tight_reps_X, indices_X = list_bars_reps_indices_X[i] - bars_Y, reps_Y, tight_reps_Y, indices_Y = list_bars_reps_indices_X[j] - - - matched_X_Y, affinity_X_Y = find_match(bars_X, bars_X_Z, indices_X, indices_X_Z, - bars_Y, bars_Y_Z, indices_Y, indices_Y_Z, dim = 1, - affinity_method = affinity_method, check_Morse = False, - check_ambiguous_deaths = False) - - list_matched_X_Y[i,j] = matched_X_Y - list_affinity_X_Y[i,j] = affinity_X_Y - - for i in range(len(list_X)) : - for j in range(i) : - aa = list_matched_X_Y[j,i] - if len(aa) > 0 : - list_matched_X_Y[i,j] = np.array(aa)[:,::-1].tolist() # reverse column order - else : - list_matched_X_Y[i,j] = [] - list_affinity_X_Y[i,j] = list_affinity_X_Y[j,i] - return list_matched_X_Y, list_affinity_X_Y, list_bars_reps_X - -##### SOME FUNCTIONS TO ENABLE TRACKING CYCLES - -def track_cycles_from_slice(list_matched_X_Y, list_affinity_X_Y, cycle, list_indices, initial_slice = 0): - ''' From a list of matched cycles between consecutive slices, obtains a list in which we store the matches - that track a particular cycle from a some chosen slice - output = [[cycle, a],[a, b], [b, c] ...] - rmk: set of indices does not include the initial slice, counts from the second slice studied''' - - tracked_cycle = [] - tracked_affinity = [] - current_match = [] - - # initialise - for i, match in enumerate(list_matched_X_Y[initial_slice]): - if match[0] == cycle: - tracked_cycle += [match] - current_match = match - tracked_affinity.append(list_affinity_X_Y[initial_slice][i]) - - - # track the cycle - for i in list_indices: - next_cycle = current_match[1] - tracked_copy = tracked_cycle.copy() - for j, match in enumerate(list_matched_X_Y[i]): - if match[0] == next_cycle: - #print(match) - tracked_cycle += [match] - current_match = match - tracked_affinity.append(list_affinity_X_Y[i][j]) - continue - if len(tracked_copy) == len(tracked_cycle): # in case we don't find a next cycle we stop tracking - - break - - return tracked_cycle, tracked_affinity diff --git a/src/scloop/plotting/_hodge.py b/src/scloop/plotting/_hodge.py index 65a6a33..c58a65b 100644 --- a/src/scloop/plotting/_hodge.py +++ b/src/scloop/plotting/_hodge.py @@ -219,7 +219,9 @@ def loop_edge_overlay( assert edge_embeddings is not None assert loop_class.coordinates_edges is not None - for rep_idx, edge_coords_raw in enumerate(loop_class.coordinates_edges): + for rep_idx, edge_coords_raw in loop_class.filter_valid( + list(enumerate(loop_class.coordinates_edges)) + ): valid_indices = loop_class.valid_edge_indices_per_rep[rep_idx] if not valid_indices: continue diff --git a/src/scloop/plotting/_homology.py b/src/scloop/plotting/_homology.py index 78eb1c2..75ae896 100644 --- a/src/scloop/plotting/_homology.py +++ b/src/scloop/plotting/_homology.py @@ -450,6 +450,7 @@ def loops( components: tuple[Index_t, Index_t] | list[Index_t] = (0, 1), ax: Axes | None = None, *, + use_refined: bool = False, pointsize: PositiveFloat = 1, figsize: tuple[PositiveFloat, PositiveFloat] = DEFAULT_FIGSIZE, dpi: PositiveFloat = DEFAULT_DPI, @@ -529,6 +530,7 @@ def _loops_for_selector( selector=selector, include_bootstrap=False, embedding_alt=emb, + use_refined=use_refined, ) else: return data._get_loop_embedding( @@ -536,6 +538,7 @@ def _loops_for_selector( include_bootstrap=True, embedding_alt=emb, keep_matches=keep_matches, + use_refined=use_refined, ) case IdxLoopInTrack(idx_track, idx_loop): return data._get_loop_embedding( @@ -544,6 +547,7 @@ def _loops_for_selector( idx_loop=idx_loop, embedding_alt=emb, keep_matches=keep_matches, + use_refined=use_refined, ) case IdxLoopInClassOrig(idx_loop_class, idx_loop): return data._get_loop_embedding( @@ -551,17 +555,20 @@ def _loops_for_selector( include_bootstrap=False, idx_loop=idx_loop, embedding_alt=emb, + use_refined=use_refined, ) case IdxClassInBootstrap(idx_bootstrap, idx_loop_class): return data._get_loop_embedding( selector=(idx_bootstrap, idx_loop_class), embedding_alt=emb, + use_refined=use_refined, ) case IdxLoopInClassInBootstrap(idx_bootstrap, idx_loop_class, idx_loop): return data._get_loop_embedding( selector=(idx_bootstrap, idx_loop_class), idx_loop=idx_loop, embedding_alt=emb, + use_refined=use_refined, ) except AssertionError: return [] diff --git a/src/scloop/tools/_loops.py b/src/scloop/tools/_loops.py index 1b35915..4bb67cf 100644 --- a/src/scloop/tools/_loops.py +++ b/src/scloop/tools/_loops.py @@ -2,7 +2,7 @@ from __future__ import annotations from contextlib import nullcontext -from typing import Annotated, Any, Literal +from typing import Annotated, Any import numpy as np from anndata import AnnData @@ -12,6 +12,7 @@ from ..data.constants import ( DEFAULT_AUTO_THRESHOLD_FACTOR, + DEFAULT_CANDIDATE_METHOD, DEFAULT_K_NEIGHBORS_CHECK_EQUIVALENCE, DEFAULT_MAX_ROWS_BOUNDARY_MATRIX, DEFAULT_MAXITER_EIGENDECOMPOSITION, @@ -19,6 +20,7 @@ DEFAULT_N_HODGE_COMPONENTS, DEFAULT_N_MAX_WORKERS, DEFAULT_N_NEIGHBORS_EDGE_EMBEDDING, + DEFAULT_PRESENCE_METHOD, DEFAULT_TIMEOUT_EIGENDECOMPOSITION, SCLOOP_META_UNS_KEY, SCLOOP_UNS_KEY, @@ -26,10 +28,12 @@ from ..data.containers import HomologyData from ..data.metadata import ScloopMeta from ..data.types import ( + CandidateMethod, Index_t, NonZeroCount_t, Percent_t, PositiveFloat, + PresenceTestMethod, Size_t, ) from ..preprocessing.downsample import sample @@ -66,7 +70,8 @@ def find_loops( bootstrap_fps_alpha: float = 1.0, bootstrap_herding_n_features: int = 1000, bootstrap_herding_seed: int | None = None, - bootstrap_candidate_method: Literal["geometric", "image"] = "geometric", + bootstrap_candidate_method: CandidateMethod = DEFAULT_CANDIDATE_METHOD, + bootstrap_presence_test_method: PresenceTestMethod = DEFAULT_PRESENCE_METHOD, require_bootstrap_homological_equivalence: bool = True, n_check_per_candidate: NonZeroCount_t = 1, max_rows_boundary_matrix: NonZeroCount_t = DEFAULT_MAX_ROWS_BOUNDARY_MATRIX, @@ -79,6 +84,7 @@ def find_loops( kwargs_bootstrap: dict[str, Any] | None = None, kwargs_loop_test: dict[str, Any] | None = None, kwargs_loop_representatives: dict[str, Any] | None = None, + kwargs_loop_refinement: dict[str, Any] | None = None, ) -> None: use_log_display = verbose and max_log_messages is not None if verbose: @@ -229,6 +235,14 @@ def find_loops( "require_homological_equivalence", require_bootstrap_homological_equivalence, ) + loop_test_kwargs = dict(kwargs_loop_test or {}) + presence_test_method = loop_test_kwargs.pop( + "presence_test_method", bootstrap_presence_test_method + ) + # only the wilcoxon presence test needs coherence score + compute_homotopy_coherence = bootstrap_kwargs.pop( + "compute_homotopy_coherence", presence_test_method == "wilcoxon" + ) if verbose: logger.info(f"Bootstrap validation: {n_bootstrap} resamples") hd._bootstrap( @@ -248,12 +262,25 @@ def find_loops( bootstrap_herding_seed=bootstrap_herding_seed, candidate_method=bootstrap_candidate_method, require_homological_equivalence=(require_bootstrap_homological_equivalence), + compute_homotopy_coherence=compute_homotopy_coherence, verbose=verbose, progress_main=progress_main, use_log_display=use_log_display, use_parallel=use_parallel, **bootstrap_kwargs, ) + """ + ========= loop refinement ========= + - refine loops against full data + =================================== + """ + if verbose: + logger.info("Refining loop representatives on full data") + hd._refine_loop_representatives( + embedding=embedding, + include_bootstrap=True, + **(kwargs_loop_refinement or {}), + ) """ ========= statistcal tests ========= @@ -264,7 +291,7 @@ def find_loops( assert hd.bootstrap_data is not None if verbose: logger.info("Statistical testing (presence + persistence)") - hd._test_loops(**(kwargs_loop_test or {})) + hd._test_loops(presence_test_method=presence_test_method, **loop_test_kwargs) if verbose: presence = hd.bootstrap_data.presence_test_result if presence is not None: diff --git a/src/scloop/utils/logging.py b/src/scloop/utils/logging.py index 005e739..d3b3256 100644 --- a/src/scloop/utils/logging.py +++ b/src/scloop/utils/logging.py @@ -453,8 +453,10 @@ def __enter__(self) -> LogCache: self._in_jupyter = _is_in_jupyter() console = self.console - if not self._in_jupyter and console is not None and ( - not console.is_terminal or console.is_dumb_terminal + if ( + not self._in_jupyter + and console is not None + and (not console.is_terminal or console.is_dumb_terminal) ): return self.cache