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,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)