11# Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan)
22from abc import ABC , abstractmethod
3+ from concurrent .futures import ThreadPoolExecutor , as_completed
34
45import numpy as np
56from anndata import AnnData
67from pydantic import BaseModel , Field , ValidationInfo , field_validator
78from pydantic .dataclasses import dataclass
89from scipy .sparse import csr_matrix , triu
10+ from scipy .spatial .distance import directed_hausdorff
911
1012from ..computing .homology import (
1113 compute_boundary_matrix_data ,
2224 edge_ids_to_rows ,
2325 extract_edges_from_coo ,
2426 loop_vertices_to_edge_ids ,
27+ nearest_neighbor_per_row ,
2528)
2629
2730
@@ -289,30 +292,74 @@ def _compute_loop_representatives(
289292 else :
290293 bootstrap_data .loop_representatives [idx_bootstrap ][loop_idx ] = loops # type: ignore[attr-defined]
291294
295+ def _assess_bootstrap_geometric_equivalence (
296+ self ,
297+ adata : AnnData ,
298+ source_class_idx : int ,
299+ target_class_idx : int ,
300+ idx_bootstrap : int = 0 ,
301+ ) -> tuple [int , int , float ]:
302+ assert self .loop_representatives is not None
303+ assert self .bootstrap_data is not None
304+ assert self .meta .preprocess .embedding_method is not None
305+
306+ if idx_bootstrap >= len (self .bootstrap_data .loop_representatives ):
307+ return (source_class_idx , target_class_idx , np .nan )
308+ if source_class_idx >= len (self .loop_representatives ):
309+ return (source_class_idx , target_class_idx , np .nan )
310+
311+ boot_loops_all = self .bootstrap_data .loop_representatives [idx_bootstrap ]
312+ if target_class_idx >= len (boot_loops_all ):
313+ return (source_class_idx , target_class_idx , np .nan )
314+
315+ source_loops = self .loop_representatives [source_class_idx ]
316+ target_loops = boot_loops_all [target_class_idx ]
317+
318+ if len (source_loops ) == 0 or len (target_loops ) == 0 :
319+ return (source_class_idx , target_class_idx , np .nan )
320+
321+ emb = adata .obsm [f"X_{ self .meta .preprocess .embedding_method } " ]
322+ distances = []
323+ for source_loop in source_loops :
324+ for target_loop in target_loops :
325+ source_coords = emb [source_loop ]
326+ target_coords = emb [target_loop ]
327+ try :
328+ dist = max (
329+ directed_hausdorff (source_coords , target_coords )[0 ],
330+ directed_hausdorff (target_coords , source_coords )[0 ],
331+ )
332+ distances .append (dist )
333+ except (ValueError , IndexError ):
334+ distances .append (np .nan )
335+
336+ mean_distance = np .nanmean (distances ) if distances else np .nan
337+ return (source_class_idx , target_class_idx , mean_distance )
338+
292339 def _assess_bootstrap_homology_equivalence (
293340 self ,
294341 source_class_idx : int ,
295342 target_class_idx : int | None = None ,
296343 idx_bootstrap : int = 0 ,
297344 n_pairs_check : int = 10 ,
298- ) -> bool :
345+ ) -> tuple [ int , int , bool ] :
299346 assert self .boundary_matrix_d1 is not None
300347 assert self .loop_representatives is not None
301348 assert self .bootstrap_data is not None
302349 if target_class_idx is None :
303350 target_class_idx = source_class_idx
304- if idx_bootstrap >= len (self .bootstrap_data .loop_representatives ): # type: ignore[attr-defined]
305- return False
351+ if idx_bootstrap >= len (self .bootstrap_data .loop_representatives ):
352+ return ( source_class_idx , target_class_idx , False )
306353 if source_class_idx >= len (self .loop_representatives ):
307- return False
308- boot_loops_all = self .bootstrap_data .loop_representatives [idx_bootstrap ] # type: ignore[attr-defined]
354+ return ( source_class_idx , target_class_idx , False )
355+ boot_loops_all = self .bootstrap_data .loop_representatives [idx_bootstrap ]
309356 if target_class_idx >= len (boot_loops_all ):
310- return False
357+ return ( source_class_idx , target_class_idx , False )
311358
312359 source_loops = self .loop_representatives [source_class_idx ]
313360 target_loops = boot_loops_all [target_class_idx ]
314361 if len (source_loops ) == 0 or len (target_loops ) == 0 :
315- return False
362+ return ( source_class_idx , target_class_idx , False )
316363
317364 mask_a = self ._loops_to_edge_mask (source_loops )
318365 mask_b = self ._loops_to_edge_mask (target_loops )
@@ -323,7 +370,7 @@ def _assess_bootstrap_homology_equivalence(
323370 loop_mask_b = mask_b ,
324371 n_pairs_check = n_pairs_check ,
325372 )
326- return any (r == 0 for r in results )
373+ return ( source_class_idx , target_class_idx , any (r == 0 for r in results ) )
327374
328375 def _bootstrap (
329376 self ,
@@ -340,6 +387,8 @@ def _bootstrap(
340387 loop_lower_t_pct : float = 5 ,
341388 loop_upper_t_pct : float = 95 ,
342389 n_pairs_check_equivalence : int = 4 ,
390+ n_max_workers : int = 4 ,
391+ k_neighbors_check_equivalence : int = 3 ,
343392 verbose : bool = True ,
344393 ** nei_kwargs ,
345394 ) -> None :
@@ -374,24 +423,65 @@ def _bootstrap(
374423 loop_upper_t_pct = loop_upper_t_pct ,
375424 )
376425 """
377- ========= geometric matching =========
378- - find loop neighbors using frechet
426+ ============= geometric matching ==== =========
427+ - find loop neighbors using hausdorff/ frechet
379428 - reduce computation load
380- ======================================
429+ ==============================================
381430 """
382- # n_source_
383- # for i in range(len(self.loop_representatives)):
384- # for j in range(len(self.bootstrap_data)):
385-
431+ n_original_loop_classes = len (self .loop_representatives )
432+ n_bootstrap_loop_classes = len (
433+ self .bootstrap_data .loop_representatives [idx_bootstrap ]
434+ )
435+
436+ if n_original_loop_classes == 0 or n_bootstrap_loop_classes == 0 :
437+ continue
438+
439+ pairwise_result_matrix = np .full (
440+ (n_original_loop_classes , n_bootstrap_loop_classes ), np .nan
441+ )
442+
443+ with ThreadPoolExecutor (max_workers = n_max_workers ) as executor :
444+ tasks = {}
445+ for i in range (n_original_loop_classes ):
446+ for j in range (n_bootstrap_loop_classes ):
447+ task = executor .submit (
448+ self ._assess_bootstrap_geometric_equivalence ,
449+ i ,
450+ j ,
451+ idx_bootstrap ,
452+ )
453+ tasks [task ] = (i , j )
454+
455+ for task in as_completed (tasks ):
456+ src_idx , tgt_idx , distance = task .result ()
457+ pairwise_result_matrix [src_idx , tgt_idx ] = distance
458+
459+ neighbor_indices , neighbor_distances = nearest_neighbor_per_row (
460+ pairwise_result_matrix , k_neighbors_check_equivalence
461+ )
386462
387463 """
388464 ========= homological matching =========
389465 - gf2 regression
390466 ========================================
391467 """
392- self ._assess_bootstrap_homology_equivalence (
393- source_class_idx = 0 ,
394- target_class_idx = 0 ,
395- idx_bootstrap = idx_bootstrap ,
396- n_pairs_check = n_pairs_check_equivalence ,
397- )
468+ with ThreadPoolExecutor (max_workers = n_max_workers ) as executor :
469+ tasks = {}
470+ for si in range (n_original_loop_classes ):
471+ for k in range (k_neighbors_check_equivalence ):
472+ tj = neighbor_indices [si , k ]
473+ if tj >= 0 :
474+ task = executor .submit (
475+ self ._assess_bootstrap_homology_equivalence ,
476+ si ,
477+ tj ,
478+ idx_bootstrap ,
479+ n_pairs_check_equivalence ,
480+ )
481+ tasks [task ] = (si , tj , neighbor_distances [si , k ])
482+
483+ for task in as_completed (tasks ):
484+ si , tj , geom_dist = tasks [task ]
485+ _ , _ , is_equivalent = task .result ()
486+ if is_equivalent :
487+ pass
0 commit comments