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.
- multineuronchat/InteractionDB/InteractionDB.py +67 -0
- multineuronchat/InteractionDB/InteractionDBRow.py +50 -0
- multineuronchat/InteractionDB/__init__.py +4 -0
- multineuronchat/MultiNeuronChat.py +447 -0
- multineuronchat/MultiNeuronChatObject.py +959 -0
- multineuronchat/__init__.py +28 -0
- multineuronchat/db/__init__.py +0 -0
- multineuronchat/loompy_utils.py +66 -0
- multineuronchat/masks.py +358 -0
- multineuronchat/normalize.py +177 -0
- multineuronchat/utils.py +159 -0
- multineuronchat/visualize.py +1029 -0
- multineuronchat-2025.11.10.dev0.dist-info/METADATA +117 -0
- multineuronchat-2025.11.10.dev0.dist-info/RECORD +17 -0
- multineuronchat-2025.11.10.dev0.dist-info/WHEEL +5 -0
- multineuronchat-2025.11.10.dev0.dist-info/licenses/LICENSE +674 -0
- multineuronchat-2025.11.10.dev0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,1029 @@
|
|
|
1
|
+
import math
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
|
|
5
|
+
from decimal import Decimal
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import xarray as xr
|
|
9
|
+
|
|
10
|
+
import pandas as pd
|
|
11
|
+
|
|
12
|
+
import seaborn as sns
|
|
13
|
+
import matplotlib as mpl
|
|
14
|
+
import matplotlib.pyplot as plt
|
|
15
|
+
import matplotlib.patches as patches
|
|
16
|
+
from matplotlib.pyplot import Line2D
|
|
17
|
+
|
|
18
|
+
from . import MultiNeuronChatObject
|
|
19
|
+
|
|
20
|
+
def plot_circle_plot(
|
|
21
|
+
axs: plt.Axes,
|
|
22
|
+
cell_cell_counts: dict[tuple[str, str], int],
|
|
23
|
+
cell_types: list[str],
|
|
24
|
+
cell_type_colors: dict[str, str] | list[str] | None = None,
|
|
25
|
+
cell_type_labels_short: dict[str, str] | None = None,
|
|
26
|
+
radius_of_nodes: float = 0.05,
|
|
27
|
+
show_labels: bool = False,
|
|
28
|
+
label_font_size: int = 12,
|
|
29
|
+
legend_axs: plt.Axes | None = None,
|
|
30
|
+
legend_marker_size: float = 13,
|
|
31
|
+
legend_marker_edge_width: float = 1,
|
|
32
|
+
legend_font_size: float = 12,
|
|
33
|
+
legend_n_cols: int = 1,
|
|
34
|
+
title: str | None = None,
|
|
35
|
+
title_font_size: float = 12,
|
|
36
|
+
arrow_thickness_factor: float = 2,
|
|
37
|
+
arrow_head_length: float = 5,
|
|
38
|
+
arrow_head_width: float = 2.5,
|
|
39
|
+
arrow_legend_axs: plt.Axes | None = None,
|
|
40
|
+
arrow_legend_font_size: float = 12,
|
|
41
|
+
arrow_legend_n_cols: int = 1,
|
|
42
|
+
arrow_normalization_factor: float | None = None,
|
|
43
|
+
show_counts: bool = False
|
|
44
|
+
) -> None:
|
|
45
|
+
"""
|
|
46
|
+
Create a circle plot showing the interactions between different cell-types.
|
|
47
|
+
|
|
48
|
+
:param axs: Matplotlib Axes object where the circle plot will be drawn.
|
|
49
|
+
:param cell_cell_counts: Dictionary with keys as tuples of (source_cell_type, receiver_cell_type) and values as the number of connections.
|
|
50
|
+
:param cell_types: List of cell-types to be included in the plot.
|
|
51
|
+
:param cell_type_colors: Dictionary or list of colors for each cell-type. If None, default colors will be used.
|
|
52
|
+
:param cell_type_labels_short: Dictionary mapping cell-types to their short labels for display.
|
|
53
|
+
:param radius_of_nodes: Radius of the nodes representing cell-types.
|
|
54
|
+
:param show_labels: Whether to show cell-type labels on the plot.
|
|
55
|
+
:param label_font_size: Font size for the cell-type labels.
|
|
56
|
+
:param legend_axs: Matplotlib Axes object where the legend will be drawn. If None, no legend is drawn.
|
|
57
|
+
:param legend_marker_size: Size of the markers in the legend.
|
|
58
|
+
:param legend_marker_edge_width: Edge width of the markers in the legend.
|
|
59
|
+
:param legend_font_size: Font size for the legend text.
|
|
60
|
+
:param legend_n_cols: Number of columns in the legend.
|
|
61
|
+
:param title: Title of the plot. If None, no title is set.
|
|
62
|
+
:param title_font_size: Font size for the title.
|
|
63
|
+
:param arrow_thickness_factor: Factor to scale the thickness of the arrows representing connections.
|
|
64
|
+
:param arrow_head_length: Length of the arrow heads.
|
|
65
|
+
:param arrow_head_width: Width of the arrow heads.
|
|
66
|
+
:param arrow_legend_axs: Matplotlib Axes object where the arrow legend will be drawn. If None, no arrow legend is drawn.
|
|
67
|
+
:param arrow_legend_font_size: Font size for the arrow legend text.
|
|
68
|
+
:param arrow_legend_n_cols: Number of columns in the arrow legend.
|
|
69
|
+
:param arrow_normalization_factor: Normalization factor for the arrow thickness. If None, the maximum number of connections is used.
|
|
70
|
+
:param show_counts: Whether to show incoming and outgoing connection counts for each cell-type.
|
|
71
|
+
|
|
72
|
+
:return: None
|
|
73
|
+
"""
|
|
74
|
+
if cell_type_colors is None:
|
|
75
|
+
cell_type_colors = list(sns.color_palette('husl', len(cell_types)))
|
|
76
|
+
elif len(cell_type_colors) != len(cell_types):
|
|
77
|
+
raise ValueError('The number of colors must match the number of cell-types.')
|
|
78
|
+
|
|
79
|
+
# Compute coordinates on a unit-circle for each cell-type
|
|
80
|
+
angle: float = 2 * np.pi / len(cell_types)
|
|
81
|
+
shift: float = 2 * (len(cell_types) * radius_of_nodes) / np.pi
|
|
82
|
+
cell_type_to_coordinates: dict[str, np.array] = {
|
|
83
|
+
cell_type: np.array([shift * np.cos(i * angle), shift * np.sin(i * angle)])
|
|
84
|
+
for i, cell_type in enumerate(cell_types)
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
for i, cell_type in enumerate(cell_types):
|
|
88
|
+
color = cell_type_colors[i] if type(cell_type_colors) is list else cell_type_colors[cell_type]
|
|
89
|
+
|
|
90
|
+
node = plt.Circle(
|
|
91
|
+
xy=cell_type_to_coordinates[cell_type],
|
|
92
|
+
radius=radius_of_nodes,
|
|
93
|
+
color=color
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
axs.add_patch(node)
|
|
97
|
+
|
|
98
|
+
# Only iterate through the connections if the cell_cell_counts is not empty
|
|
99
|
+
if len(cell_cell_counts) > 0:
|
|
100
|
+
if arrow_normalization_factor is None:
|
|
101
|
+
arrow_normalization_factor: float = max(cell_cell_counts.values())
|
|
102
|
+
|
|
103
|
+
for source_cell_type, receiver_cell_type in cell_cell_counts.keys():
|
|
104
|
+
source_pos: np.array = cell_type_to_coordinates[source_cell_type]
|
|
105
|
+
receiver_pos: np.array = cell_type_to_coordinates[receiver_cell_type]
|
|
106
|
+
|
|
107
|
+
# Get the number of connections
|
|
108
|
+
number_of_connections: float = cell_cell_counts[(source_cell_type, receiver_cell_type)]
|
|
109
|
+
|
|
110
|
+
if number_of_connections > arrow_normalization_factor:
|
|
111
|
+
# Send out a warning that the normalization factor is too low
|
|
112
|
+
warnings.warn(f'The number of connections between {source_cell_type} and {receiver_cell_type} is higher than the normalization factor. Please choose a value higher or equal to {number_of_connections} to ensure that the arrow thickness is proportional to the number of connections.')
|
|
113
|
+
|
|
114
|
+
normalized_number_of_connections: float = (number_of_connections / arrow_normalization_factor) * arrow_thickness_factor
|
|
115
|
+
|
|
116
|
+
if source_cell_type == receiver_cell_type:
|
|
117
|
+
length_of_origin_vector: float = np.linalg.norm(source_pos)
|
|
118
|
+
|
|
119
|
+
normalized_radius_vector: np.array = (source_pos / length_of_origin_vector) * radius_of_nodes
|
|
120
|
+
|
|
121
|
+
angle_to_rotate: float = 80 * (np.pi / 180)
|
|
122
|
+
x, y = normalized_radius_vector
|
|
123
|
+
rotation_clockwise_vector: np.array = np.array([
|
|
124
|
+
np.cos(-angle_to_rotate) * x - np.sin(-angle_to_rotate) * y,
|
|
125
|
+
np.sin(-angle_to_rotate) * x + np.cos(-angle_to_rotate) * y
|
|
126
|
+
])
|
|
127
|
+
rotation_anticlockwise_vector: np.array = np.array([
|
|
128
|
+
np.cos(angle_to_rotate) * x - np.sin(angle_to_rotate) * y,
|
|
129
|
+
np.sin(angle_to_rotate) * x + np.cos(angle_to_rotate) * y
|
|
130
|
+
])
|
|
131
|
+
|
|
132
|
+
start_arrow_pos: np.array = source_pos + rotation_anticlockwise_vector
|
|
133
|
+
end_arrow_pos: np.array = receiver_pos + rotation_clockwise_vector
|
|
134
|
+
|
|
135
|
+
edge: patches.FancyArrowPatch = patches.FancyArrowPatch(
|
|
136
|
+
posA=start_arrow_pos,
|
|
137
|
+
posB=end_arrow_pos,
|
|
138
|
+
arrowstyle=f'-|>,head_length={arrow_head_length},head_width={arrow_head_width}',
|
|
139
|
+
connectionstyle=f'arc3,rad={-2}',
|
|
140
|
+
color='#000000',
|
|
141
|
+
linewidth=normalized_number_of_connections,
|
|
142
|
+
)
|
|
143
|
+
|
|
144
|
+
else:
|
|
145
|
+
connecting_vector: np.array = receiver_pos - source_pos
|
|
146
|
+
length: float = np.linalg.norm(connecting_vector)
|
|
147
|
+
radius_vector: np.array = (connecting_vector / length) * radius_of_nodes
|
|
148
|
+
|
|
149
|
+
source_arrow_pos: np.array = source_pos + radius_vector
|
|
150
|
+
receiver_arrow_pos: np.array = receiver_pos - radius_vector
|
|
151
|
+
|
|
152
|
+
# Determine direction to change angle of the circle
|
|
153
|
+
theta1: float = np.arctan2(source_arrow_pos[1], source_arrow_pos[0])
|
|
154
|
+
theta2: float = np.arctan2(receiver_arrow_pos[1], receiver_arrow_pos[0])
|
|
155
|
+
delta_theta: float = theta2 - theta1
|
|
156
|
+
|
|
157
|
+
if delta_theta > np.pi:
|
|
158
|
+
delta_theta = delta_theta - 2 * np.pi
|
|
159
|
+
elif delta_theta < -np.pi:
|
|
160
|
+
delta_theta = delta_theta + 2 * np.pi
|
|
161
|
+
|
|
162
|
+
clockwise = 1 if delta_theta < 0 else -2
|
|
163
|
+
|
|
164
|
+
# Check if the source cell-type is adjacent to the receiver cell; then ignore the arc
|
|
165
|
+
index_cell_type_a = cell_types.index(source_cell_type)
|
|
166
|
+
index_cell_type_b = cell_types.index(receiver_cell_type)
|
|
167
|
+
if index_cell_type_a == index_cell_type_b + 1 or index_cell_type_a == index_cell_type_b - 1:
|
|
168
|
+
clockwise = -1
|
|
169
|
+
|
|
170
|
+
edge: patches.FancyArrowPatch = patches.FancyArrowPatch(
|
|
171
|
+
posA=source_arrow_pos,
|
|
172
|
+
posB=receiver_arrow_pos,
|
|
173
|
+
arrowstyle=f'-|>,head_length={arrow_head_length},head_width={arrow_head_width}',
|
|
174
|
+
connectionstyle=f'arc3,rad={clockwise * 0.25}',
|
|
175
|
+
color='#000000',
|
|
176
|
+
linewidth=normalized_number_of_connections,
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
axs.add_patch(edge)
|
|
180
|
+
|
|
181
|
+
if show_labels:
|
|
182
|
+
for cell_type in cell_types:
|
|
183
|
+
label = cell_type_labels_short[cell_type] if cell_type_labels_short is not None else cell_type
|
|
184
|
+
|
|
185
|
+
axs.text(
|
|
186
|
+
x=cell_type_to_coordinates[cell_type][0],
|
|
187
|
+
y=cell_type_to_coordinates[cell_type][1],
|
|
188
|
+
s=label,
|
|
189
|
+
ha='center',
|
|
190
|
+
va='center',
|
|
191
|
+
fontsize=label_font_size,
|
|
192
|
+
)
|
|
193
|
+
|
|
194
|
+
# If user wants the counts, show them above (incoming) and below (outgoing)
|
|
195
|
+
if show_counts:
|
|
196
|
+
# Sum of incoming
|
|
197
|
+
incoming_count = sum(
|
|
198
|
+
count
|
|
199
|
+
for (src, dst), count in cell_cell_counts.items()
|
|
200
|
+
if dst == cell_type
|
|
201
|
+
)
|
|
202
|
+
# Sum of outgoing
|
|
203
|
+
outgoing_count = sum(
|
|
204
|
+
count
|
|
205
|
+
for (src, dst), count in cell_cell_counts.items()
|
|
206
|
+
if src == cell_type
|
|
207
|
+
)
|
|
208
|
+
|
|
209
|
+
# Above the label (incoming)
|
|
210
|
+
axs.text(
|
|
211
|
+
x=cell_type_to_coordinates[cell_type][0],
|
|
212
|
+
y=cell_type_to_coordinates[cell_type][1] + radius_of_nodes*0.4, # small offset upward
|
|
213
|
+
s=str(incoming_count),
|
|
214
|
+
ha='center',
|
|
215
|
+
va='bottom',
|
|
216
|
+
fontsize=label_font_size,
|
|
217
|
+
color='black'
|
|
218
|
+
)
|
|
219
|
+
# Below the label (outgoing)
|
|
220
|
+
axs.text(
|
|
221
|
+
x=cell_type_to_coordinates[cell_type][0],
|
|
222
|
+
y=cell_type_to_coordinates[cell_type][1] - radius_of_nodes*0.4, # small offset downward
|
|
223
|
+
s=str(outgoing_count),
|
|
224
|
+
ha='center',
|
|
225
|
+
va='top',
|
|
226
|
+
fontsize=label_font_size,
|
|
227
|
+
color='black'
|
|
228
|
+
)
|
|
229
|
+
|
|
230
|
+
axs.set_xlim(-shift - radius_of_nodes * 2.5, shift + radius_of_nodes * 2.5)
|
|
231
|
+
axs.set_ylim(-shift - radius_of_nodes * 2.5, shift + radius_of_nodes * 2.5)
|
|
232
|
+
|
|
233
|
+
axs.axis('off')
|
|
234
|
+
|
|
235
|
+
if title is not None:
|
|
236
|
+
axs.set_title(title, fontsize=title_font_size)
|
|
237
|
+
|
|
238
|
+
if legend_axs is not None:
|
|
239
|
+
legend_patches = []
|
|
240
|
+
|
|
241
|
+
for i, cell_type in enumerate(cell_types):
|
|
242
|
+
if cell_type_labels_short is not None:
|
|
243
|
+
label = f'{cell_type_labels_short[cell_type]} - {cell_type}'
|
|
244
|
+
else:
|
|
245
|
+
label = cell_type
|
|
246
|
+
|
|
247
|
+
color = cell_type_colors[i] if type(cell_type_colors) is list else cell_type_colors[cell_type]
|
|
248
|
+
|
|
249
|
+
legend_patch = Line2D(
|
|
250
|
+
[0], [0],
|
|
251
|
+
label=label,
|
|
252
|
+
marker='o',
|
|
253
|
+
markersize=legend_marker_size,
|
|
254
|
+
markerfacecolor=color,
|
|
255
|
+
markeredgecolor='black',
|
|
256
|
+
markeredgewidth=legend_marker_edge_width,
|
|
257
|
+
linestyle='',
|
|
258
|
+
)
|
|
259
|
+
legend_patches.append(legend_patch)
|
|
260
|
+
|
|
261
|
+
# center legend
|
|
262
|
+
legend_axs.legend(handles=legend_patches, loc='center', prop={'size': legend_font_size}, ncol=legend_n_cols)
|
|
263
|
+
|
|
264
|
+
legend_axs.axis('off')
|
|
265
|
+
|
|
266
|
+
if arrow_legend_axs is not None:
|
|
267
|
+
arrow_legend_patches = []
|
|
268
|
+
|
|
269
|
+
def get_integer_steps(min_value, max_value, max_steps=5):
|
|
270
|
+
unique_values = np.arange(min_value, max_value + 1)
|
|
271
|
+
num_unique_values = len(unique_values)
|
|
272
|
+
if num_unique_values <= max_steps:
|
|
273
|
+
steps = unique_values
|
|
274
|
+
else:
|
|
275
|
+
indices = np.linspace(0, num_unique_values - 1, num=max_steps)
|
|
276
|
+
indices = np.round(indices).astype(int)
|
|
277
|
+
steps = unique_values[indices]
|
|
278
|
+
return steps.tolist()
|
|
279
|
+
|
|
280
|
+
# Check if there are even connections
|
|
281
|
+
if cell_cell_counts.values():
|
|
282
|
+
max_number_of_connections = max(cell_cell_counts.values())
|
|
283
|
+
min_number_of_connections = min(cell_cell_counts.values())
|
|
284
|
+
|
|
285
|
+
# Generate integer steps for the legend
|
|
286
|
+
arrow_step_size = get_integer_steps(
|
|
287
|
+
min_number_of_connections,
|
|
288
|
+
max_number_of_connections,
|
|
289
|
+
max_steps=5
|
|
290
|
+
)
|
|
291
|
+
|
|
292
|
+
for arrow_number in arrow_step_size:
|
|
293
|
+
normalized_arrow_number = (arrow_number / arrow_normalization_factor) * arrow_thickness_factor
|
|
294
|
+
|
|
295
|
+
arrow_legend_patch = Line2D(
|
|
296
|
+
[0], [0],
|
|
297
|
+
color='#000000',
|
|
298
|
+
linewidth=normalized_arrow_number,
|
|
299
|
+
linestyle='-',
|
|
300
|
+
)
|
|
301
|
+
arrow_legend_patches.append(arrow_legend_patch)
|
|
302
|
+
|
|
303
|
+
arrow_legend_axs.legend(
|
|
304
|
+
arrow_legend_patches,
|
|
305
|
+
[f'{int(arrow_number)}' for arrow_number in arrow_step_size],
|
|
306
|
+
loc='center',
|
|
307
|
+
prop={'size': arrow_legend_font_size},
|
|
308
|
+
ncol=arrow_legend_n_cols
|
|
309
|
+
)
|
|
310
|
+
|
|
311
|
+
arrow_legend_axs.axis('off')
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def plot_p_value_differential_communication_circle_plot(
|
|
315
|
+
axs: plt.Axes,
|
|
316
|
+
mnc_object: MultiNeuronChatObject,
|
|
317
|
+
significance_test_to_use: str,
|
|
318
|
+
significance_threshold: float = 0.05,
|
|
319
|
+
use_adj_p_values: bool = True,
|
|
320
|
+
cell_types: list[str] | None = None,
|
|
321
|
+
cell_type_colors: dict[str, str] | list[str] | None = None,
|
|
322
|
+
cell_type_labels_short: dict[str, str] | None = None,
|
|
323
|
+
radius_of_nodes: float = 0.05,
|
|
324
|
+
show_labels: bool = False,
|
|
325
|
+
label_font_size: int = 12,
|
|
326
|
+
legend_axs: plt.Axes | None = None,
|
|
327
|
+
legend_marker_size: float = 13,
|
|
328
|
+
legend_marker_edge_width: float = 1,
|
|
329
|
+
legend_font_size: float = 12,
|
|
330
|
+
legend_n_cols: int = 1,
|
|
331
|
+
title: str | None = None,
|
|
332
|
+
title_font_size: float = 12,
|
|
333
|
+
arrow_thickness_factor: float = 2,
|
|
334
|
+
arrow_head_length: float = 5,
|
|
335
|
+
arrow_head_width: float = 2.5,
|
|
336
|
+
arrow_legend_axs: plt.Axes | None = None,
|
|
337
|
+
arrow_legend_font_size: float = 12,
|
|
338
|
+
arrow_legend_n_cols: int = 1,
|
|
339
|
+
arrow_normalization_factor: float | None = None,
|
|
340
|
+
show_counts: bool = False
|
|
341
|
+
):
|
|
342
|
+
"""
|
|
343
|
+
Create a circle plot showing the significant interactions between different cell-types based on p-values from a MultiNeuronChatObject.
|
|
344
|
+
|
|
345
|
+
:param axs: Matplotlib Axes object where the circle plot will be drawn.
|
|
346
|
+
:param mnc_object: MultiNeuronChatObject containing the p-values for differential communication analysis.
|
|
347
|
+
:param significance_test_to_use: Name of the significance test to use for plotting (e.g., 'KS', 'Anderson', etc.).
|
|
348
|
+
:param significance_threshold: Significance threshold for p-values to consider an interaction significant (default: 0.05).
|
|
349
|
+
:param use_adj_p_values: Whether to use adjusted p-values (default: True).
|
|
350
|
+
:param cell_types: List of cell-types to be included in the plot.
|
|
351
|
+
:param cell_type_colors: Dictionary or list of colors for each cell-type. If None, default colors will be used.
|
|
352
|
+
:param cell_type_labels_short: Dictionary mapping cell-types to their short labels for display.
|
|
353
|
+
:param radius_of_nodes: Radius of the nodes representing cell-types.
|
|
354
|
+
:param show_labels: Whether to show cell-type labels on the plot.
|
|
355
|
+
:param label_font_size: Font size for the cell-type labels.
|
|
356
|
+
:param legend_axs: Matplotlib Axes object where the legend will be drawn. If None, no legend is drawn.
|
|
357
|
+
:param legend_marker_size: Size of the markers in the legend.
|
|
358
|
+
:param legend_marker_edge_width: Edge width of the markers in the legend.
|
|
359
|
+
:param legend_font_size: Font size for the legend text.
|
|
360
|
+
:param legend_n_cols: Number of columns in the legend.
|
|
361
|
+
:param title: Title of the plot. If None, no title is set.
|
|
362
|
+
:param title_font_size: Font size for the title.
|
|
363
|
+
:param arrow_thickness_factor: Factor to scale the thickness of the arrows representing connections.
|
|
364
|
+
:param arrow_head_length: Length of the arrow heads.
|
|
365
|
+
:param arrow_head_width: Width of the arrow heads.
|
|
366
|
+
:param arrow_legend_axs: Matplotlib Axes object where the arrow legend will be drawn. If None, no arrow legend is drawn.
|
|
367
|
+
:param arrow_legend_font_size: Font size for the arrow legend text.
|
|
368
|
+
:param arrow_legend_n_cols: Number of columns in the arrow legend.
|
|
369
|
+
:param arrow_normalization_factor: Normalization factor for the arrow thickness. If None, the maximum number of connections is used.
|
|
370
|
+
:param show_counts: Whether to show incoming and outgoing connection counts for each cell-type.
|
|
371
|
+
|
|
372
|
+
:return: None
|
|
373
|
+
"""
|
|
374
|
+
if cell_types is None:
|
|
375
|
+
# Get the number of cell-types
|
|
376
|
+
cell_types: list[str] = list(set(mnc_object.source_cell_types).union(set(mnc_object.receiver_cell_types)))
|
|
377
|
+
cell_types.sort()
|
|
378
|
+
else:
|
|
379
|
+
# Check if cell_types are valid
|
|
380
|
+
if not all([cell_type in mnc_object.source_cell_types or cell_type in mnc_object.receiver_cell_types for cell_type in cell_types]):
|
|
381
|
+
not_present_cell_types: set[str] = set(cell_types).difference(set(mnc_object.source_cell_types).union(set(mnc_object.receiver_cell_types)))
|
|
382
|
+
raise ValueError(f'All cell-types must be present in the MultiNeuronChatObject!'
|
|
383
|
+
f'The following cell_types are not present: {not_present_cell_types}')
|
|
384
|
+
|
|
385
|
+
|
|
386
|
+
if use_adj_p_values and significance_test_to_use not in mnc_object.p_values_adj.keys():
|
|
387
|
+
raise ValueError(f'The significance test "{significance_test_to_use}" was not run.')
|
|
388
|
+
elif significance_test_to_use not in mnc_object.p_values.keys():
|
|
389
|
+
raise ValueError(f'The significance test "{significance_test_to_use}" was not run.')
|
|
390
|
+
|
|
391
|
+
# Get the p-values to use for creating the graph
|
|
392
|
+
if use_adj_p_values:
|
|
393
|
+
p_values: xr.DataArray = mnc_object.p_values_adj[significance_test_to_use]
|
|
394
|
+
else:
|
|
395
|
+
p_values: xr.DataArray = mnc_object.p_values[significance_test_to_use]
|
|
396
|
+
|
|
397
|
+
significant_interactions: xr.DataArray = p_values < significance_threshold
|
|
398
|
+
|
|
399
|
+
# Sum over the ligand target interactions to get the counts of significant interactions
|
|
400
|
+
significant_interactions_counts: xr.DataArray = significant_interactions.sum(dim='interaction')
|
|
401
|
+
significant_interactions_counts_idx: np.array = np.where(significant_interactions_counts > 0)
|
|
402
|
+
|
|
403
|
+
# Get the cell-cell counts
|
|
404
|
+
source_receiver_pair_counts_dict: dict[tuple[str, str], int] = {
|
|
405
|
+
(significant_interactions.coords['source'][source_cell_idx].values.item(), significant_interactions.coords['receiver'][receiver_cell_idx].values.item()): significant_interactions_counts[source_cell_idx, receiver_cell_idx].values.item()
|
|
406
|
+
for source_cell_idx, receiver_cell_idx in zip(*significant_interactions_counts_idx)
|
|
407
|
+
}
|
|
408
|
+
|
|
409
|
+
# Get the cell-type colors
|
|
410
|
+
plot_circle_plot(
|
|
411
|
+
axs=axs,
|
|
412
|
+
cell_cell_counts=source_receiver_pair_counts_dict,
|
|
413
|
+
cell_types=cell_types,
|
|
414
|
+
cell_type_colors=cell_type_colors,
|
|
415
|
+
cell_type_labels_short=cell_type_labels_short,
|
|
416
|
+
radius_of_nodes=radius_of_nodes,
|
|
417
|
+
show_labels=show_labels,
|
|
418
|
+
label_font_size=label_font_size,
|
|
419
|
+
legend_axs=legend_axs,
|
|
420
|
+
legend_marker_size=legend_marker_size,
|
|
421
|
+
legend_marker_edge_width=legend_marker_edge_width,
|
|
422
|
+
legend_font_size=legend_font_size,
|
|
423
|
+
legend_n_cols=legend_n_cols,
|
|
424
|
+
title=title,
|
|
425
|
+
title_font_size=title_font_size,
|
|
426
|
+
arrow_thickness_factor=arrow_thickness_factor,
|
|
427
|
+
arrow_legend_axs=arrow_legend_axs,
|
|
428
|
+
arrow_legend_font_size=arrow_legend_font_size,
|
|
429
|
+
arrow_head_width=arrow_head_width,
|
|
430
|
+
arrow_head_length=arrow_head_length,
|
|
431
|
+
arrow_legend_n_cols=arrow_legend_n_cols,
|
|
432
|
+
arrow_normalization_factor=arrow_normalization_factor,
|
|
433
|
+
show_counts=show_counts
|
|
434
|
+
)
|
|
435
|
+
|
|
436
|
+
|
|
437
|
+
def plot_wasserstein_ranked_differential_communication_circle_plot(
|
|
438
|
+
axs: plt.Axes,
|
|
439
|
+
wasserstein_distances: xr.DataArray,
|
|
440
|
+
top_n_edges: int = 10,
|
|
441
|
+
cell_types: list[str] | None = None,
|
|
442
|
+
cell_type_colors: dict[str, str] | list[str] | None = None,
|
|
443
|
+
cell_type_labels_short: dict[str, str] | None = None,
|
|
444
|
+
radius_of_nodes: float = 0.05,
|
|
445
|
+
show_labels: bool = False,
|
|
446
|
+
label_font_size: int = 12,
|
|
447
|
+
legend_axs: plt.Axes | None = None,
|
|
448
|
+
legend_marker_size: float = 13,
|
|
449
|
+
legend_marker_edge_width: float = 1,
|
|
450
|
+
legend_font_size: float = 12,
|
|
451
|
+
legend_n_cols: int = 1,
|
|
452
|
+
title: str | None = None,
|
|
453
|
+
title_font_size: float = 12,
|
|
454
|
+
arrow_thickness_factor: float = 2,
|
|
455
|
+
arrow_head_length: float = 5,
|
|
456
|
+
arrow_head_width: float = 2.5,
|
|
457
|
+
arrow_legend_axs: plt.Axes | None = None,
|
|
458
|
+
arrow_legend_font_size: float = 12,
|
|
459
|
+
arrow_legend_n_cols: int = 1,
|
|
460
|
+
arrow_normalization_factor: float | None = None,
|
|
461
|
+
show_counts: bool = False,
|
|
462
|
+
):
|
|
463
|
+
"""
|
|
464
|
+
Create a circle plot showing the top N interactions between different cell-types based on Wasserstein distances.
|
|
465
|
+
|
|
466
|
+
:param axs: Matplotlib Axes object where the circle plot will be drawn.
|
|
467
|
+
:param wasserstein_distances: xarray DataArray containing the Wasserstein distances between cell-types.
|
|
468
|
+
:param top_n_edges: Number of top interactions (edges) to display in the plot (default: 10).
|
|
469
|
+
:param cell_types: List of cell-types to be included in the plot.
|
|
470
|
+
:param cell_type_colors: Dictionary or list of colors for each cell-type. If None, default colors will be used.
|
|
471
|
+
:param cell_type_labels_short: Dictionary mapping cell-types to their short labels for display.
|
|
472
|
+
:param radius_of_nodes: Radius of the nodes representing cell-types.
|
|
473
|
+
:param show_labels: Whether to show cell-type labels on the plot.
|
|
474
|
+
:param label_font_size: Font size for the cell-type labels.
|
|
475
|
+
:param legend_axs: Matplotlib Axes object where the legend will be drawn. If None, no legend is drawn.
|
|
476
|
+
:param legend_marker_size: Size of the markers in the legend.
|
|
477
|
+
:param legend_marker_edge_width: Edge width of the markers in the legend.
|
|
478
|
+
:param legend_font_size: Font size for the legend text.
|
|
479
|
+
:param legend_n_cols: Number of columns in the legend.
|
|
480
|
+
:param title: Title of the plot. If None, no title is set.
|
|
481
|
+
:param title_font_size: Font size for the title.
|
|
482
|
+
:param arrow_thickness_factor: Factor to scale the thickness of the arrows representing connections.
|
|
483
|
+
:param arrow_head_length: Length of the arrow heads.
|
|
484
|
+
:param arrow_head_width: Width of the arrow heads.
|
|
485
|
+
:param arrow_legend_axs: Matplotlib Axes object where the arrow legend will be drawn. If None, no arrow legend is drawn.
|
|
486
|
+
:param arrow_legend_font_size: Font size for the arrow legend text.
|
|
487
|
+
:param arrow_legend_n_cols: Number of columns in the arrow legend.
|
|
488
|
+
:param arrow_normalization_factor: Normalization factor for the arrow thickness. If None, the maximum number of connections is used.
|
|
489
|
+
:param show_counts: Whether to show incoming and outgoing connection counts for each cell-type.
|
|
490
|
+
|
|
491
|
+
:return: None
|
|
492
|
+
"""
|
|
493
|
+
if cell_types is None:
|
|
494
|
+
# Get the number of cell-types
|
|
495
|
+
cell_types: list[str] = list(set(wasserstein_distances.coords['source'].values.tolist()).union(
|
|
496
|
+
set(wasserstein_distances.coords['receiver'].values.tolist())))
|
|
497
|
+
cell_types.sort()
|
|
498
|
+
else:
|
|
499
|
+
# Check if cell_types are valid
|
|
500
|
+
if not all([cell_type in wasserstein_distances.coords['source'].values.tolist() or cell_type in
|
|
501
|
+
wasserstein_distances.coords['receiver'].values.tolist() for cell_type in cell_types]):
|
|
502
|
+
not_present_cell_types: set[str] = set(cell_types).difference(
|
|
503
|
+
set(wasserstein_distances.coords['source'].values.tolist()).union(
|
|
504
|
+
set(wasserstein_distances.coords['receiver'].values.tolist())))
|
|
505
|
+
raise ValueError(f'All cell-types must be present in the MultiNeuronChatObject!'
|
|
506
|
+
f'The following cell_types are not present: {not_present_cell_types}')
|
|
507
|
+
|
|
508
|
+
if cell_type_colors is None:
|
|
509
|
+
cell_type_colors = list(sns.color_palette('husl', len(cell_types)))
|
|
510
|
+
elif len(cell_type_colors) != len(cell_types):
|
|
511
|
+
raise ValueError('The number of colors must match the number of cell-types.')
|
|
512
|
+
|
|
513
|
+
# Globally sorted wasserstein distances indices
|
|
514
|
+
sorted_wasserstein_distances_idxs = np.argsort(wasserstein_distances.values.flatten())[::-1]
|
|
515
|
+
# Remove NaNs
|
|
516
|
+
sorted_wasserstein_distances_idxs = sorted_wasserstein_distances_idxs[
|
|
517
|
+
~np.isnan(wasserstein_distances.values.flatten()[sorted_wasserstein_distances_idxs])]
|
|
518
|
+
|
|
519
|
+
# Get the top_n_edges indices
|
|
520
|
+
top_n_edges_idxs = sorted_wasserstein_distances_idxs[:top_n_edges]
|
|
521
|
+
|
|
522
|
+
# Unravel the indices
|
|
523
|
+
source_idxs, receiver_idxs, interaction_idxs = np.unravel_index(top_n_edges_idxs, wasserstein_distances.shape)
|
|
524
|
+
|
|
525
|
+
# Get the source and receiver cell-types
|
|
526
|
+
source_cell_types = wasserstein_distances.coords['source'].values[source_idxs]
|
|
527
|
+
receiver_cell_types = wasserstein_distances.coords['receiver'].values[receiver_idxs]
|
|
528
|
+
interaction_cell_types = wasserstein_distances.coords['interaction'].values[interaction_idxs]
|
|
529
|
+
|
|
530
|
+
df: pd.DataFrame = pd.DataFrame({
|
|
531
|
+
'source': source_cell_types,
|
|
532
|
+
'receiver': receiver_cell_types,
|
|
533
|
+
'interaction': interaction_cell_types,
|
|
534
|
+
'wasserstein distance': wasserstein_distances.values.flatten()[top_n_edges_idxs]
|
|
535
|
+
})
|
|
536
|
+
|
|
537
|
+
# Count the source-receiver pairs
|
|
538
|
+
source_receiver_pair_counts: pd.Series = df.groupby(['source', 'receiver']).size()
|
|
539
|
+
|
|
540
|
+
# Convert series of source-receiver pair counts to a dictionary of tuples to int
|
|
541
|
+
source_receiver_pair_counts_dict: dict[tuple[str, str], int] = source_receiver_pair_counts.to_dict()
|
|
542
|
+
|
|
543
|
+
plot_circle_plot(
|
|
544
|
+
axs=axs,
|
|
545
|
+
cell_cell_counts=source_receiver_pair_counts_dict,
|
|
546
|
+
cell_types=cell_types,
|
|
547
|
+
cell_type_colors=cell_type_colors,
|
|
548
|
+
cell_type_labels_short=cell_type_labels_short,
|
|
549
|
+
radius_of_nodes=radius_of_nodes,
|
|
550
|
+
show_labels=show_labels,
|
|
551
|
+
label_font_size=label_font_size,
|
|
552
|
+
legend_axs=legend_axs,
|
|
553
|
+
legend_marker_size=legend_marker_size,
|
|
554
|
+
legend_marker_edge_width=legend_marker_edge_width,
|
|
555
|
+
legend_font_size=legend_font_size,
|
|
556
|
+
legend_n_cols=legend_n_cols,
|
|
557
|
+
title=title,
|
|
558
|
+
title_font_size=title_font_size,
|
|
559
|
+
arrow_thickness_factor=arrow_thickness_factor,
|
|
560
|
+
arrow_legend_axs=arrow_legend_axs,
|
|
561
|
+
arrow_legend_font_size=arrow_legend_font_size,
|
|
562
|
+
arrow_head_width=arrow_head_width,
|
|
563
|
+
arrow_head_length=arrow_head_length,
|
|
564
|
+
arrow_legend_n_cols=arrow_legend_n_cols,
|
|
565
|
+
arrow_normalization_factor=arrow_normalization_factor,
|
|
566
|
+
show_counts=show_counts
|
|
567
|
+
)
|
|
568
|
+
|
|
569
|
+
|
|
570
|
+
def plot_aula_medica_plot(
|
|
571
|
+
axs: plt.Axes,
|
|
572
|
+
mnc_object: MultiNeuronChatObject,
|
|
573
|
+
wasserstein_distance_matrix: xr.DataArray,
|
|
574
|
+
statistical_test: str = 'KS',
|
|
575
|
+
use_adjusted_p_values: bool = True,
|
|
576
|
+
cell_types: list[str] | list[tuple[str, list[str]]] | None = None,
|
|
577
|
+
cell_type_labels_short: dict[str, str] | None = None,
|
|
578
|
+
ligand_target_interactions: list[str] | None = None,
|
|
579
|
+
scale: float = 1.0,
|
|
580
|
+
triangle_line_width: float = 0.5,
|
|
581
|
+
source_cell_type_split_line_width: float = 1.0,
|
|
582
|
+
marker_size: float = 10,
|
|
583
|
+
tick_font_size: float = 20,
|
|
584
|
+
wasserstein_color_map: str = 'seagreen',
|
|
585
|
+
p_value_color_map: str = 'salmon',
|
|
586
|
+
not_tested_color: str = '#CCCCCC',
|
|
587
|
+
wasserstein_legend_axs: plt.Axes | None = None,
|
|
588
|
+
p_value_legend_axs: plt.Axes | None = None,
|
|
589
|
+
wasserstein_rounding: int | None = None,
|
|
590
|
+
p_value_rounding: int | None = None,
|
|
591
|
+
annotate_source_cell_types: bool = True,
|
|
592
|
+
annotate_receiver_cell_types: bool = True,
|
|
593
|
+
annotate_interactions: bool = True,
|
|
594
|
+
annotate_significance: bool = True,
|
|
595
|
+
significance_threshold: float = 0.05,
|
|
596
|
+
max_log_p_value: float | None = None,
|
|
597
|
+
max_wasserstein_distance: float | None = None,
|
|
598
|
+
):
|
|
599
|
+
"""
|
|
600
|
+
Create an Aula Medica plot showing the Wasserstein distances and p-values for differential communication analysis.
|
|
601
|
+
|
|
602
|
+
:param axs: Matplotlib Axes object where the Aula Medica plot will be drawn.
|
|
603
|
+
:param mnc_object: MultiNeuronChatObject containing the p-values for differential communication analysis.
|
|
604
|
+
:param wasserstein_distance_matrix: xarray DataArray containing the Wasserstein distances between cell-types.
|
|
605
|
+
:param statistical_test: Name of the statistical test to use for plotting (e.g 'KS', 'Anderson', etc.).
|
|
606
|
+
:param use_adjusted_p_values: Whether to use adjusted p-values (default: True).
|
|
607
|
+
:param cell_types: List of cell-types or list of tuples (source_cell_type, list of receiver_cell_types) to be included in the plot. If a list is provided all interactions between these cells will be plotted. If a list of tuples is given, the first entry of the tuple represents the sender-cell-type and the second entry the list of receiver-cell-types. If None, all cell-types will be used, and the second entry is a list of all the receiving cell-types to be plotted.
|
|
608
|
+
:param cell_type_labels_short: Dictionary mapping cell-types to their short labels for display.
|
|
609
|
+
:param ligand_target_interactions: List of ligand-target interactions to be included in the plot.
|
|
610
|
+
:param scale: Scale factor for the plot.
|
|
611
|
+
:param triangle_line_width: Line width for the triangles in the plot.
|
|
612
|
+
:param source_cell_type_split_line_width: Line width for the split lines between source cell-types.
|
|
613
|
+
:param marker_size: Size of the markers for annotations.
|
|
614
|
+
:param tick_font_size: Font size for the tick labels.
|
|
615
|
+
:param wasserstein_color_map: Color map for the Wasserstein distances.
|
|
616
|
+
:param p_value_color_map: Color map for the p-values.
|
|
617
|
+
:param not_tested_color: Color for interactions that were not tested.
|
|
618
|
+
:param wasserstein_legend_axs: Matplotlib Axes object where the Wasserstein legend will be drawn. If None, no legend is drawn.
|
|
619
|
+
:param p_value_legend_axs: Matplotlib Axes object where the p-value legend will be drawn. If None, no legend is drawn.
|
|
620
|
+
:param wasserstein_rounding: Number of decimal places to round the Wasserstein distances in the legend. If None, no rounding is applied.
|
|
621
|
+
:param p_value_rounding: Number of decimal places to round the p-values in the legend. If None, no rounding is applied.
|
|
622
|
+
:param annotate_source_cell_types: Whether to annotate the source cell-types.
|
|
623
|
+
:param annotate_receiver_cell_types: Whether to annotate the receiver cell-types.
|
|
624
|
+
:param annotate_interactions: Whether to annotate the ligand-target interactions.
|
|
625
|
+
:param annotate_significance: Whether to annotate the significance based on the significance threshold.
|
|
626
|
+
:param significance_threshold: Significance threshold for p-values to consider an interaction significant (default: 0.05).
|
|
627
|
+
:param max_log_p_value: Maximum -log10(p-value) to use for normalization. This is especially useful if you plot several lines of the Aula Medica plot to have the same color scale. If None, the maximum value in the data is used.
|
|
628
|
+
:param max_wasserstein_distance: Maximum Wasserstein distance to use for normalization. If None, the maximum value in the data is used.
|
|
629
|
+
|
|
630
|
+
:return: None
|
|
631
|
+
"""
|
|
632
|
+
# If no cell-types are provided -> use all cell-types
|
|
633
|
+
all_cell_types: list[str] = list(set(mnc_object.source_cell_types + mnc_object.receiver_cell_types))
|
|
634
|
+
all_cell_types.sort()
|
|
635
|
+
if cell_types is None:
|
|
636
|
+
cell_types: list[tuple[str, list[str]]] = [
|
|
637
|
+
(cell_type, all_cell_types)
|
|
638
|
+
for cell_type in all_cell_types
|
|
639
|
+
]
|
|
640
|
+
else: # If cell-types are provided -> check if the selection is valid
|
|
641
|
+
if type(cell_types[0]) is str:
|
|
642
|
+
# The provided type is a list of strings -> first check if all cell-types are valid, then convert to a list of tuples
|
|
643
|
+
if not all([cell_type in all_cell_types for cell_type in cell_types]):
|
|
644
|
+
raise ValueError('Invalid cell types')
|
|
645
|
+
|
|
646
|
+
cell_types: list[tuple[str, list[str]]] = [
|
|
647
|
+
(cell_type, cell_types)
|
|
648
|
+
for cell_type in cell_types
|
|
649
|
+
]
|
|
650
|
+
elif type(cell_types[0]) is tuple:
|
|
651
|
+
# The provided type is a list of tuples -> check if all cell-types are valid
|
|
652
|
+
selected_cell_types: list[str] = list(set(
|
|
653
|
+
[cell_type for cell_type, _ in cell_types] +
|
|
654
|
+
[cell_type for _, receiver_cell_types in cell_types for cell_type in receiver_cell_types]
|
|
655
|
+
))
|
|
656
|
+
|
|
657
|
+
if not all([cell_type in all_cell_types for cell_type in selected_cell_types]):
|
|
658
|
+
raise ValueError('Invalid cell types')
|
|
659
|
+
|
|
660
|
+
# If no ligand-target interactions are provided -> use all ligand-target interactions
|
|
661
|
+
if ligand_target_interactions is None:
|
|
662
|
+
ligand_target_interactions = mnc_object.interaction_names
|
|
663
|
+
else: # If ligand-target interactions are provided -> check if the selection is valid
|
|
664
|
+
if not all([interaction in mnc_object.interaction_names for interaction in ligand_target_interactions]):
|
|
665
|
+
raise ValueError('Invalid ligand-target interactions')
|
|
666
|
+
|
|
667
|
+
if statistical_test not in ['KS', 'Anderson', 'CVM', 'Wilcoxon', 'MannWhitneyU']:
|
|
668
|
+
raise ValueError('Invalid statistical test')
|
|
669
|
+
elif (not use_adjusted_p_values) and (statistical_test not in mnc_object.p_values.keys()):
|
|
670
|
+
raise ValueError('The p-values have not been yet calculated for the selected statistical test')
|
|
671
|
+
elif use_adjusted_p_values and (statistical_test not in mnc_object.p_values_adj.keys()):
|
|
672
|
+
raise ValueError('The adjusted p-values have not been yet calculated for the selected statistical test')
|
|
673
|
+
|
|
674
|
+
n_x: int = sum([len(cell_type[1]) for cell_type in cell_types])
|
|
675
|
+
n_y: int = len(ligand_target_interactions)
|
|
676
|
+
|
|
677
|
+
p_values: xr.DataArray = mnc_object.p_values_adj[statistical_test] if use_adjusted_p_values else \
|
|
678
|
+
mnc_object.p_values[statistical_test]
|
|
679
|
+
neg_log_p_values: xr.DataArray = -np.log10(p_values)
|
|
680
|
+
|
|
681
|
+
if max_log_p_value is None:
|
|
682
|
+
max_log_p_value: float = np.nanmax(neg_log_p_values.to_numpy())
|
|
683
|
+
|
|
684
|
+
if max_wasserstein_distance is None:
|
|
685
|
+
max_wasserstein_distance: float = np.nanmax(wasserstein_distance_matrix.to_numpy())
|
|
686
|
+
|
|
687
|
+
# Define the points of the triangle to draw
|
|
688
|
+
uniform_lower_triangle_corners_x: np.array = np.array([0, 1, 0, 0]) * scale
|
|
689
|
+
uniform_lower_triangle_corners_y: np.array = np.array([0, 0, 1, 0]) * scale
|
|
690
|
+
uniform_upper_triangle_corners_x: np.array = np.array([1, 1, 0, 1]) * scale
|
|
691
|
+
uniform_upper_triangle_corners_y: np.array = np.array([0, 1, 1, 0]) * scale
|
|
692
|
+
|
|
693
|
+
triangle_positions_x: np.array = np.arange(n_x) * scale
|
|
694
|
+
triangle_positions_y: np.array = np.arange(n_y) * scale
|
|
695
|
+
|
|
696
|
+
triangle_coordinates: np.array = np.stack(np.meshgrid(triangle_positions_x, triangle_positions_y)).transpose(2, 1,
|
|
697
|
+
0)
|
|
698
|
+
|
|
699
|
+
wasserstein_cmap: sns.color_palette = sns.light_palette(wasserstein_color_map, as_cmap=True)
|
|
700
|
+
p_value_cmap: sns.color_palette = sns.light_palette(p_value_color_map, as_cmap=True)
|
|
701
|
+
|
|
702
|
+
# Iterate through
|
|
703
|
+
x_pos: int = 0
|
|
704
|
+
for source_idx, (source_cell_type, receiver_cell_types) in enumerate(cell_types):
|
|
705
|
+
for receiver_idx, receiver_cell_type in enumerate(receiver_cell_types):
|
|
706
|
+
for interaction_idx, interaction in enumerate(ligand_target_interactions):
|
|
707
|
+
y_pos: int = interaction_idx
|
|
708
|
+
|
|
709
|
+
p_value: float = p_values.loc[
|
|
710
|
+
{'source': source_cell_type, 'receiver': receiver_cell_type,
|
|
711
|
+
'interaction': interaction}].values.item()
|
|
712
|
+
neg_log_p_value: float = neg_log_p_values.loc[
|
|
713
|
+
{'source': source_cell_type, 'receiver': receiver_cell_type,
|
|
714
|
+
'interaction': interaction}].values.item()
|
|
715
|
+
wasserstein_distance: float = wasserstein_distance_matrix.loc[
|
|
716
|
+
{'source': source_cell_type, 'receiver': receiver_cell_type,
|
|
717
|
+
'interaction': interaction}].values.item()
|
|
718
|
+
|
|
719
|
+
normalized_neg_log_p_value: float = neg_log_p_value / max_log_p_value
|
|
720
|
+
normalized_wasserstein_distance: float = wasserstein_distance / max_wasserstein_distance
|
|
721
|
+
|
|
722
|
+
x_start, y_start = triangle_coordinates[x_pos, y_pos]
|
|
723
|
+
|
|
724
|
+
axs.plot(
|
|
725
|
+
uniform_lower_triangle_corners_x + x_start,
|
|
726
|
+
uniform_lower_triangle_corners_y + y_start,
|
|
727
|
+
color='black',
|
|
728
|
+
linewidth=triangle_line_width
|
|
729
|
+
)
|
|
730
|
+
axs.plot(
|
|
731
|
+
uniform_upper_triangle_corners_x + x_start,
|
|
732
|
+
uniform_upper_triangle_corners_y + y_start,
|
|
733
|
+
color='black',
|
|
734
|
+
linewidth=triangle_line_width
|
|
735
|
+
)
|
|
736
|
+
|
|
737
|
+
if math.isnan(normalized_neg_log_p_value):
|
|
738
|
+
p_value_color: str = not_tested_color
|
|
739
|
+
else:
|
|
740
|
+
p_value_color: tuple = p_value_cmap(normalized_neg_log_p_value)
|
|
741
|
+
|
|
742
|
+
if math.isnan(wasserstein_distance):
|
|
743
|
+
wasserstein_color: str = not_tested_color
|
|
744
|
+
else:
|
|
745
|
+
wasserstein_color: tuple = wasserstein_cmap(normalized_wasserstein_distance)
|
|
746
|
+
|
|
747
|
+
axs.fill(
|
|
748
|
+
uniform_lower_triangle_corners_x + x_start,
|
|
749
|
+
uniform_lower_triangle_corners_y + y_start,
|
|
750
|
+
color=wasserstein_color,
|
|
751
|
+
)
|
|
752
|
+
axs.fill(
|
|
753
|
+
uniform_upper_triangle_corners_x + x_start,
|
|
754
|
+
uniform_upper_triangle_corners_y + y_start,
|
|
755
|
+
color=p_value_color,
|
|
756
|
+
)
|
|
757
|
+
|
|
758
|
+
if annotate_significance and p_value < significance_threshold:
|
|
759
|
+
mid_point_x = mid_point_y = 0.75 * scale
|
|
760
|
+
axs.scatter(
|
|
761
|
+
x=x_start + mid_point_x,
|
|
762
|
+
y=y_start + mid_point_y,
|
|
763
|
+
marker='.',
|
|
764
|
+
color='black',
|
|
765
|
+
s=marker_size,
|
|
766
|
+
)
|
|
767
|
+
|
|
768
|
+
x_pos += 1
|
|
769
|
+
|
|
770
|
+
# Annotate x-axis
|
|
771
|
+
axs.set_xticks(triangle_positions_x + (scale / 2))
|
|
772
|
+
if annotate_receiver_cell_types:
|
|
773
|
+
receiver_cell_type_list: list[str] = [cell_type for _, receiver_cell_types in cell_types for cell_type in
|
|
774
|
+
receiver_cell_types]
|
|
775
|
+
if cell_type_labels_short is not None:
|
|
776
|
+
receiver_cell_type_list = [cell_type_labels_short[cell_type] for cell_type in receiver_cell_type_list]
|
|
777
|
+
|
|
778
|
+
axs.set_xticklabels(
|
|
779
|
+
receiver_cell_type_list,
|
|
780
|
+
fontsize=tick_font_size,
|
|
781
|
+
rotation=90
|
|
782
|
+
)
|
|
783
|
+
else:
|
|
784
|
+
axs.set_xticklabels([])
|
|
785
|
+
|
|
786
|
+
if annotate_source_cell_types:
|
|
787
|
+
current_count_cell_types: int = 0
|
|
788
|
+
for i, (source_cell_type, receiver_cell_types) in enumerate(cell_types):
|
|
789
|
+
n_receiver_cell_types: int = len(receiver_cell_types)
|
|
790
|
+
label_x_position = (current_count_cell_types + (n_receiver_cell_types / 2)) * scale
|
|
791
|
+
|
|
792
|
+
current_count_cell_types += n_receiver_cell_types
|
|
793
|
+
|
|
794
|
+
if cell_type_labels_short is not None:
|
|
795
|
+
source_cell_type = cell_type_labels_short[source_cell_type]
|
|
796
|
+
|
|
797
|
+
# Plot text at the top of the plot
|
|
798
|
+
axs.text(
|
|
799
|
+
label_x_position,
|
|
800
|
+
n_y * scale + 0.5,
|
|
801
|
+
source_cell_type,
|
|
802
|
+
fontsize=tick_font_size,
|
|
803
|
+
ha='center'
|
|
804
|
+
)
|
|
805
|
+
|
|
806
|
+
# Plot lines that separate the cell-types
|
|
807
|
+
current_count_cell_types: int = 0
|
|
808
|
+
for i, (source_cell_type, receiver_cell_types) in enumerate(cell_types[:-1]):
|
|
809
|
+
current_count_cell_types += len(receiver_cell_types)
|
|
810
|
+
axs.plot(
|
|
811
|
+
[current_count_cell_types * scale, current_count_cell_types * scale],
|
|
812
|
+
[0, n_y * scale],
|
|
813
|
+
color='black',
|
|
814
|
+
linewidth=source_cell_type_split_line_width
|
|
815
|
+
)
|
|
816
|
+
|
|
817
|
+
# Annotate y-axis
|
|
818
|
+
axs.set_yticks(triangle_positions_y + (scale / 2))
|
|
819
|
+
if annotate_interactions:
|
|
820
|
+
axs.set_yticklabels(
|
|
821
|
+
ligand_target_interactions,
|
|
822
|
+
fontsize=tick_font_size
|
|
823
|
+
)
|
|
824
|
+
else:
|
|
825
|
+
axs.set_yticklabels([])
|
|
826
|
+
|
|
827
|
+
axs.grid(False)
|
|
828
|
+
|
|
829
|
+
axs.set_xlim(0, n_x * scale)
|
|
830
|
+
axs.set_ylim(0, n_y * scale)
|
|
831
|
+
|
|
832
|
+
# Remove borders
|
|
833
|
+
axs.spines['top'].set_visible(False)
|
|
834
|
+
axs.spines['right'].set_visible(False)
|
|
835
|
+
axs.spines['bottom'].set_visible(False)
|
|
836
|
+
axs.spines['left'].set_visible(False)
|
|
837
|
+
|
|
838
|
+
axs.set_aspect('equal')
|
|
839
|
+
|
|
840
|
+
if wasserstein_legend_axs is not None:
|
|
841
|
+
# Define the normalization from 0 to 1
|
|
842
|
+
norm = mpl.colors.Normalize(vmin=0, vmax=max_wasserstein_distance)
|
|
843
|
+
wasserstein_cb = mpl.colorbar.ColorbarBase(
|
|
844
|
+
wasserstein_legend_axs, # The axis to draw the colorbar on
|
|
845
|
+
cmap=wasserstein_cmap, # Your custom colormap
|
|
846
|
+
norm=norm, # Normalization from 0 to 1
|
|
847
|
+
orientation='vertical' # Orientation of the colorbar ('vertical' or 'horizontal')
|
|
848
|
+
)
|
|
849
|
+
|
|
850
|
+
wasserstein_tick_labels = [0, max_wasserstein_distance / 2, max_wasserstein_distance]
|
|
851
|
+
if wasserstein_rounding is not None:
|
|
852
|
+
wasserstein_tick_labels = [round(label, wasserstein_rounding) for label in wasserstein_tick_labels]
|
|
853
|
+
|
|
854
|
+
# Optionally, set the label and ticks
|
|
855
|
+
wasserstein_cb.set_label('Wasserstein distance')
|
|
856
|
+
wasserstein_cb.set_ticks(wasserstein_tick_labels) # Set custom tick positions
|
|
857
|
+
wasserstein_cb.set_ticklabels(wasserstein_tick_labels) # Set custom tick labels
|
|
858
|
+
|
|
859
|
+
if p_value_legend_axs is not None:
|
|
860
|
+
# Define the normalization from 0 to 1
|
|
861
|
+
norm = mpl.colors.Normalize(vmin=0, vmax=max_log_p_value)
|
|
862
|
+
|
|
863
|
+
# Create the colorbar on your custom axis
|
|
864
|
+
p_val_cb = mpl.colorbar.ColorbarBase(
|
|
865
|
+
p_value_legend_axs, # The axis to draw the colorbar on
|
|
866
|
+
cmap=p_value_cmap, # Your custom colormap
|
|
867
|
+
norm=norm, # Normalization from 0 to 1
|
|
868
|
+
orientation='vertical' # Orientation of the colorbar ('vertical' or 'horizontal')
|
|
869
|
+
)
|
|
870
|
+
|
|
871
|
+
p_value_tick_labels = [0, max_log_p_value / 2, max_log_p_value]
|
|
872
|
+
if p_value_rounding is not None:
|
|
873
|
+
p_value_tick_labels = [round(label, p_value_rounding) for label in p_value_tick_labels]
|
|
874
|
+
|
|
875
|
+
# Optionally, set the label and ticks
|
|
876
|
+
if use_adjusted_p_values:
|
|
877
|
+
p_val_cb.set_label('-log10(adjusted p-value)')
|
|
878
|
+
else:
|
|
879
|
+
p_val_cb.set_label('-log10(p-value)')
|
|
880
|
+
|
|
881
|
+
p_val_cb.set_ticks(p_value_tick_labels) # Set custom tick positions
|
|
882
|
+
p_val_cb.set_ticklabels(p_value_tick_labels)
|
|
883
|
+
|
|
884
|
+
|
|
885
|
+
def plot_communication_score_distribution(
|
|
886
|
+
axs: plt.Axes,
|
|
887
|
+
mnc_object: MultiNeuronChatObject,
|
|
888
|
+
cell_type_pair: tuple[str, str],
|
|
889
|
+
ligand_target_interaction: str,
|
|
890
|
+
n_bins: int = 25,
|
|
891
|
+
plot_p_value: bool = False,
|
|
892
|
+
use_adjusted_p_value: bool = True,
|
|
893
|
+
statistical_test: str | None = None,
|
|
894
|
+
show_legend: bool = False,
|
|
895
|
+
legend_position: str = 'best',
|
|
896
|
+
legend_axs: plt.Axes | None = None,
|
|
897
|
+
condition_colors: tuple[str, str] | None = None,
|
|
898
|
+
alpha: float = 0.5,
|
|
899
|
+
title_font_size: float = 20,
|
|
900
|
+
annotation_font_size: float = 20,
|
|
901
|
+
tick_font_size: float = 20,
|
|
902
|
+
p_value_font_size: float = 15,
|
|
903
|
+
legend_font_size: float = 20,
|
|
904
|
+
min_x: float = 0,
|
|
905
|
+
max_x: float = -1,
|
|
906
|
+
title: str = '',
|
|
907
|
+
):
|
|
908
|
+
"""
|
|
909
|
+
Plot the distribution of communication scores between two cell types for a specific ligand-target interaction.
|
|
910
|
+
|
|
911
|
+
:param axs: Matplotlib Axes object where the histogram will be drawn.
|
|
912
|
+
:param mnc_object: MultiNeuronChatObject containing the communication scores.
|
|
913
|
+
:param cell_type_pair: Tuple containing the source and receiver cell types.
|
|
914
|
+
:param ligand_target_interaction: Name of the ligand-target interaction to plot.
|
|
915
|
+
:param n_bins: Number of bins for the histogram (default: 25).
|
|
916
|
+
:param plot_p_value: Whether to plot the p-value on the histogram (default: False).
|
|
917
|
+
:param use_adjusted_p_value: Whether to use adjusted p-values (default: True).
|
|
918
|
+
:param statistical_test: Name of the statistical test to use for p-value retrieval (e.g., 'KS', 'Anderson', etc.). Required if plot_p_value is True.
|
|
919
|
+
:param show_legend: Whether to show the legend (default: False).
|
|
920
|
+
:param legend_position: Position of the legend (default: 'best').
|
|
921
|
+
:param legend_axs: Matplotlib Axes object where the legend will be drawn. If None, the legend will be drawn on the main axes.
|
|
922
|
+
:param condition_colors: Tuple containing the colors for the two conditions. If None, default colors will be used.
|
|
923
|
+
:param alpha: Transparency level for the histogram bars (default: 0.5).
|
|
924
|
+
:param title_font_size: Font size for the plot title (default: 20).
|
|
925
|
+
:param annotation_font_size: Font size for the axis labels (default: 20).
|
|
926
|
+
:param tick_font_size: Font size for the tick labels (default: 20).
|
|
927
|
+
:param p_value_font_size: Font size for the p-value text (default: 15).
|
|
928
|
+
:param legend_font_size: Font size for the legend text (default: 20).
|
|
929
|
+
:param min_x: Minimum x-axis value for the histogram (default: 0).
|
|
930
|
+
:param max_x: Maximum x-axis value for the histogram. If -1, it will be set to the maximum communication score (default: -1).
|
|
931
|
+
:param title: Title of the plot (default: '').
|
|
932
|
+
|
|
933
|
+
:return: None
|
|
934
|
+
"""
|
|
935
|
+
if plot_p_value:
|
|
936
|
+
if statistical_test is None:
|
|
937
|
+
raise ValueError("You selected to plot the p-value of this comparision, but have not selected a statistical test! Please select a valid statistical test!")
|
|
938
|
+
|
|
939
|
+
if use_adjusted_p_value and (statistical_test not in mnc_object.p_values_adj.keys()):
|
|
940
|
+
raise ValueError(f'The statistical test "{statistical_test}" is not available in the current MultiNeuronChat adjusted p-values! Please select one of the valid tests: {mnc_object.p_values_adj.keys()}')
|
|
941
|
+
elif (not use_adjusted_p_value) and (statistical_test not in mnc_object.p_values.keys()):
|
|
942
|
+
raise ValueError(f'The statistical test "{statistical_test}" is not available in the current MultiNeuronChat p-values! Please select one of the valid tests: {mnc_object.p_values.keys()}')
|
|
943
|
+
|
|
944
|
+
# Check if the selected cell-type pair and ligand-target pair exists in the MultiNeuronChat object
|
|
945
|
+
mnc_cell_types: set[str] = set(mnc_object.source_cell_types + mnc_object.receiver_cell_types)
|
|
946
|
+
|
|
947
|
+
if (not cell_type_pair[0] in mnc_cell_types) or (not cell_type_pair[1] in mnc_cell_types):
|
|
948
|
+
raise ValueError(f'One of the two selected cell-types {cell_type_pair} is not valid! Available cell-types: {mnc_cell_types}')
|
|
949
|
+
|
|
950
|
+
if ligand_target_interaction not in mnc_object.interaction_names:
|
|
951
|
+
raise ValueError(f'The selected ligand-target interaction {ligand_target_interaction} is not valid! Available interactions: {mnc_object.interaction_names}')
|
|
952
|
+
|
|
953
|
+
condition_a_communication_scores: np.array = mnc_object.communication_score_dict[mnc_object.condition_names[0]].sel(
|
|
954
|
+
source=cell_type_pair[0],
|
|
955
|
+
receiver=cell_type_pair[1],
|
|
956
|
+
interaction=ligand_target_interaction,
|
|
957
|
+
).to_numpy()
|
|
958
|
+
condition_b_communication_scores: np.array = mnc_object.communication_score_dict[mnc_object.condition_names[1]].sel(
|
|
959
|
+
source=cell_type_pair[0],
|
|
960
|
+
receiver=cell_type_pair[1],
|
|
961
|
+
interaction=ligand_target_interaction,
|
|
962
|
+
).to_numpy()
|
|
963
|
+
|
|
964
|
+
condition_a_communication_scores = condition_a_communication_scores[~np.isnan(condition_a_communication_scores)]
|
|
965
|
+
condition_b_communication_scores = condition_b_communication_scores[~np.isnan(condition_b_communication_scores)]
|
|
966
|
+
|
|
967
|
+
# Uniform binning
|
|
968
|
+
if max_x < 0:
|
|
969
|
+
max_x = max(np.max(condition_a_communication_scores), np.max(condition_b_communication_scores))
|
|
970
|
+
|
|
971
|
+
bins = np.linspace(min_x, max_x, n_bins + 1)
|
|
972
|
+
|
|
973
|
+
axs.hist(
|
|
974
|
+
x=condition_a_communication_scores,
|
|
975
|
+
bins=bins,
|
|
976
|
+
color=condition_colors[0],
|
|
977
|
+
alpha=alpha,
|
|
978
|
+
label=mnc_object.condition_names[0]
|
|
979
|
+
)
|
|
980
|
+
axs.hist(
|
|
981
|
+
x=condition_b_communication_scores,
|
|
982
|
+
bins=bins,
|
|
983
|
+
color=condition_colors[1],
|
|
984
|
+
alpha=alpha,
|
|
985
|
+
label=mnc_object.condition_names[1]
|
|
986
|
+
)
|
|
987
|
+
|
|
988
|
+
top_y_value: float = axs.get_ylim()[1]
|
|
989
|
+
y_ticks: np.array = np.arange(0, top_y_value)
|
|
990
|
+
axs.set_yticks(y_ticks)
|
|
991
|
+
axs.set_yticklabels(y_ticks)
|
|
992
|
+
|
|
993
|
+
axs.tick_params(axis='both', which='major', labelsize=tick_font_size)
|
|
994
|
+
|
|
995
|
+
if plot_p_value:
|
|
996
|
+
if use_adjusted_p_value:
|
|
997
|
+
p_value: float = mnc_object.p_values_adj[statistical_test].sel(
|
|
998
|
+
source=cell_type_pair[0],
|
|
999
|
+
receiver=cell_type_pair[1],
|
|
1000
|
+
interaction=ligand_target_interaction,
|
|
1001
|
+
).values.tolist()
|
|
1002
|
+
|
|
1003
|
+
else:
|
|
1004
|
+
p_value: float = mnc_object.p_values[statistical_test].sel(
|
|
1005
|
+
source=cell_type_pair[0],
|
|
1006
|
+
receiver=cell_type_pair[1],
|
|
1007
|
+
interaction=ligand_target_interaction,
|
|
1008
|
+
).values.tolist()
|
|
1009
|
+
axs.text(
|
|
1010
|
+
x=0.8,
|
|
1011
|
+
y=0.8,
|
|
1012
|
+
s=f'{"adj. " if use_adjusted_p_value else ""}p-value:\n{Decimal(p_value):.4E}',
|
|
1013
|
+
ha='center',
|
|
1014
|
+
fontsize=p_value_font_size,
|
|
1015
|
+
transform=axs.transAxes
|
|
1016
|
+
)
|
|
1017
|
+
|
|
1018
|
+
|
|
1019
|
+
axs.set_ylabel('#Donors', fontsize=annotation_font_size)
|
|
1020
|
+
axs.set_xlabel('Communication score', fontsize=annotation_font_size)
|
|
1021
|
+
axs.set_title(title, fontsize=title_font_size)
|
|
1022
|
+
|
|
1023
|
+
if show_legend:
|
|
1024
|
+
if legend_axs is not None:
|
|
1025
|
+
legend_axs.plot([], [], marker='s', color=condition_colors[0], label=mnc_object.condition_names[0])
|
|
1026
|
+
legend_axs.plot([], [], marker='s', color=condition_colors[0], label=mnc_object.condition_names[1])
|
|
1027
|
+
legend_axs.legend(loc='center', fontsize=legend_font_size)
|
|
1028
|
+
else:
|
|
1029
|
+
axs.legend(loc=legend_position, fontsize=legend_font_size)
|