multineuronchat 2025.11.10.dev0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,959 @@
1
+ import os
2
+
3
+ import pickle
4
+
5
+ import numpy as np
6
+ import xarray as xr
7
+
8
+ from scipy import stats
9
+ from scipy.stats import PermutationMethod
10
+
11
+ from .InteractionDB import InteractionDB
12
+
13
+ from .utils import gene_filter_and_subject_wise_normalize_dataset
14
+ from .MultiNeuronChat import compute_avg_expression, compute_communication_score_matrix
15
+
16
+ from typing import Tuple, Optional, List, Dict, Set, Any, Union
17
+
18
+
19
+ class MultiNeuronChatObject:
20
+ def __init__(self,
21
+ condition_label_column: str,
22
+ condition_names: Tuple[str, str],
23
+ subject_label_column: str,
24
+ cell_type_label_column: str,
25
+ db: str,
26
+ interaction_db: Optional[InteractionDB] = None):
27
+ if (db not in ['human', 'mouse', 'human_extended']) and (not os.path.exists(db)):
28
+ raise ValueError('db must be either "human", "human_extended", or "mouse" '
29
+ 'or a valid path to an interaction database')
30
+
31
+ self.condition_label_column: str = condition_label_column
32
+ self.condition_names: Tuple[str, str] = condition_names
33
+
34
+ self.subject_label_column: str = subject_label_column
35
+ self.cell_type_label_column: str = cell_type_label_column
36
+
37
+ self.db: str = db
38
+
39
+ if interaction_db is None:
40
+ self.interaction_db: InteractionDB = InteractionDB(db=db)
41
+ else:
42
+ self.interaction_db: InteractionDB = interaction_db
43
+ self.gene_set: Set[str] = self.interaction_db.get_set_of_genes()
44
+
45
+ self.mean_type: Optional[str] = None
46
+ self.trim_mean_fraction: Optional[float] = None
47
+
48
+ # Info about the data
49
+ self.source_cell_types: List[str] = []
50
+ self.receiver_cell_types: List[str] = []
51
+ self.interaction_names: List[str] = []
52
+ self.__n_cell_types: int = 0
53
+ self.__n_interactions: int = 0
54
+
55
+ self.avg_expression_per_condition_and_subject_dict: Optional[Dict[str, Dict[str, xr.DataArray]]] = None
56
+
57
+ self.communication_scores_per_condition_and_subject_dict: Optional[Dict[str, Dict[str, xr.DataArray]]] = None
58
+ self.ligand_abundance_per_condition_and_subject_dict: Optional[Dict[str, Dict[str, xr.DataArray]]] = None
59
+ self.target_abundance_per_condition_and_subject_dict: Optional[Dict[str, Dict[str, xr.DataArray]]] = None
60
+
61
+ self.communication_score_dict: Optional[Dict[str, xr.DataArray]] = None
62
+ self.ligand_abundance_dict: Optional[Dict[str, xr.DataArray]] = None
63
+ self.target_abundance_dict: Optional[Dict[str, xr.DataArray]] = None
64
+
65
+ self.p_values: Optional[Dict[str, xr.DataArray]] = None
66
+ self.p_values_adj: Optional[Dict[str, xr.DataArray]] = None
67
+ self.statistics: Optional[Dict[str, Dict[str, xr.DataArray]]] = None
68
+
69
+ def compute_communication_scores(
70
+ self,
71
+ path_to_data_loom: str,
72
+ path_to_subject_wise_max_normalized: Optional[str] = None,
73
+ chunk_size: Optional[int] = 1024,
74
+ n_processes: Optional[int] = 1,
75
+ min_n_cells_threshold: Optional[int] = -1,
76
+ mean_type: str = 'tri_mean',
77
+ gene_label_row: str = 'Gene',
78
+ trim_mean_fraction: Optional[float] = None,
79
+ verbose: Optional[bool] = False
80
+ ):
81
+ """
82
+ Compute the communication scores for each subject and condition.
83
+
84
+ :param path_to_data_loom: path to the loom file containing the data
85
+ :param path_to_subject_wise_max_normalized: path to the loom file containing the subject wise max normalized data. If this is not provided, it will be generated automatically. If the file already exists the subject wise max normalization will not be recomputed.
86
+ :param chunk_size: chunk size to use for processing the loom file (can be optimized for memory efficiency)
87
+ :param n_processes: number of processes to use for parallelization
88
+ :param min_n_cells_threshold: minimum number of cells required within a donor for a cell-type to be considered
89
+ :param mean_type: type of mean to use for averaging the expression values ('mean', 'tri_mean', or 'trim_mean')
90
+ :param gene_label_row: name of the row in the loom file that contains the gene labels
91
+ :param trim_mean_fraction: fraction to trim when using the trim mean (only required if mean_type is 'trim_mean')
92
+ :param verbose: whether to print progress messages
93
+
94
+ :return: None
95
+ """
96
+ if not os.path.exists(path_to_data_loom):
97
+ raise FileNotFoundError(f"File {path_to_data_loom} not found")
98
+
99
+ if mean_type not in ['mean', 'tri_mean', 'trim_mean']:
100
+ raise ValueError('mean_type must be either "mean", "tri_mean", or "trim_mean"')
101
+ if mean_type == 'trim_mean' and trim_mean_fraction is None:
102
+ raise ValueError('If mean_type is "trim_mean", trim_mean_fraction must be specified')
103
+
104
+ # If the path to the subject wise max normalized data is not provided, generate it
105
+ if path_to_subject_wise_max_normalized is None:
106
+ path_to_subject_wise_max_normalized = path_to_data_loom.replace('.loom', '_MultiNeuronChat.loom')
107
+
108
+ # if path_to_subject_wise_max_normalized does not exist -> generate it
109
+ if not os.path.exists(path_to_subject_wise_max_normalized):
110
+ # Filter genes and max normalize each subject
111
+ gene_filter_and_subject_wise_normalize_dataset(
112
+ path_to_loom=path_to_data_loom,
113
+ gene_set=self.gene_set,
114
+ subject_label_column=self.subject_label_column,
115
+ path_to_normalized_loom=path_to_subject_wise_max_normalized,
116
+ gene_label_row=gene_label_row,
117
+ chunk_size=chunk_size,
118
+ verbose=verbose
119
+ )
120
+
121
+ # Compute avg expression for each subject
122
+ self.avg_expression_per_condition_and_subject_dict: Dict[str, Dict[str, xr.DataArray]] = compute_avg_expression(
123
+ path_to_loom=path_to_subject_wise_max_normalized,
124
+ condition_label_column=self.condition_label_column,
125
+ condition_label_a=self.condition_names[0],
126
+ condition_label_b=self.condition_names[1],
127
+
128
+ subject_label_column=self.subject_label_column,
129
+ cell_type_label_column=self.cell_type_label_column,
130
+
131
+ min_n_cells_threshold=min_n_cells_threshold,
132
+
133
+ mean_type=mean_type,
134
+ trim_mean_fraction=trim_mean_fraction,
135
+
136
+ gene_label_row=gene_label_row,
137
+
138
+ n_processes=n_processes
139
+ )
140
+
141
+ # Compute interaction score matrix for each subject
142
+ res: Tuple[Dict[str, Dict[str, xr.DataArray]], Dict[str, Dict[str, xr.DataArray]], Dict[str, Dict[str, xr.DataArray]]] = compute_communication_score_matrix(
143
+ avg_expression_dict=self.avg_expression_per_condition_and_subject_dict,
144
+ interaction_db=self.interaction_db,
145
+ n_processes=n_processes
146
+ )
147
+ self.communication_scores_per_condition_and_subject_dict = res[0]
148
+ self.ligand_abundance_per_condition_and_subject_dict = res[1]
149
+ self.target_abundance_per_condition_and_subject_dict = res[2]
150
+
151
+ # TODO this is super ugly...
152
+ # Set cell_types and interaction_names
153
+ tmp_condition_dict = self.communication_scores_per_condition_and_subject_dict[self.condition_names[0]]
154
+ tmp_subject = list(tmp_condition_dict.keys())[0]
155
+ self.source_cell_types, self.receiver_cell_types, self.interaction_names = tmp_condition_dict[
156
+ tmp_subject].coords.values()
157
+ self.source_cell_types = self.source_cell_types.to_numpy().tolist()
158
+ self.receiver_cell_types = self.receiver_cell_types.to_numpy().tolist()
159
+ self.interaction_names = self.interaction_names.to_numpy().tolist()
160
+
161
+ self.__n_cell_types = len(self.source_cell_types)
162
+ self.__n_interactions = len(self.interaction_names)
163
+
164
+ # Combine the distributions
165
+ res: Tuple[Dict[str, xr.DataArray], Dict[str, xr.DataArray], Dict[str, xr.DataArray]] = self.__combine_distributions()
166
+ self.communication_score_dict = res[0]
167
+ self.ligand_abundance_dict = res[1]
168
+ self.target_abundance_dict = res[2]
169
+
170
+ def __combine_distributions(self) -> Tuple[Dict[str, xr.DataArray], Dict[str, xr.DataArray], Dict[str, xr.DataArray]]:
171
+ """
172
+ Combine the communication scores, ligand abundance, and target abundance for each condition into a single xarray DataArray.
173
+
174
+ :return: Tuple of three dictionaries containing the combined communication scores, ligand abundance, and target abundance for each condition
175
+ """
176
+ combined_communication_score_dict: Dict[str, xr.DataArray] = {}
177
+ combined_ligand_abundance_dict: Dict[str, xr.DataArray] = {}
178
+ combined_target_abundance_dict: Dict[str, xr.DataArray] = {}
179
+
180
+ for condition in self.communication_scores_per_condition_and_subject_dict.keys():
181
+ subject_dict_communication_scores: Dict[str, xr.DataArray] = self.communication_scores_per_condition_and_subject_dict[condition]
182
+ subject_dict_ligand_abundance: Dict[str, xr.DataArray] = self.ligand_abundance_per_condition_and_subject_dict[condition]
183
+ subject_dict_target_abundance: Dict[str, xr.DataArray] = self.target_abundance_per_condition_and_subject_dict[condition]
184
+
185
+ subject_list: List[str] = list(subject_dict_communication_scores.keys())
186
+ n_subjects: int = len(subject_list)
187
+
188
+ condition_communication_score_np: np.array = np.zeros(shape=(n_subjects, self.__n_cell_types, self.__n_cell_types, self.__n_interactions))
189
+ condition_ligand_abundance_np: np.array = np.zeros(shape=(n_subjects, self.__n_cell_types, self.__n_interactions))
190
+ condition_target_abundance_np: np.array = np.zeros(shape=(n_subjects, self.__n_cell_types, self.__n_interactions))
191
+
192
+ for i, subject in enumerate(subject_list):
193
+ condition_communication_score_np[i, :, :, :] = subject_dict_communication_scores[subject]
194
+ condition_ligand_abundance_np[i, :, :] = subject_dict_ligand_abundance[subject]
195
+ condition_target_abundance_np[i, :, :] = subject_dict_target_abundance[subject]
196
+
197
+ condition_array: xr.DataArray = xr.DataArray(
198
+ data=condition_communication_score_np,
199
+ dims=['subject', 'source', 'receiver', 'interaction'],
200
+ coords=[subject_list, self.source_cell_types, self.receiver_cell_types, self.interaction_names]
201
+ )
202
+
203
+ condition_array_ligand_abundance: xr.DataArray = xr.DataArray(
204
+ data=condition_ligand_abundance_np,
205
+ dims=['subject', 'source', 'interaction'],
206
+ coords=[subject_list, self.source_cell_types, self.interaction_names]
207
+ )
208
+
209
+ condition_array_target_abundance: xr.DataArray = xr.DataArray(
210
+ data=condition_target_abundance_np,
211
+ dims=['subject', 'receiver', 'interaction'],
212
+ coords=[subject_list, self.receiver_cell_types, self.interaction_names]
213
+ )
214
+
215
+ combined_communication_score_dict[condition] = condition_array
216
+ combined_ligand_abundance_dict[condition] = condition_array_ligand_abundance
217
+ combined_target_abundance_dict[condition] = condition_array_target_abundance
218
+
219
+ return combined_communication_score_dict, combined_ligand_abundance_dict, combined_target_abundance_dict
220
+
221
+ def __get_condition_specific_communication_scores_for_p_value_test(
222
+ self,
223
+ condition: str,
224
+ source: Union[str, int],
225
+ receiver: Union[str, int],
226
+ interaction: Union[str, int]
227
+ ) -> np.ndarray:
228
+ """
229
+ Get the communication scores for a specific source, receiver, and interaction for a specific condition.
230
+ Filter out all NaN values.
231
+
232
+ :param condition: The condition for which the communication scores should be extracted
233
+ :param source: the source cell-type
234
+ :param receiver: the receiver cell-type
235
+ :param interaction: the ligand-target interaction to extract
236
+ :return: numpy array containing the communication scores for the specified condition
237
+ """
238
+ # check if dtype is all the same
239
+ if not ((isinstance(source, str) and isinstance(receiver, str) and isinstance(interaction, str)) or \
240
+ (isinstance(source, int) and isinstance(receiver, int) and isinstance(interaction, int)) or \
241
+ (isinstance(source, np.integer) and isinstance(receiver, np.integer) and isinstance(interaction,
242
+ np.integer))):
243
+ raise ValueError('source, receiver, and interaction must be of the same type!')
244
+
245
+ if isinstance(source, str):
246
+ condition_communication_scores = self.communication_score_dict[condition].loc[:, source, receiver, interaction].to_numpy()
247
+ elif isinstance(source, int) or isinstance(source, np.integer):
248
+ condition_communication_scores = self.communication_score_dict[condition][:, source, receiver, interaction].to_numpy()
249
+ else:
250
+ raise ValueError('source must be either a string or an integer!')
251
+
252
+ # Filter out NaN values
253
+ condition_communication_scores = condition_communication_scores[~np.isnan(condition_communication_scores)]
254
+
255
+ return condition_communication_scores
256
+
257
+ def __get_communication_scores_for_p_value_test(
258
+ self,
259
+ source: Union[str, int],
260
+ receiver: Union[str, int],
261
+ interaction: Union[str, int]
262
+ ) -> Tuple[np.ndarray, np.ndarray]:
263
+ """
264
+ Get the communication scores for a specific source, receiver, and interaction. Filter out all NaN values.
265
+ :param source: the source cell-type
266
+ :param receiver: the receiver cell-type
267
+ :param interaction: the ligand-target interaction to extract
268
+ :return: Tuple of two numpy arrays containing the communication scores for the two conditions
269
+ """
270
+ res: tuple[np.ndarray, np.ndarray] = (
271
+ self.__get_condition_specific_communication_scores_for_p_value_test(
272
+ condition=self.condition_names[0],
273
+ source=source,
274
+ receiver=receiver,
275
+ interaction=interaction
276
+ ),
277
+ self.__get_condition_specific_communication_scores_for_p_value_test(
278
+ condition=self.condition_names[1],
279
+ source=source,
280
+ receiver=receiver,
281
+ interaction=interaction
282
+ )
283
+ )
284
+
285
+ return res
286
+
287
+ @staticmethod
288
+ def __correct_p_values(
289
+ p_values: np.ndarray,
290
+ method: str = 'by'
291
+ ) -> np.ndarray:
292
+ """
293
+ Correct p-values using the Benjamini-Yekutieli method and return the corrected p-values.
294
+ Exclude any nan values from the correction.
295
+ :param p_values: the p-values to correct
296
+ :param method: the method to use for correction (default: 'by')
297
+ :return: the corrected p-values with nan values where the p-values were nan
298
+ """
299
+ if method not in ['by', 'bh']:
300
+ raise ValueError('method must be either "by" or "bh"')
301
+
302
+ # initiate array with nan values
303
+ p_values_corrected_full = np.ones_like(p_values) * np.nan
304
+
305
+ # select only non nan values for correction
306
+ non_nan_p_values = p_values[~np.isnan(p_values)]
307
+
308
+ # correct p-values
309
+ p_values_corrected = stats.false_discovery_control(
310
+ ps=non_nan_p_values,
311
+ method=method
312
+ )
313
+
314
+ # insert corrected p-values into full array
315
+ p_values_corrected_full[~np.isnan(p_values)] = p_values_corrected
316
+
317
+ return p_values_corrected_full
318
+
319
+ @property
320
+ def _cube_shape(self) -> tuple[int, int, int]:
321
+ """Shape of a single (source, receiver, interaction) cube."""
322
+ return tuple(self.communication_score_dict[self.condition_names[0]].shape[1:])
323
+
324
+ def _validate_or_make_mask(self, mask: Optional[Union[np.ndarray, xr.DataArray]]) -> np.ndarray:
325
+ """Ensure mask has correct shape and dtype; create a full-True mask if None."""
326
+ shape = self._cube_shape
327
+ if mask is None:
328
+ return np.ones(shape=shape, dtype=bool)
329
+ if isinstance(mask, xr.DataArray):
330
+ mask_arr = mask.values
331
+ else:
332
+ mask_arr = mask
333
+ if mask_arr.shape != shape:
334
+ raise ValueError('The shape of the mask does not match the shape of the communication scores')
335
+ return mask_arr.astype(bool, copy=False)
336
+
337
+
338
+ def correct_p_values(
339
+ self,
340
+ statistical_test: Optional[Union[str, List[str]]] = None,
341
+ method: str = 'by'
342
+ ) -> Union[xr.DataArray, Dict[str, xr.DataArray]]:
343
+ """
344
+ This function corrects the p-values using the Benjamini-Yekutieli method. The corrected p-values are stored in
345
+ the MultiNeuronChatObject. If statistical_test is None, all p-values are corrected. If statistical_test is a
346
+ string, only the p-values of the specified test are corrected. If statistical_test is a list of strings, only
347
+ the p-values of the specified tests are corrected.
348
+
349
+ :param statistical_test: the test for which the p-values should be corrected
350
+ :param method: the method to use for correction (default: 'by')
351
+ :return: the corrected p-values as a DataArray or a dictionary of DataArrays
352
+ """
353
+ if statistical_test is None or isinstance(statistical_test, list):
354
+ # compute significance for all tests that have been computed or the specific list of tests
355
+ if statistical_test is None:
356
+ statistical_test = list(self.p_values.keys())
357
+
358
+ not_included_tests = [test for test in statistical_test if test not in self.p_values.keys()]
359
+
360
+ if len(not_included_tests) > 0:
361
+ raise ValueError(f'Tests {not_included_tests} have not been computed yet')
362
+
363
+ self.p_values_adj = {}
364
+
365
+ for test in statistical_test:
366
+ p_values = self.p_values[test].values
367
+ p_values_adj = self.__correct_p_values(
368
+ p_values,
369
+ method=method
370
+ )
371
+
372
+ p_values_adj_xr = xr.DataArray(
373
+ data=p_values_adj,
374
+ dims=['source', 'receiver', 'interaction'],
375
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
376
+ )
377
+
378
+ self.p_values_adj[test] = p_values_adj_xr
379
+
380
+ # if only one test was computed, return the p-values directly
381
+ if len(statistical_test) == 1:
382
+ return self.p_values_adj[statistical_test[0]]
383
+ # if multiple tests were computed, return the p-values as a dictionary
384
+ return self.p_values_adj
385
+ elif isinstance(statistical_test, str):
386
+ # compute significance for the specified test
387
+ if statistical_test not in self.p_values.keys():
388
+ raise ValueError(f'Test {statistical_test} has not been computed yet')
389
+
390
+ p_values = self.p_values[statistical_test].values
391
+ p_values_adj = self.__correct_p_values(
392
+ p_values,
393
+ method=method
394
+ )
395
+
396
+ p_values_adj_xr = xr.DataArray(
397
+ data=p_values_adj,
398
+ dims=['source', 'receiver', 'interaction'],
399
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
400
+ )
401
+
402
+ self.p_values_adj[statistical_test] = p_values_adj_xr
403
+
404
+ return p_values_adj_xr
405
+ else:
406
+ raise ValueError('If you define statistical_test, it must be either a string or a list of strings,'
407
+ 'where each string is a valid test')
408
+
409
+ @staticmethod
410
+ def _needs_permutation(p: float, coarse_levels: Optional[set] = None) -> bool:
411
+ """Return True if we should recompute via permutation.
412
+ We only trigger when p is non-finite, equals 0.0, or equals a known coarse/clamped level (e.g., 0.001 for Anderson–Darling).
413
+ """
414
+ return not np.isfinite(p) or p == 0.0 or (coarse_levels and (p in coarse_levels))
415
+
416
+ @staticmethod
417
+ def _ensure_nonzero_p(p: float, n_resamples: Optional[int] = None) -> float:
418
+ """Ensure returned p-value is strictly positive without imposing an arbitrary lower bound.
419
+ If p==0 from a permutation test, return 1/(n_resamples+1). Otherwise, map exact 0.0 to the
420
+ smallest positive float.
421
+ """
422
+ if p == 0.0:
423
+ if n_resamples is not None and n_resamples > 0:
424
+ return 1.0 / (n_resamples + 1)
425
+ # fallback for analytical edge cases
426
+ return np.nextafter(0.0, 1.0)
427
+ return p
428
+
429
+ @staticmethod
430
+ def __perm_pvalue_independent(
431
+ A: np.ndarray,
432
+ B: np.ndarray,
433
+ stat_fn,
434
+ n_resamples: int = 10_000,
435
+ alternative: str = 'two-sided',
436
+ random_state: Optional[int] = None
437
+ ) -> float:
438
+ """Return permutation p-value for an independent two-sample statistic.
439
+ Uses scipy.stats.permutation_test under the hood.
440
+ """
441
+ res = stats.permutation_test(
442
+ (A, B),
443
+ statistic=stat_fn,
444
+ permutation_type='independent',
445
+ n_resamples=n_resamples,
446
+ alternative=alternative,
447
+ random_state=random_state,
448
+ )
449
+ return float(res.pvalue)
450
+
451
+ def compute_significance(
452
+ self,
453
+ statistical_test: Optional[str] = 'KS',
454
+ mask: Optional[Union[np.ndarray, xr.DataArray]] = None,
455
+ n_resamples: Optional[int] = 10_000,
456
+ random_state: Optional[int] = None
457
+ ) -> xr.DataArray:
458
+ """
459
+ Compute the p-values for the specified statistical test. The p-values are stored in the MultiNeuronChatObject.
460
+ If mask is None, the significance is tested for all interactions and cell-type pairs. If a mask is provided,
461
+ the significance is only computed for the specified interactions and cell-type pairs, setting all other p-values
462
+ to np.nan.
463
+
464
+ :param statistical_test: the statistical test to compute the p-values for (default: 'KS', i.e. Kolmogorov-Smirnov).
465
+ :param mask: a mask to specify for which interactions and cell-type pairs the significance should be computed.
466
+ :param n_resamples: number of resamples for permutation fallback (default 10_000 if None)
467
+ :param random_state: random state for reproducibility (default None)
468
+ :return: the p-values as a DataArray.
469
+ """
470
+ mask = self._validate_or_make_mask(mask)
471
+
472
+ if statistical_test not in ['KS', 'Anderson', 'CVM', 'MannWhitneyU']:
473
+ raise ValueError('Statistical_test must be either "KS", "Anderson", "CVM", or "MannWhitneyU"')
474
+
475
+ if statistical_test == 'KS':
476
+ return self.__compute_p_values_KS(
477
+ mask=mask,
478
+ n_resamples=n_resamples,
479
+ random_state=random_state
480
+ )
481
+ elif statistical_test == 'Anderson':
482
+ return self.__compute_p_values_Anderson(
483
+ mask=mask,
484
+ n_resamples=n_resamples,
485
+ random_state=random_state
486
+ )
487
+ elif statistical_test == 'CVM':
488
+ return self.__compute_p_values_CVM(
489
+ mask=mask,
490
+ n_resamples=n_resamples,
491
+ random_state=random_state
492
+ )
493
+ elif statistical_test == 'MannWhitneyU':
494
+ return self.__compute_p_values_MannWhitneyU(
495
+ mask=mask,
496
+ n_resamples=n_resamples,
497
+ random_state=random_state
498
+ )
499
+
500
+ def __compute_p_values_KS(
501
+ self,
502
+ mask: Optional[Union[np.ndarray, xr.DataArray]] = None,
503
+ n_resamples: Optional[int] = 10_000,
504
+ random_state: Optional[int] = None
505
+ ) -> xr.DataArray:
506
+ mask = self._validate_or_make_mask(mask)
507
+
508
+ # Set up numpy arrays with NaN values
509
+ shape = (self.__n_cell_types, self.__n_cell_types, self.__n_interactions)
510
+ p_values = np.full(shape, np.nan, dtype=float)
511
+ test_statistic = np.full(shape, np.nan, dtype=float)
512
+ test_statistic_location = np.full(shape, np.nan, dtype=float)
513
+ test_statistic_sign = np.full(shape, np.nan, dtype=float)
514
+
515
+ # Compute p-values where mask is True
516
+ idx = np.where(mask)
517
+
518
+ for source, receiver, interaction in zip(*idx):
519
+ condition_a_communication_scores, condition_b_communication_scores = self.__get_communication_scores_for_p_value_test(
520
+ source=source,
521
+ receiver=receiver,
522
+ interaction=interaction
523
+ )
524
+
525
+ # Check if any of the two conditions has no communication scores
526
+ # If so, we have to set the p-value and the test statistic to np.nan as
527
+ # the KS-Test cannot be computed on empty distributions
528
+
529
+ if len(condition_a_communication_scores) == 0 or len(condition_b_communication_scores) == 0:
530
+ continue
531
+
532
+ # Check if both distributions are only zeros
533
+ # If this is the case: set p-value and test statistic to np.nan
534
+ if np.max(condition_a_communication_scores) == 0 and np.max(condition_b_communication_scores) == 0:
535
+ continue
536
+
537
+ ks_statistic = stats.ks_2samp(
538
+ data1=condition_a_communication_scores,
539
+ data2=condition_b_communication_scores,
540
+ )
541
+
542
+ # store observed stats
543
+ test_statistic[source, receiver, interaction] = ks_statistic.statistic
544
+ test_statistic_location[source, receiver, interaction] = ks_statistic.statistic_location
545
+ test_statistic_sign[source, receiver, interaction] = ks_statistic.statistic_sign
546
+
547
+ p = float(ks_statistic.pvalue)
548
+ if self._needs_permutation(p):
549
+ p_perm = self.__perm_pvalue_independent(
550
+ condition_a_communication_scores,
551
+ condition_b_communication_scores,
552
+ stat_fn=lambda a, b: stats.ks_2samp(a, b).statistic,
553
+ n_resamples=n_resamples,
554
+ alternative='two-sided',
555
+ random_state=random_state
556
+ )
557
+ p_final = self._ensure_nonzero_p(p_perm, n_resamples)
558
+ else:
559
+ p_final = self._ensure_nonzero_p(p)
560
+ p_values[source, receiver, interaction] = p_final
561
+
562
+ p_values_xr = xr.DataArray(
563
+ data=p_values,
564
+ dims=['source', 'receiver', 'interaction'],
565
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
566
+ )
567
+
568
+ test_statistic_xr: xr.DataArray = xr.DataArray(
569
+ data=test_statistic,
570
+ dims=['source', 'receiver', 'interaction'],
571
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
572
+ )
573
+ test_statistic_location_xr: xr.DataArray = xr.DataArray(
574
+ data=test_statistic_location,
575
+ dims=['source', 'receiver', 'interaction'],
576
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
577
+ )
578
+ test_statistic_sign_xr: xr.DataArray = xr.DataArray(
579
+ data=test_statistic_sign,
580
+ dims=['source', 'receiver', 'interaction'],
581
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
582
+ )
583
+
584
+ if self.p_values is None:
585
+ self.p_values = {}
586
+ self.p_values_adj = {}
587
+ self.statistics = {}
588
+
589
+ self.p_values['KS'] = p_values_xr
590
+ self.statistics['KS'] = {
591
+ 'statistics': test_statistic_xr,
592
+ 'location': test_statistic_location_xr,
593
+ 'sign': test_statistic_sign_xr
594
+ }
595
+
596
+ return p_values_xr
597
+
598
+ def __compute_p_values_Anderson(
599
+ self,
600
+ mask: Optional[Union[np.ndarray, xr.DataArray]] = None,
601
+ permutation_method: Optional[bool] = False,
602
+ n_resamples: Optional[int] = 10_000,
603
+ random_state: Optional[int] = None
604
+ ) -> xr.DataArray:
605
+ mask = self._validate_or_make_mask(mask)
606
+
607
+ if permutation_method and n_resamples < 1:
608
+ raise ValueError('The number of resamples must be greater than 0')
609
+
610
+ # Set up numpy arrays with NaN values
611
+ shape = (self.__n_cell_types, self.__n_cell_types, self.__n_interactions)
612
+ p_values = np.full(shape, np.nan, dtype=float)
613
+ test_statistic = np.full(shape, np.nan, dtype=float)
614
+
615
+ # Compute p-values where mask is True
616
+ idx = np.where(mask)
617
+
618
+ for source, receiver, interaction in zip(*idx):
619
+ condition_a_communication_scores, condition_b_communication_scores = self.__get_communication_scores_for_p_value_test(
620
+ source=source,
621
+ receiver=receiver,
622
+ interaction=interaction
623
+ )
624
+
625
+ # The Anderson Darling Test does work if all values of the data are the same
626
+ # Therefore we have to check if the data is the same and if so we have to skip the test
627
+ # IMPORTANT: this includes the case when both distributions are the same, e.g., both are completely zero
628
+ if np.unique(np.hstack([condition_a_communication_scores, condition_b_communication_scores])).shape[0] == 1:
629
+ continue
630
+
631
+ # TODO Check this; I had a bug, where when both samples are only one sample long, the Anderson test does not work.
632
+ # This seems to be the case because of an scipy implementation detail where they arrange N-1 sampes, i.e.,
633
+ # if N=1, then they have 0 samples. This then lead to an error for the Anderson test.
634
+ # Therefore, I have decided to restrict the test to the case where both samples have at least 2 samples.
635
+ # This is not ideal, but I think it is the best solution for now.
636
+ if len(condition_a_communication_scores) < 2 or len(condition_b_communication_scores) < 2:
637
+ continue
638
+
639
+ # First try the default (analytical/approximate) unless permutation was explicitly requested
640
+ if not permutation_method:
641
+ anderson_statistic = stats.anderson_ksamp(
642
+ samples=[
643
+ condition_a_communication_scores,
644
+ condition_b_communication_scores,
645
+ ]
646
+ )
647
+ p = float(anderson_statistic.pvalue)
648
+ # SciPy's approximate p can be clipped (~<=0.001). If clipped or non-finite, redo via permutation.
649
+ if self._needs_permutation(p, coarse_levels={0.001}):
650
+ anderson_statistic = stats.anderson_ksamp(
651
+ samples=[
652
+ condition_a_communication_scores,
653
+ condition_b_communication_scores,
654
+ ],
655
+ method=PermutationMethod(
656
+ n_resamples=n_resamples,
657
+ random_state=random_state
658
+ )
659
+ )
660
+ else:
661
+ anderson_statistic = stats.anderson_ksamp(
662
+ samples=[
663
+ condition_a_communication_scores,
664
+ condition_b_communication_scores
665
+ ],
666
+ method=PermutationMethod(
667
+ n_resamples=n_resamples,
668
+ random_state=random_state
669
+ )
670
+ )
671
+
672
+ p_values[source, receiver, interaction] = self._ensure_nonzero_p(float(anderson_statistic.pvalue), n_resamples if permutation_method else None)
673
+ test_statistic[source, receiver, interaction] = float(anderson_statistic.statistic)
674
+
675
+ p_values_anderson_xr = xr.DataArray(
676
+ data=p_values,
677
+ dims=['source', 'receiver', 'interaction'],
678
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
679
+ )
680
+
681
+ test_statistic_anderson_xr: xr.DataArray = xr.DataArray(
682
+ data=test_statistic,
683
+ dims=['source', 'receiver', 'interaction'],
684
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
685
+ )
686
+
687
+ if self.p_values is None:
688
+ self.p_values = {}
689
+ self.p_values_adj = {}
690
+ self.statistics = {}
691
+
692
+ self.p_values['Anderson'] = p_values_anderson_xr
693
+ self.statistics['Anderson'] = {
694
+ 'statistics': test_statistic_anderson_xr
695
+ }
696
+
697
+ return p_values_anderson_xr
698
+
699
+ def __compute_p_values_CVM(
700
+ self,
701
+ mask: Optional[Union[np.ndarray, xr.DataArray]] = None,
702
+ n_resamples: Optional[int] = 10_000,
703
+ random_state: Optional[int] = None
704
+ ) -> xr.DataArray:
705
+ mask = self._validate_or_make_mask(mask)
706
+
707
+ # Set up numpy arrays with NaN values
708
+ shape = (self.__n_cell_types, self.__n_cell_types, self.__n_interactions)
709
+ p_values = np.full(shape, np.nan, dtype=float)
710
+ test_statistic = np.full(shape, np.nan, dtype=float)
711
+
712
+ # Compute p-values where mask is True
713
+ idx = np.where(mask)
714
+
715
+ for source, receiver, interaction in zip(*idx):
716
+ condition_a_communication_scores, condition_b_communication_scores = self.__get_communication_scores_for_p_value_test(
717
+ source=source,
718
+ receiver=receiver,
719
+ interaction=interaction
720
+ )
721
+
722
+ # The CVM Test does work if all values of the data are the same
723
+ # Therefore we have to check if the data is the same and if so we have to skip the test
724
+ # IMPORTANT: this includes the case when both distributions are the same, e.g., both are completely zero
725
+ if np.unique(np.hstack([condition_a_communication_scores, condition_b_communication_scores])).shape[0] == 1:
726
+ continue
727
+
728
+ # The CVM Test requires that there are at least two observations in each array
729
+ if len(condition_a_communication_scores) < 2 or len(condition_b_communication_scores) < 2:
730
+ continue
731
+
732
+ cvm_statistic = stats.cramervonmises_2samp(
733
+ x=condition_a_communication_scores,
734
+ y=condition_b_communication_scores,
735
+ method='auto'
736
+ )
737
+ # Always store the observed statistic
738
+ test_statistic[source, receiver, interaction] = cvm_statistic.statistic
739
+
740
+ # Fallback criteria: exact zero, non-finite, or suspiciously tiny asymptotic p (underflow)
741
+ p = float(cvm_statistic.pvalue)
742
+ if self._needs_permutation(p):
743
+ # Use permutation p-value with the same statistic definition
744
+ p_perm = self.__perm_pvalue_independent(
745
+ condition_a_communication_scores,
746
+ condition_b_communication_scores,
747
+ stat_fn=lambda a, b: stats.cramervonmises_2samp(x=a, y=b, method='auto').statistic,
748
+ n_resamples=n_resamples,
749
+ alternative='two-sided',
750
+ random_state=random_state
751
+ )
752
+ p_final = self._ensure_nonzero_p(p_perm, n_resamples)
753
+ else:
754
+ p_final = self._ensure_nonzero_p(p)
755
+ p_values[source, receiver, interaction] = p_final
756
+
757
+ p_values_cvm_xr = xr.DataArray(
758
+ data=p_values,
759
+ dims=['source', 'receiver', 'interaction'],
760
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
761
+ )
762
+
763
+ test_statistic_cvm_xr: xr.DataArray = xr.DataArray(
764
+ data=test_statistic,
765
+ dims=['source', 'receiver', 'interaction'],
766
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
767
+ )
768
+
769
+ if self.p_values is None:
770
+ self.p_values = {}
771
+ self.p_values_adj = {}
772
+ self.statistics = {}
773
+
774
+ self.p_values['CVM'] = p_values_cvm_xr
775
+ self.statistics['CVM'] = {
776
+ 'statistics': test_statistic_cvm_xr
777
+ }
778
+
779
+ return p_values_cvm_xr
780
+
781
+ def __compute_p_values_MannWhitneyU(
782
+ self,
783
+ mask: Optional[Union[np.ndarray, xr.DataArray]] = None,
784
+ n_resamples: Optional[int] = 10_000,
785
+ random_state: Optional[int] = None
786
+ ) -> xr.DataArray:
787
+ mask = self._validate_or_make_mask(mask)
788
+
789
+ shape = (self.__n_cell_types, self.__n_cell_types, self.__n_interactions)
790
+ p_values = np.full(shape, np.nan, dtype=float)
791
+ test_statistic = np.full(shape, np.nan, dtype=float)
792
+
793
+ # Compute p-values where mask is True
794
+ idx = np.where(mask)
795
+
796
+ for source, receiver, interaction in zip(*idx):
797
+ condition_a_communication_scores, condition_b_communication_scores = self.__get_communication_scores_for_p_value_test(
798
+ source=source,
799
+ receiver=receiver,
800
+ interaction=interaction
801
+ )
802
+
803
+ if np.unique(np.hstack([condition_a_communication_scores, condition_b_communication_scores])).shape[0] == 1:
804
+ continue
805
+
806
+ if len(condition_a_communication_scores) == 0 or len(condition_b_communication_scores) == 0:
807
+ continue
808
+
809
+ mannwhitneyu_statistic = stats.mannwhitneyu(
810
+ x=condition_a_communication_scores,
811
+ y=condition_b_communication_scores,
812
+ alternative='two-sided',
813
+ )
814
+
815
+ p = float(mannwhitneyu_statistic.pvalue)
816
+ test_statistic[source, receiver, interaction] = float(mannwhitneyu_statistic.statistic)
817
+
818
+ if self._needs_permutation(p):
819
+ res_perm = stats.mannwhitneyu(
820
+ x=condition_a_communication_scores,
821
+ y=condition_b_communication_scores,
822
+ alternative='two-sided',
823
+ method=PermutationMethod(
824
+ n_resamples=n_resamples,
825
+ random_state=random_state
826
+ )
827
+ )
828
+ p_final = self._ensure_nonzero_p(float(res_perm.pvalue), n_resamples)
829
+ else:
830
+ p_final = self._ensure_nonzero_p(p)
831
+ p_values[source, receiver, interaction] = p_final
832
+
833
+ p_values_mannwhitneyu_xr = xr.DataArray(
834
+ data=p_values,
835
+ dims=['source', 'receiver', 'interaction'],
836
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
837
+ )
838
+
839
+ test_statistic_mannwhitneyu_xr: xr.DataArray = xr.DataArray(
840
+ data=test_statistic,
841
+ dims=['source', 'receiver', 'interaction'],
842
+ coords=[self.source_cell_types, self.receiver_cell_types, self.interaction_names]
843
+ )
844
+
845
+ if self.p_values is None:
846
+ self.p_values = {}
847
+ self.p_values_adj = {}
848
+ self.statistics = {}
849
+
850
+ self.p_values['MannWhitneyU'] = p_values_mannwhitneyu_xr
851
+ self.statistics['MannWhitneyU'] = {
852
+ 'statistics': test_statistic_mannwhitneyu_xr
853
+ }
854
+
855
+ return p_values_mannwhitneyu_xr
856
+
857
+ def save(self, path_to_file: str):
858
+ """
859
+ Save the MultiNeuronChatObject to a pickle file.
860
+
861
+ :param path_to_file: path to the pickle file
862
+ """
863
+ dir_name: str = os.path.dirname(path_to_file)
864
+ if not os.path.exists(dir_name):
865
+ os.makedirs(dir_name)
866
+
867
+ # Save all variables as a dictionary and finally as a pickle file
868
+ save_dict: Dict[str, Any] = {
869
+ '__n_cell_types': self.__n_cell_types,
870
+ '__n_interactions': self.__n_interactions,
871
+
872
+ 'condition_label_column': self.condition_label_column,
873
+ 'condition_names': self.condition_names,
874
+
875
+ 'subject_label_column': self.subject_label_column,
876
+ 'cell_type_label_column': self.cell_type_label_column,
877
+
878
+ 'db': self.db,
879
+ 'interaction_db': self.interaction_db,
880
+ 'gene_set': self.gene_set,
881
+
882
+ 'mean_type': self.mean_type,
883
+ 'trim_mean_fraction': self.trim_mean_fraction,
884
+
885
+ 'source_cell_types': self.source_cell_types,
886
+ 'receiver_cell_types': self.receiver_cell_types,
887
+ 'interaction_names': self.interaction_names,
888
+
889
+ 'avg_expression_per_condition_and_subject_dict': self.avg_expression_per_condition_and_subject_dict,
890
+
891
+ 'communication_scores_per_condition_and_subject_dict': self.communication_scores_per_condition_and_subject_dict,
892
+ 'ligand_abundance_per_condition_and_subject_dict': self.ligand_abundance_per_condition_and_subject_dict,
893
+ 'target_abundance_per_condition_and_subject_dict': self.target_abundance_per_condition_and_subject_dict,
894
+
895
+ 'communication_score_dict': self.communication_score_dict,
896
+ 'ligand_abundance_dict': self.ligand_abundance_dict,
897
+ 'target_abundance_dict': self.target_abundance_dict,
898
+
899
+ 'p_values': self.p_values,
900
+ 'p_values_adj': self.p_values_adj,
901
+ 'statistics': self.statistics,
902
+ }
903
+
904
+ with open(path_to_file, 'wb') as file:
905
+ pickle.dump(save_dict, file)
906
+
907
+ @staticmethod
908
+ def load(path_to_file: str):
909
+ """
910
+ Load a MultiNeuronChatObject from a pickle file.
911
+
912
+ :param path_to_file: path to the pickle file
913
+ :return: MultiNeuronChatObject loaded from the pickle file
914
+ """
915
+ if not os.path.exists(path_to_file):
916
+ raise FileNotFoundError(f'The file with path {path_to_file} could not be found.')
917
+
918
+ with open(path_to_file, 'rb') as file:
919
+ load_dict: Dict[str, Any] = pickle.load(file)
920
+
921
+ mnc_obj: MultiNeuronChatObject = MultiNeuronChatObject(
922
+ condition_label_column=load_dict['condition_label_column'],
923
+ condition_names=load_dict['condition_names'],
924
+ subject_label_column=load_dict['subject_label_column'],
925
+ cell_type_label_column=load_dict['cell_type_label_column'],
926
+ db=load_dict['db'],
927
+ interaction_db=load_dict['interaction_db'],
928
+ )
929
+
930
+ mnc_obj.__n_cell_types = load_dict['__n_cell_types']
931
+ mnc_obj.__n_interactions = load_dict['__n_interactions']
932
+
933
+ mnc_obj.gene_set = load_dict['gene_set']
934
+
935
+ mnc_obj.mean_type = load_dict['mean_type']
936
+ mnc_obj.trim_mean_fraction = load_dict['trim_mean_fraction']
937
+
938
+ mnc_obj.source_cell_types = load_dict['source_cell_types']
939
+ mnc_obj.receiver_cell_types = load_dict['receiver_cell_types']
940
+ mnc_obj.interaction_names = load_dict['interaction_names']
941
+
942
+ mnc_obj.avg_expression_per_condition_and_subject_dict = load_dict['avg_expression_per_condition_and_subject_dict']
943
+
944
+ # Ensure compatibility with prior versions
945
+ if 'ligand_abundance_per_condition_and_subject_dict' in load_dict.keys():
946
+ mnc_obj.ligand_abundance_per_condition_and_subject_dict = load_dict['ligand_abundance_per_condition_and_subject_dict']
947
+ mnc_obj.target_abundance_per_condition_and_subject_dict = load_dict['target_abundance_per_condition_and_subject_dict']
948
+
949
+ mnc_obj.ligand_abundance_dict = load_dict['ligand_abundance_dict']
950
+ mnc_obj.target_abundance_dict = load_dict['target_abundance_dict']
951
+
952
+ mnc_obj.communication_scores_per_condition_and_subject_dict = load_dict['communication_scores_per_condition_and_subject_dict']
953
+ mnc_obj.communication_score_dict = load_dict['communication_score_dict']
954
+
955
+ mnc_obj.p_values = load_dict['p_values']
956
+ mnc_obj.p_values_adj = load_dict['p_values_adj']
957
+ mnc_obj.statistics = load_dict['statistics']
958
+
959
+ return mnc_obj