multineuronchat 2026.8.18.dev0__tar.gz → 2026.9.1.dev0__tar.gz

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.
Files changed (37) hide show
  1. {multineuronchat-2026.8.18.dev0/src/multineuronchat.egg-info → multineuronchat-2026.9.1.dev0}/PKG-INFO +1 -1
  2. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/pyproject.toml +1 -1
  3. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/MultiNeuronChatObject.py +149 -78
  4. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0/src/multineuronchat.egg-info}/PKG-INFO +1 -1
  5. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat.egg-info/SOURCES.txt +12 -1
  6. multineuronchat-2026.9.1.dev0/tests/test_avg_expression.py +272 -0
  7. multineuronchat-2026.9.1.dev0/tests/test_communication_score.py +190 -0
  8. multineuronchat-2026.9.1.dev0/tests/test_correction.py +147 -0
  9. multineuronchat-2026.9.1.dev0/tests/test_golden_master.py +95 -0
  10. multineuronchat-2026.9.1.dev0/tests/test_interaction_db.py +77 -0
  11. multineuronchat-2026.9.1.dev0/tests/test_masks.py +231 -0
  12. multineuronchat-2026.9.1.dev0/tests/test_normalize.py +179 -0
  13. multineuronchat-2026.9.1.dev0/tests/test_object_save_load.py +142 -0
  14. multineuronchat-2026.9.1.dev0/tests/test_power.py +225 -0
  15. multineuronchat-2026.9.1.dev0/tests/test_significance.py +301 -0
  16. multineuronchat-2026.9.1.dev0/tests/test_utils_filter.py +117 -0
  17. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/LICENSE +0 -0
  18. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/README.md +0 -0
  19. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/setup.cfg +0 -0
  20. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/InteractionDB/InteractionDB.py +0 -0
  21. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/InteractionDB/InteractionDBRow.py +0 -0
  22. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/InteractionDB/__init__.py +0 -0
  23. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/MultiNeuronChat.py +0 -0
  24. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/__init__.py +0 -0
  25. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/db/__init__.py +0 -0
  26. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/db/interactionDB_human.pkl +0 -0
  27. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/db/interactionDB_human_extended.pkl +0 -0
  28. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/db/interactionDB_mouse.pkl +0 -0
  29. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/loompy_utils.py +0 -0
  30. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/masks.py +0 -0
  31. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/normalize.py +0 -0
  32. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/power.py +0 -0
  33. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/utils.py +0 -0
  34. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat/visualize.py +0 -0
  35. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat.egg-info/dependency_links.txt +0 -0
  36. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat.egg-info/requires.txt +0 -0
  37. {multineuronchat-2026.8.18.dev0 → multineuronchat-2026.9.1.dev0}/src/multineuronchat.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: multineuronchat
3
- Version: 2026.8.18.dev0
3
+ Version: 2026.9.1.dev0
4
4
  Summary: MultiNeuronChat is a Python library for inferring condition‑related changes in synaptic cell‑cell communication from scRNA-/snRNA-Seq datasets in case vs control study designs.
5
5
  Author-email: Gianluca Volkmer <gianluca.volkmer@ki.se>
6
6
  License-Expression: GPL-3.0
@@ -11,7 +11,7 @@ description = "MultiNeuronChat is a Python library for inferring condition‑rel
11
11
  readme = "README.md"
12
12
  license = "GPL-3.0"
13
13
  license-files = ["LICENSE"]
14
- version = "2026.8.18.dev0"
14
+ version = "2026.9.1.dev0"
15
15
  requires-python = '>=3.10'
16
16
  dependencies = [
17
17
  "numpy==1.26.4",
@@ -4,6 +4,8 @@ import pickle
4
4
 
5
5
  import warnings
6
6
 
7
+ from multiprocessing import Pool
8
+
7
9
  import numpy as np
8
10
  import xarray as xr
9
11
 
@@ -19,6 +21,112 @@ from .MultiNeuronChat import compute_avg_expression, compute_communication_score
19
21
  from typing import Tuple, Optional, List, Dict, Set, Any, Union
20
22
 
21
23
 
24
+ def _compute_anderson_pvalue_for_triple(
25
+ task: Tuple[Any, Any, Any, np.ndarray, np.ndarray, bool, int, int, bool, Optional[int]]
26
+ ) -> Optional[Tuple[Any, Any, Any, float, float, bool]]:
27
+ """
28
+ Compute the Anderson-Darling p-value and test statistic for a single triple.
29
+
30
+ This is the body of the per-triple loop in ``MultiNeuronChatObject.__compute_p_values_Anderson``, lifted out to a
31
+ module-level function so it can be dispatched to a ``multiprocessing.Pool``. It is kept at module scope (rather than a
32
+ closure or a bound method) so it stays picklable under the 'spawn' start method used on macOS, and it receives the two
33
+ pre-extracted score arrays instead of the whole object so the parent does not have to pickle its state per task.
34
+
35
+ The computation is identical to the serial path, including which triples are skipped and how randomness is seeded:
36
+ the first stage uses the shared run-level ``random_state`` and the escalation stage uses ``_derive_seed`` keyed on the
37
+ triple's position. Both are independent of the order in which triples are processed, so the parallel result is
38
+ bit-identical to the serial one regardless of scheduling.
39
+
40
+ :param task: a tuple ``(source, receiver, interaction, scores_a, scores_b, permutation_method, n_resamples_stage_one,
41
+ n_resamples_stage_two, escalate, random_state)``
42
+ :return: ``(source, receiver, interaction, p_value, test_statistic, escalated)``, or None if the triple is degenerate
43
+ (all scores identical, or fewer than two scores in either condition) and must be left as NaN
44
+ """
45
+ (source, receiver, interaction,
46
+ condition_a_communication_scores, condition_b_communication_scores,
47
+ permutation_method,
48
+ n_resamples_stage_one, n_resamples_stage_two,
49
+ escalate, random_state) = task
50
+
51
+ # The Anderson Darling Test does work if all values of the data are the same
52
+ # Therefore we have to check if the data is the same and if so we have to skip the test
53
+ # IMPORTANT: this includes the case when both distributions are the same, e.g., both are completely zero
54
+ if np.unique(np.hstack([condition_a_communication_scores, condition_b_communication_scores])).shape[0] == 1:
55
+ return None
56
+
57
+ # Restrict to the case where both samples have at least 2 samples (SciPy's anderson_ksamp arranges N-1 samples, so a
58
+ # single-value sample yields 0 and errors). Mirrors the serial guard.
59
+ if len(condition_a_communication_scores) < 2 or len(condition_b_communication_scores) < 2:
60
+ return None
61
+
62
+ # Track whether the reported p-value came from a permutation run, and at what depth, so the escalation trigger and the
63
+ # zero-guard below both refer to the resolution that actually produced it.
64
+ n_resamples_used: Optional[int] = None
65
+
66
+ # First try the default (analytical/approximate) unless permutation was explicitly requested
67
+ if not permutation_method:
68
+ anderson_statistic = stats.anderson_ksamp(
69
+ samples=[
70
+ condition_a_communication_scores,
71
+ condition_b_communication_scores,
72
+ ]
73
+ )
74
+ p = float(anderson_statistic.pvalue)
75
+ # SciPy's approximate p can be clipped (~<=0.001). If clipped or non-finite, redo via permutation.
76
+ if MultiNeuronChatObject._needs_permutation(p, coarse_levels={0.001}):
77
+ anderson_statistic = stats.anderson_ksamp(
78
+ samples=[
79
+ condition_a_communication_scores,
80
+ condition_b_communication_scores,
81
+ ],
82
+ method=PermutationMethod(
83
+ n_resamples=n_resamples_stage_one,
84
+ random_state=random_state
85
+ )
86
+ )
87
+ n_resamples_used = n_resamples_stage_one
88
+ else:
89
+ anderson_statistic = stats.anderson_ksamp(
90
+ samples=[
91
+ condition_a_communication_scores,
92
+ condition_b_communication_scores
93
+ ],
94
+ method=PermutationMethod(
95
+ n_resamples=n_resamples_stage_one,
96
+ random_state=random_state
97
+ )
98
+ )
99
+ n_resamples_used = n_resamples_stage_one
100
+
101
+ # Second stage: this triple's p-value is pinned at the resolution of the first stage and carries no information beyond
102
+ # "smaller than 1 / (R + 1)". Re-run it deeply enough to survive the correction, with fresh draws from a derived seed,
103
+ # and replace the first-stage value outright.
104
+ escalated: bool = False
105
+ if escalate and MultiNeuronChatObject._is_at_monte_carlo_floor(float(anderson_statistic.pvalue), n_resamples_used):
106
+ anderson_statistic = stats.anderson_ksamp(
107
+ samples=[
108
+ condition_a_communication_scores,
109
+ condition_b_communication_scores,
110
+ ],
111
+ method=PermutationMethod(
112
+ n_resamples=n_resamples_stage_two,
113
+ random_state=MultiNeuronChatObject._derive_seed(
114
+ random_state=random_state,
115
+ source=source,
116
+ receiver=receiver,
117
+ interaction=interaction
118
+ )
119
+ )
120
+ )
121
+ n_resamples_used = n_resamples_stage_two
122
+ escalated = True
123
+
124
+ p_value: float = MultiNeuronChatObject._ensure_nonzero_p(float(anderson_statistic.pvalue), n_resamples_used)
125
+ test_statistic: float = float(anderson_statistic.statistic)
126
+
127
+ return source, receiver, interaction, p_value, test_statistic, escalated
128
+
129
+
22
130
  class MultiNeuronChatObject:
23
131
  def __init__(self,
24
132
  condition_label_column: str,
@@ -602,7 +710,8 @@ class MultiNeuronChatObject:
602
710
  warn_if_underpowered: Optional[bool] = True,
603
711
  correction_method: Optional[str] = 'by',
604
712
  alpha: Optional[float] = 0.05,
605
- escalate_permutation: Optional[bool] = True
713
+ escalate_permutation: Optional[bool] = True,
714
+ n_processes: Optional[int] = 1
606
715
  ) -> xr.DataArray:
607
716
  """
608
717
  Compute the p-values for the specified statistical test. The p-values are stored in the MultiNeuronChatObject.
@@ -630,6 +739,11 @@ class MultiNeuronChatObject:
630
739
  floor with the deeper resample budget the correction requires (default: True).
631
740
  Only affects the 'Anderson' test, the only one whose analytical p-value is clamped
632
741
  and therefore always routed through the permutation fallback.
742
+ :param n_processes: number of processes to parallelise the per-triple loop across (default: 1, i.e. serial).
743
+ Only the 'Anderson' test is parallelised, since it is the only one that always runs the
744
+ (expensive) permutation fallback; the per-triple work is embarrassingly parallel and the
745
+ result is identical to the serial run because every triple's randomness is seeded
746
+ independently of the scheduling order.
633
747
  :return: the p-values as a DataArray.
634
748
  """
635
749
  mask = self._validate_or_make_mask(mask)
@@ -683,7 +797,8 @@ class MultiNeuronChatObject:
683
797
  random_state=random_state,
684
798
  escalate_permutation=escalate_permutation,
685
799
  correction_method=correction_method,
686
- alpha=alpha
800
+ alpha=alpha,
801
+ n_processes=n_processes
687
802
  )
688
803
  elif statistical_test == 'CVM':
689
804
  return self.__compute_p_values_CVM(
@@ -813,7 +928,8 @@ class MultiNeuronChatObject:
813
928
  random_state: Optional[int] = None,
814
929
  escalate_permutation: Optional[bool] = True,
815
930
  correction_method: Optional[str] = 'by',
816
- alpha: Optional[float] = 0.05
931
+ alpha: Optional[float] = 0.05,
932
+ n_processes: Optional[int] = 1
817
933
  ) -> xr.DataArray:
818
934
  """
819
935
  Compute Anderson-Darling p-values, resampling deeply enough that they can survive multiple-testing correction.
@@ -847,6 +963,9 @@ class MultiNeuronChatObject:
847
963
  :param escalate_permutation: whether to run the second stage for floor-pinned triples (default: True)
848
964
  :param correction_method: the correction the p-values are intended for, used to size the second stage
849
965
  :param alpha: the significance level the p-values are intended to be judged at
966
+ :param n_processes: number of processes to distribute the per-triple loop across (default: 1, i.e. serial). Each
967
+ triple is independent and its randomness is seeded independently of the scheduling order, so
968
+ the result is identical to the serial run for any number of processes.
850
969
  :return: the p-values as a DataArray
851
970
  """
852
971
  mask = self._validate_or_make_mask(mask)
@@ -886,91 +1005,43 @@ class MultiNeuronChatObject:
886
1005
 
887
1006
  n_escalated: int = 0
888
1007
 
1008
+ # Build one independent task per tested triple. The scores are extracted here (cheap xarray slicing) rather than
1009
+ # inside the workers so the workers never need the object itself. See the TODO preserved in the worker: both
1010
+ # samples must have at least 2 values, which the worker checks and skips otherwise.
1011
+ tasks: List[Tuple[Any, Any, Any, np.ndarray, np.ndarray, bool, int, int, bool, Optional[int]]] = []
889
1012
  for source, receiver, interaction in zip(*idx):
890
1013
  condition_a_communication_scores, condition_b_communication_scores = self.__get_communication_scores_for_p_value_test(
891
1014
  source=source,
892
1015
  receiver=receiver,
893
1016
  interaction=interaction
894
1017
  )
1018
+ tasks.append((
1019
+ source, receiver, interaction,
1020
+ condition_a_communication_scores, condition_b_communication_scores,
1021
+ bool(permutation_method),
1022
+ n_resamples_stage_one, n_resamples_stage_two,
1023
+ escalate, random_state
1024
+ ))
1025
+
1026
+ # The per-triple work is embarrassingly parallel and dominated by the permutation resampling of the few escalated
1027
+ # triples, so a process pool with dynamic scheduling (chunksize=1) keeps every worker busy on the heavy triples
1028
+ # rather than letting one straggler chunk serialise the run. Results are written by (source, receiver,
1029
+ # interaction) index, so completion order is irrelevant and the output is identical for any n_processes.
1030
+ if n_processes is None or n_processes <= 1:
1031
+ results = [_compute_anderson_pvalue_for_triple(task) for task in tasks]
1032
+ else:
1033
+ with Pool(processes=n_processes) as pool:
1034
+ results = list(pool.imap_unordered(_compute_anderson_pvalue_for_triple, tasks, chunksize=1))
895
1035
 
896
- # The Anderson Darling Test does work if all values of the data are the same
897
- # Therefore we have to check if the data is the same and if so we have to skip the test
898
- # IMPORTANT: this includes the case when both distributions are the same, e.g., both are completely zero
899
- if np.unique(np.hstack([condition_a_communication_scores, condition_b_communication_scores])).shape[0] == 1:
900
- continue
901
-
902
- # TODO Check this; I had a bug, where when both samples are only one sample long, the Anderson test does not work.
903
- # This seems to be the case because of an scipy implementation detail where they arrange N-1 sampes, i.e.,
904
- # if N=1, then they have 0 samples. This then lead to an error for the Anderson test.
905
- # Therefore, I have decided to restrict the test to the case where both samples have at least 2 samples.
906
- # This is not ideal, but I think it is the best solution for now.
907
- if len(condition_a_communication_scores) < 2 or len(condition_b_communication_scores) < 2:
1036
+ for result in results:
1037
+ if result is None:
908
1038
  continue
909
-
910
- # Track whether the reported p-value came from a permutation run, and at what depth, so the escalation
911
- # trigger and the zero-guard below both refer to the resolution that actually produced it.
912
- n_resamples_used: Optional[int] = None
913
-
914
- # First try the default (analytical/approximate) unless permutation was explicitly requested
915
- if not permutation_method:
916
- anderson_statistic = stats.anderson_ksamp(
917
- samples=[
918
- condition_a_communication_scores,
919
- condition_b_communication_scores,
920
- ]
921
- )
922
- p = float(anderson_statistic.pvalue)
923
- # SciPy's approximate p can be clipped (~<=0.001). If clipped or non-finite, redo via permutation.
924
- if self._needs_permutation(p, coarse_levels={0.001}):
925
- anderson_statistic = stats.anderson_ksamp(
926
- samples=[
927
- condition_a_communication_scores,
928
- condition_b_communication_scores,
929
- ],
930
- method=PermutationMethod(
931
- n_resamples=n_resamples_stage_one,
932
- random_state=random_state
933
- )
934
- )
935
- n_resamples_used = n_resamples_stage_one
936
- else:
937
- anderson_statistic = stats.anderson_ksamp(
938
- samples=[
939
- condition_a_communication_scores,
940
- condition_b_communication_scores
941
- ],
942
- method=PermutationMethod(
943
- n_resamples=n_resamples_stage_one,
944
- random_state=random_state
945
- )
946
- )
947
- n_resamples_used = n_resamples_stage_one
948
-
949
- # Second stage: this triple's p-value is pinned at the resolution of the first stage and carries no
950
- # information beyond "smaller than 1 / (R + 1)". Re-run it deeply enough to survive the correction, with
951
- # fresh draws from a derived seed, and replace the first-stage value outright.
952
- if escalate and self._is_at_monte_carlo_floor(float(anderson_statistic.pvalue), n_resamples_used):
953
- anderson_statistic = stats.anderson_ksamp(
954
- samples=[
955
- condition_a_communication_scores,
956
- condition_b_communication_scores,
957
- ],
958
- method=PermutationMethod(
959
- n_resamples=n_resamples_stage_two,
960
- random_state=self._derive_seed(
961
- random_state=random_state,
962
- source=source,
963
- receiver=receiver,
964
- interaction=interaction
965
- )
966
- )
967
- )
968
- n_resamples_used = n_resamples_stage_two
1039
+ source, receiver, interaction, p_value, test_statistic_value, escalated = result
1040
+ p_values[source, receiver, interaction] = p_value
1041
+ test_statistic[source, receiver, interaction] = test_statistic_value
1042
+ if escalated:
969
1043
  n_escalated += 1
970
1044
 
971
- p_values[source, receiver, interaction] = self._ensure_nonzero_p(float(anderson_statistic.pvalue), n_resamples_used)
972
- test_statistic[source, receiver, interaction] = float(anderson_statistic.statistic)
973
-
974
1045
  # If escalation was disabled or capped below what the correction needs, p-values can still be pinned at the
975
1046
  # Monte-Carlo floor. Those triples are unreportable no matter how large their effect, so say so explicitly
976
1047
  # rather than letting them silently fail the correction.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: multineuronchat
3
- Version: 2026.8.18.dev0
3
+ Version: 2026.9.1.dev0
4
4
  Summary: MultiNeuronChat is a Python library for inferring condition‑related changes in synaptic cell‑cell communication from scRNA-/snRNA-Seq datasets in case vs control study designs.
5
5
  Author-email: Gianluca Volkmer <gianluca.volkmer@ki.se>
6
6
  License-Expression: GPL-3.0
@@ -21,4 +21,15 @@ src/multineuronchat/InteractionDB/__init__.py
21
21
  src/multineuronchat/db/__init__.py
22
22
  src/multineuronchat/db/interactionDB_human.pkl
23
23
  src/multineuronchat/db/interactionDB_human_extended.pkl
24
- src/multineuronchat/db/interactionDB_mouse.pkl
24
+ src/multineuronchat/db/interactionDB_mouse.pkl
25
+ tests/test_avg_expression.py
26
+ tests/test_communication_score.py
27
+ tests/test_correction.py
28
+ tests/test_golden_master.py
29
+ tests/test_interaction_db.py
30
+ tests/test_masks.py
31
+ tests/test_normalize.py
32
+ tests/test_object_save_load.py
33
+ tests/test_power.py
34
+ tests/test_significance.py
35
+ tests/test_utils_filter.py
@@ -0,0 +1,272 @@
1
+ """Locks for the average-expression stage (``MultiNeuronChat.py``).
2
+
3
+ ``compute_subject_specific_avg_expression`` summarizes each gene across one
4
+ subject's cells of a given type, with three mean flavors:
5
+
6
+ * ``'mean'`` — plain arithmetic mean;
7
+ * ``'tri_mean'`` — Tukey's trimean ``(q1 + 2*q2 + q3) / 4`` using
8
+ ``np.percentile`` (linear interpolation) — the default;
9
+ * ``'trim_mean'`` — ``scipy.stats.trim_mean`` (needs ``trim_mean_fraction``).
10
+
11
+ Each is pinned against a direct oracle on hand-known count columns. Also locked:
12
+ the NaN-on-too-few-cells behavior (with its ``UserWarning``), the empty
13
+ cell-type/subject combination (NaN + ``UserWarning``), and that cell types come
14
+ out **sorted**.
15
+
16
+ ``compute_avg_expression`` fans the per-subject function over subjects and groups
17
+ the results by condition, dropping subjects in neither condition. Its default
18
+ ``mean_type='trimean'`` is a *bug* (the accepted values are ``'tri_mean'`` /
19
+ ``'mean'`` / ``'trim_mean'``), so calling it with defaults raises ``ValueError``;
20
+ that broken default is deliberately locked in here, not fixed.
21
+ """
22
+
23
+ import numpy as np
24
+ import pytest
25
+ from scipy import stats
26
+
27
+ import helpers
28
+ from multineuronchat.MultiNeuronChat import (
29
+ compute_avg_expression,
30
+ compute_subject_specific_avg_expression,
31
+ )
32
+
33
+
34
+ # One subject (s1/ctrl), one cell type (A), four cells. Two genes with values
35
+ # chosen so mean != trimean != trim_mean.
36
+ _ARITH_GENES = ["G1", "G2"]
37
+ _ARITH_COUNTS = np.array(
38
+ [
39
+ [1.0, 2.0, 3.0, 10.0], # G1
40
+ [0.0, 1.0, 1.0, 2.0], # G2
41
+ ]
42
+ )
43
+
44
+
45
+ def _arith_loom(tmp_path):
46
+ return helpers.make_loom(
47
+ tmp_path / "arith.loom",
48
+ counts=_ARITH_COUNTS,
49
+ gene_names=_ARITH_GENES,
50
+ subjects=["s1", "s1", "s1", "s1"],
51
+ conditions=["ctrl", "ctrl", "ctrl", "ctrl"],
52
+ cell_types=["A", "A", "A", "A"],
53
+ )
54
+
55
+
56
+ def test_mean_matches_numpy_mean(tmp_path):
57
+ loom = _arith_loom(tmp_path)
58
+ avg = compute_subject_specific_avg_expression(
59
+ loom, subject_id="s1", subject_label_column="subject",
60
+ cell_type_label_column="cell_type", mean_type="mean",
61
+ )
62
+ expected = _ARITH_COUNTS.mean(axis=1)
63
+ np.testing.assert_allclose(avg.sel(cell_types="A").values, expected, rtol=1e-12)
64
+
65
+
66
+ def test_tri_mean_matches_percentile_oracle(tmp_path):
67
+ loom = _arith_loom(tmp_path)
68
+ avg = compute_subject_specific_avg_expression(
69
+ loom, subject_id="s1", subject_label_column="subject",
70
+ cell_type_label_column="cell_type", mean_type="tri_mean",
71
+ )
72
+ q1 = np.percentile(_ARITH_COUNTS, 25, axis=1)
73
+ q2 = np.percentile(_ARITH_COUNTS, 50, axis=1)
74
+ q3 = np.percentile(_ARITH_COUNTS, 75, axis=1)
75
+ expected = (q1 + 2 * q2 + q3) / 4
76
+ np.testing.assert_allclose(avg.sel(cell_types="A").values, expected, rtol=1e-12)
77
+
78
+
79
+ def test_tri_mean_is_the_default(tmp_path):
80
+ loom = _arith_loom(tmp_path)
81
+ default = compute_subject_specific_avg_expression(
82
+ loom, subject_id="s1", subject_label_column="subject",
83
+ cell_type_label_column="cell_type",
84
+ )
85
+ explicit = compute_subject_specific_avg_expression(
86
+ loom, subject_id="s1", subject_label_column="subject",
87
+ cell_type_label_column="cell_type", mean_type="tri_mean",
88
+ )
89
+ np.testing.assert_allclose(default.values, explicit.values, rtol=1e-12)
90
+
91
+
92
+ def test_trim_mean_matches_scipy_oracle(tmp_path):
93
+ loom = _arith_loom(tmp_path)
94
+ avg = compute_subject_specific_avg_expression(
95
+ loom, subject_id="s1", subject_label_column="subject",
96
+ cell_type_label_column="cell_type", mean_type="trim_mean",
97
+ trim_mean_fraction=0.25,
98
+ )
99
+ expected = stats.trim_mean(_ARITH_COUNTS, 0.25, axis=1)
100
+ np.testing.assert_allclose(avg.sel(cell_types="A").values, expected, rtol=1e-12)
101
+
102
+
103
+ # Richer loom: subject s1 (ctrl) has cell type A x3 and B x1; subject s2 (case)
104
+ # has A x2 and NO B. Used for threshold / empty-combo / sorting locks.
105
+ def _mixed_loom(tmp_path):
106
+ return helpers.make_loom(
107
+ tmp_path / "mixed.loom",
108
+ counts=np.array([[1.0, 2.0, 3.0, 9.0, 4.0, 6.0]]), # single gene G1
109
+ gene_names=["G1"],
110
+ subjects=["s1", "s1", "s1", "s1", "s2", "s2"],
111
+ conditions=["ctrl", "ctrl", "ctrl", "ctrl", "case", "case"],
112
+ cell_types=["A", "A", "A", "B", "A", "A"],
113
+ )
114
+
115
+
116
+ def test_min_n_cells_threshold_sets_undersized_celltype_nan(tmp_path):
117
+ loom = _mixed_loom(tmp_path)
118
+ with pytest.warns(UserWarning):
119
+ avg = compute_subject_specific_avg_expression(
120
+ loom, subject_id="s1", subject_label_column="subject",
121
+ cell_type_label_column="cell_type", mean_type="mean",
122
+ min_n_cells_threshold=2,
123
+ )
124
+ # A has 3 cells (>=2) -> computed; B has 1 cell (<2) -> NaN.
125
+ np.testing.assert_allclose(avg.sel(cell_types="A").values, [2.0], rtol=1e-12)
126
+ assert np.isnan(avg.sel(cell_types="B").values).all()
127
+
128
+
129
+ def test_missing_celltype_for_subject_is_nan_and_warns(tmp_path):
130
+ loom = _mixed_loom(tmp_path)
131
+ # s2 has cell type A but no B -> B column is NaN and a UserWarning fires.
132
+ with pytest.warns(UserWarning):
133
+ avg = compute_subject_specific_avg_expression(
134
+ loom, subject_id="s2", subject_label_column="subject",
135
+ cell_type_label_column="cell_type", mean_type="mean",
136
+ )
137
+ np.testing.assert_allclose(avg.sel(cell_types="A").values, [5.0], rtol=1e-12)
138
+ assert np.isnan(avg.sel(cell_types="B").values).all()
139
+
140
+
141
+ def test_cell_types_are_sorted(tmp_path):
142
+ loom = helpers.make_loom(
143
+ tmp_path / "sorted.loom",
144
+ counts=np.array([[1.0, 2.0, 3.0]]),
145
+ gene_names=["G1"],
146
+ subjects=["s1", "s1", "s1"],
147
+ conditions=["ctrl", "ctrl", "ctrl"],
148
+ cell_types=["C", "A", "B"], # deliberately unsorted
149
+ )
150
+ avg = compute_subject_specific_avg_expression(
151
+ loom, subject_id="s1", subject_label_column="subject",
152
+ cell_type_label_column="cell_type", mean_type="mean",
153
+ )
154
+ assert list(avg.coords["cell_types"].values) == ["A", "B", "C"]
155
+
156
+
157
+ # --------------------------------------------------------------------------- #
158
+ # Validation of the per-subject function.
159
+ # --------------------------------------------------------------------------- #
160
+ def test_missing_file_raises(tmp_path):
161
+ with pytest.raises(FileNotFoundError):
162
+ compute_subject_specific_avg_expression(
163
+ str(tmp_path / "nope.loom"), subject_id="s1",
164
+ subject_label_column="subject", cell_type_label_column="cell_type",
165
+ )
166
+
167
+
168
+ def test_non_loom_extension_raises(tmp_path):
169
+ bogus = tmp_path / "data.txt"
170
+ bogus.write_text("not a loom")
171
+ with pytest.raises(ValueError):
172
+ compute_subject_specific_avg_expression(
173
+ str(bogus), subject_id="s1",
174
+ subject_label_column="subject", cell_type_label_column="cell_type",
175
+ )
176
+
177
+
178
+ def test_bad_mean_type_raises(tmp_path):
179
+ loom = _arith_loom(tmp_path)
180
+ with pytest.raises(ValueError):
181
+ compute_subject_specific_avg_expression(
182
+ loom, subject_id="s1", subject_label_column="subject",
183
+ cell_type_label_column="cell_type", mean_type="median",
184
+ )
185
+
186
+
187
+ def test_trim_mean_without_fraction_raises(tmp_path):
188
+ loom = _arith_loom(tmp_path)
189
+ with pytest.raises(ValueError):
190
+ compute_subject_specific_avg_expression(
191
+ loom, subject_id="s1", subject_label_column="subject",
192
+ cell_type_label_column="cell_type", mean_type="trim_mean",
193
+ )
194
+
195
+
196
+ def test_empty_subject_id_raises(tmp_path):
197
+ loom = _arith_loom(tmp_path)
198
+ with pytest.raises(ValueError):
199
+ compute_subject_specific_avg_expression(
200
+ loom, subject_id="", subject_label_column="subject",
201
+ cell_type_label_column="cell_type",
202
+ )
203
+
204
+
205
+ # --------------------------------------------------------------------------- #
206
+ # compute_avg_expression orchestration.
207
+ # --------------------------------------------------------------------------- #
208
+ def _three_condition_loom(tmp_path):
209
+ # s1 -> ctrl, s2 -> case, s3 -> other (must be dropped).
210
+ return helpers.make_loom(
211
+ tmp_path / "three.loom",
212
+ counts=np.array([[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]]),
213
+ gene_names=["G1"],
214
+ subjects=["s1", "s1", "s2", "s2", "s3", "s3"],
215
+ conditions=["ctrl", "ctrl", "case", "case", "other", "other"],
216
+ cell_types=["A", "A", "A", "A", "A", "A"],
217
+ )
218
+
219
+
220
+ def test_avg_expression_groups_by_condition_and_drops_outsiders(tmp_path):
221
+ loom = _three_condition_loom(tmp_path)
222
+ result = compute_avg_expression(
223
+ path_to_loom=loom,
224
+ condition_label_column="condition",
225
+ condition_label_a="ctrl", condition_label_b="case",
226
+ subject_label_column="subject", cell_type_label_column="cell_type",
227
+ mean_type="mean", n_processes=1,
228
+ )
229
+ assert set(result.keys()) == {"ctrl", "case"}
230
+ assert set(result["ctrl"].keys()) == {"s1"}
231
+ assert set(result["case"].keys()) == {"s2"}
232
+ # s3 (condition 'other') appears nowhere.
233
+ assert "s3" not in result["ctrl"] and "s3" not in result["case"]
234
+ np.testing.assert_allclose(result["ctrl"]["s1"].sel(cell_types="A").values, [1.5], rtol=1e-12)
235
+
236
+
237
+ def test_avg_expression_default_mean_type_is_broken_and_raises(tmp_path):
238
+ """Locked bug: the default ``mean_type='trimean'`` is not an accepted value
239
+ (the valid spelling is ``'tri_mean'``), so a defaults call raises ValueError."""
240
+ loom = _three_condition_loom(tmp_path)
241
+ with pytest.raises(ValueError):
242
+ compute_avg_expression(
243
+ path_to_loom=loom,
244
+ condition_label_column="condition",
245
+ condition_label_a="ctrl", condition_label_b="case",
246
+ subject_label_column="subject", cell_type_label_column="cell_type",
247
+ n_processes=1,
248
+ )
249
+
250
+
251
+ def test_avg_expression_equal_conditions_raises(tmp_path):
252
+ loom = _three_condition_loom(tmp_path)
253
+ with pytest.raises(ValueError):
254
+ compute_avg_expression(
255
+ path_to_loom=loom,
256
+ condition_label_column="condition",
257
+ condition_label_a="ctrl", condition_label_b="ctrl",
258
+ subject_label_column="subject", cell_type_label_column="cell_type",
259
+ mean_type="mean", n_processes=1,
260
+ )
261
+
262
+
263
+ def test_avg_expression_non_positive_n_processes_raises(tmp_path):
264
+ loom = _three_condition_loom(tmp_path)
265
+ with pytest.raises(ValueError):
266
+ compute_avg_expression(
267
+ path_to_loom=loom,
268
+ condition_label_column="condition",
269
+ condition_label_a="ctrl", condition_label_b="case",
270
+ subject_label_column="subject", cell_type_label_column="cell_type",
271
+ mean_type="mean", n_processes=0,
272
+ )