1414
1515def get_lr_expr (
1616 adata : sc .AnnData ,
17- lr : pd .DataFrame ,
17+ lr_df : pd .DataFrame ,
1818 radius : float ,
1919 aggr : str = "mean" ,
2020 binary : bool = False ,
21+ threshold : float = 0.01 ,
2122) -> tuple [torch .Tensor , pd .DataFrame ]:
2223 r"""
2324 Get ligand-receptor activity from spatial data.
@@ -26,23 +27,24 @@ def get_lr_expr(
2627
2728 Args:
2829 adata: AnnData object
29- lr : Ligand receptor pairs, must contain columns `ligand.symbol` and `receptor.symbol`
30+ lr_df : Ligand receptor pairs, must contain columns `ligand.symbol` and `receptor.symbol`
3031 radius: float radius for valid ligand
3132 aggr: str aggregation method of ligand expression, default is `mean`
3233 binary: bool if convert the expression to binary, default is `False`
34+ threshold: float threshold to filter low-expressed LR pairs, default is `0.01`
3335 Returns:
3436 LR activity: Ligand-receptor pair activity
35- lr : filtered ligand receptor pairs
37+ lr_df : filtered ligand receptor pairs
3638 Note:
3739 Not support batched data yet.
3840 """
3941 # check
40- assert "ligand.symbol" in lr .columns
41- assert "receptor.symbol" in lr .columns
42+ assert "ligand.symbol" in lr_df .columns
43+ assert "receptor.symbol" in lr_df .columns
4244
4345 # filter non-LR genes (reduce memory usage)
4446 lr_genes = []
45- for _ , row in lr .iterrows ():
47+ for _ , row in lr_df .iterrows ():
4648 ligand = row ["ligand.symbol" ]
4749 if "," in ligand :
4850 lr_genes += ligand .split ("," )
@@ -57,20 +59,21 @@ def get_lr_expr(
5759 lr_genes = list (set (lr_genes ))
5860 lr_idx = [x in lr_genes for x in adata .var .index ]
5961 adata = adata [:, lr_idx ]
60- sc . pp . filter_genes ( adata , min_cells = 30 )
62+
6163 logger .info (f"Find { adata .n_vars } LR-related genes." )
6264
6365 # merge neighbor expr
6466 edge_index = build_graph (adata .obsm ["spatial" ], radius = radius , mode = "radius" )
6567 aggr_net = SimpleConv (aggr = aggr , combine_root = None )
6668 expr = torch .from_numpy (adata .X .toarray ()) if isspmatrix (adata .X ) else torch .from_numpy (adata .X )
69+ # expr = expr.to(torch.float32)
6770 expr_neighbors = aggr_net (expr , edge_index )
6871
6972 # get the ligand and receptor gene index in adata.var
7073 lr_activity_list = []
7174 gene_set = set (adata .var .index )
72- lr_filter = np .zeros (len (lr ), dtype = bool )
73- for i , lr_pair in lr .iterrows ():
75+ lr_filter = np .zeros (len (lr_df ), dtype = bool )
76+ for i , lr_pair in lr_df .iterrows ():
7477 ligand = lr_pair ["ligand.symbol" ]
7578 receptor = lr_pair ["receptor.symbol" ]
7679 ligand = [ligand ] if "," not in ligand else ligand .split ("," )
@@ -82,10 +85,16 @@ def get_lr_expr(
8285 receptor_idx = [adata .var .index .get_loc (x ) for x in receptor if x in gene_set ]
8386 ligand_expr = expr_neighbors [:, ligand_idx ].mean (dim = 1 ) # from neighbor
8487 receptor_expr = expr [:, receptor_idx ].mean (dim = 1 ) # from cell
88+ lr_activity = ligand_expr * receptor_expr
89+ if (lr_activity > 0 ).float ().mean ().item () < threshold :
90+ continue
8591 lr_activity_list .append (ligand_expr * receptor_expr )
8692 lr_filter [i ] = True
93+ if len (lr_activity_list ) == 0 :
94+ logger .warning (f"No ligand-receptor pairs (threshold: { threshold } )." )
95+ return None , None
8796 lr_activity = torch .stack (lr_activity_list , dim = 1 )
8897 if binary :
8998 lr_activity = (lr_activity > 0 ).float ()
90- lr = lr [lr_filter ]
91- return lr_activity , lr
99+ lr_df = lr_df [lr_filter ]
100+ return lr_activity , lr_df
0 commit comments