pyCoReGraph 0.0.1a1__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.
- coregraph/__init__.py +57 -0
- coregraph/analysis/__init__.py +45 -0
- coregraph/analysis/bins.py +23 -0
- coregraph/analysis/correlations.py +110 -0
- coregraph/analysis/fdr.py +100 -0
- coregraph/analysis/queries.py +118 -0
- coregraph/analysis/subset.py +64 -0
- coregraph/graphics/__init__.py +33 -0
- coregraph/graphics/animations.py +56 -0
- coregraph/graphics/colors.py +9 -0
- coregraph/graphics/network.py +214 -0
- coregraph/graphics/plots.py +85 -0
- coregraph/models/__init__.py +11 -0
- coregraph/models/coregraph.py +221 -0
- coregraph/models/cors.py +63 -0
- coregraph/models/graphics.py +38 -0
- coregraph/models/stargraph.py +41 -0
- coregraph/utils/__init__.py +23 -0
- coregraph/utils/decorators.py +39 -0
- coregraph/utils/formatting.py +16 -0
- coregraph/utils/validation.py +25 -0
- pycoregraph-0.0.1a1.dist-info/METADATA +65 -0
- pycoregraph-0.0.1a1.dist-info/RECORD +26 -0
- pycoregraph-0.0.1a1.dist-info/WHEEL +5 -0
- pycoregraph-0.0.1a1.dist-info/licenses/LICENSE +21 -0
- pycoregraph-0.0.1a1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,214 @@
|
|
|
1
|
+
## NETWORK ##
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import networkx as nx
|
|
7
|
+
import matplotlib.pyplot as plt
|
|
8
|
+
import matplotlib.cm as cm
|
|
9
|
+
import matplotlib.colors as mcolors
|
|
10
|
+
|
|
11
|
+
from matplotlib.animation import FuncAnimation
|
|
12
|
+
|
|
13
|
+
from coregraph.analysis.queries import (get_all_correlations, get_mac)
|
|
14
|
+
from coregraph.utils.decorators import (requires_cors, requires_FDR)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _build_correlation_tensor(C, genes, correlation_threshold):
|
|
18
|
+
'''
|
|
19
|
+
Build tensor (n_genes, n_genes, n_steps) containing all dynamic correlations.
|
|
20
|
+
'''
|
|
21
|
+
n = len(genes)
|
|
22
|
+
cor_tensor = np.zeros((n, n, C.cors.n_steps), dtype=C.cors.dtype)
|
|
23
|
+
|
|
24
|
+
for i, regulator in enumerate(genes):
|
|
25
|
+
for j, target in enumerate(genes):
|
|
26
|
+
if i == j:
|
|
27
|
+
continue
|
|
28
|
+
|
|
29
|
+
ref_lag_cor = get_mac(C, regulator, target)
|
|
30
|
+
ref = ref_lag_cor["ref"]
|
|
31
|
+
|
|
32
|
+
all_cors = get_all_correlations(C, regulator, target, ref)
|
|
33
|
+
all_cors[np.abs(all_cors) < correlation_threshold] = 0
|
|
34
|
+
|
|
35
|
+
cor_tensor[i, j, :] = all_cors[:C.cors.n_steps]
|
|
36
|
+
|
|
37
|
+
return cor_tensor
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def _build_adjacency_tensor(cor_tensor):
|
|
41
|
+
'''
|
|
42
|
+
Keep strongest direction only between pairs.
|
|
43
|
+
'''
|
|
44
|
+
adjacency = np.abs(cor_tensor) > np.abs(cor_tensor.transpose(1, 0, 2))
|
|
45
|
+
filtered = cor_tensor.copy()
|
|
46
|
+
filtered[~adjacency] = 0
|
|
47
|
+
|
|
48
|
+
return filtered
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _compute_node_sizes(C, genes, node_ratio=1):
|
|
52
|
+
'''
|
|
53
|
+
Compute node size evolution through time.
|
|
54
|
+
'''
|
|
55
|
+
gene_indices = [C.gene_to_idx[g] for g in genes]
|
|
56
|
+
|
|
57
|
+
return (node_ratio * 1000 * C.cors.LAG_scaled[:C.cors.n_steps, gene_indices])
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def _build_graph_layout(genes):
|
|
61
|
+
'''
|
|
62
|
+
Create stable graph layout.
|
|
63
|
+
'''
|
|
64
|
+
G = nx.DiGraph()
|
|
65
|
+
G.add_nodes_from(range(len(genes)))
|
|
66
|
+
pos = nx.spring_layout(G, seed=33)
|
|
67
|
+
|
|
68
|
+
return G, pos
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def _compute_edge_style(G, cmap, norm, edge_ratio):
|
|
72
|
+
'''
|
|
73
|
+
Compute edge colors and widths.
|
|
74
|
+
'''
|
|
75
|
+
weights = nx.get_edge_attributes(G, 'weight')
|
|
76
|
+
edge_colors = [cmap(norm(w)) for w in weights.values()]
|
|
77
|
+
edge_widths = [abs(w) * edge_ratio for w in weights.values()]
|
|
78
|
+
|
|
79
|
+
return (weights, edge_colors, edge_widths)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def _draw_network_frame(ax, G, pos, genes, adjacency_matrix, node_sizes, cmap, norm, edge_ratio, font_size, min_target_margin, min_source_margin):
|
|
83
|
+
'''
|
|
84
|
+
Draw single network frame.
|
|
85
|
+
'''
|
|
86
|
+
ax.clear()
|
|
87
|
+
G.clear_edges()
|
|
88
|
+
n = len(genes)
|
|
89
|
+
|
|
90
|
+
# Add edges
|
|
91
|
+
for i in range(n):
|
|
92
|
+
for j in range(n):
|
|
93
|
+
weight = adjacency_matrix[i, j]
|
|
94
|
+
if weight != 0:
|
|
95
|
+
G.add_edge(i, j, weight=weight)
|
|
96
|
+
|
|
97
|
+
# Edge style
|
|
98
|
+
(weights, edge_colors, edge_widths) = _compute_edge_style(G, cmap, norm, edge_ratio)
|
|
99
|
+
|
|
100
|
+
# Draw nodes
|
|
101
|
+
nx.draw_networkx_nodes(G, pos, ax=ax, node_size=node_sizes, node_color='skyblue')
|
|
102
|
+
|
|
103
|
+
# Draw labels
|
|
104
|
+
nx.draw_networkx_labels(G, pos, labels={i: gene for i, gene in enumerate(genes)}, ax=ax, font_size=font_size)
|
|
105
|
+
|
|
106
|
+
# Draw edges
|
|
107
|
+
edges = nx.draw_networkx_edges(G, pos, ax=ax, edge_color=edge_colors, width=edge_widths, arrowstyle='-|>', min_target_margin=min_target_margin, min_source_margin=min_source_margin,)
|
|
108
|
+
|
|
109
|
+
for edge in edges:
|
|
110
|
+
edge.set_joinstyle("miter")
|
|
111
|
+
edge.set_capstyle("butt")
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
@requires_cors
|
|
115
|
+
@requires_FDR
|
|
116
|
+
def animate_network(
|
|
117
|
+
C: CoReGraph,
|
|
118
|
+
genes: list[str],
|
|
119
|
+
fdr_threshold: float=0.01,
|
|
120
|
+
correlation_threshold: float | None = None,
|
|
121
|
+
edge_ratio: float=1,
|
|
122
|
+
node_ratio: float=1,
|
|
123
|
+
min_target_margin: float=15,
|
|
124
|
+
min_source_margin: float=15,
|
|
125
|
+
font_size: int=8,
|
|
126
|
+
figsize: tuple=(6, 6)
|
|
127
|
+
):
|
|
128
|
+
'''
|
|
129
|
+
Animate dynamic gene regulatory network.
|
|
130
|
+
|
|
131
|
+
Parameters
|
|
132
|
+
----------
|
|
133
|
+
C : CoreGraph
|
|
134
|
+
CoreGraph object.
|
|
135
|
+
|
|
136
|
+
genes : list[str]
|
|
137
|
+
Genes to include in network.
|
|
138
|
+
|
|
139
|
+
fdr_threshold : float
|
|
140
|
+
FDR threshold used to infer minimum correlation threshold.
|
|
141
|
+
|
|
142
|
+
correlation_threshold : float | None
|
|
143
|
+
Correlation threshold.
|
|
144
|
+
|
|
145
|
+
edge_ratio : float
|
|
146
|
+
Edge width scaling factor.
|
|
147
|
+
|
|
148
|
+
node_ratio : float
|
|
149
|
+
Node size scaling factor.
|
|
150
|
+
|
|
151
|
+
figsize : tuple
|
|
152
|
+
Figure size.
|
|
153
|
+
'''
|
|
154
|
+
# Validate threshold
|
|
155
|
+
if correlation_threshold is None:
|
|
156
|
+
correlation_threshold = np.min(C.cors.FDR.loc[C.cors.FDR["FDR"] < fdr_threshold, "cors"])
|
|
157
|
+
|
|
158
|
+
# Filter genes
|
|
159
|
+
genes = [g for g in genes if g in C.gene_to_idx]
|
|
160
|
+
|
|
161
|
+
# Build correlation tensor and graph
|
|
162
|
+
cor_tensor = _build_correlation_tensor(C, genes, correlation_threshold)
|
|
163
|
+
cor_tensor = _build_adjacency_tensor(cor_tensor)
|
|
164
|
+
node_sizes = _compute_node_sizes(C, genes, node_ratio=node_ratio)
|
|
165
|
+
G, pos = _build_graph_layout(genes)
|
|
166
|
+
|
|
167
|
+
# Figure
|
|
168
|
+
fig, ax = plt.subplots(figsize=figsize)
|
|
169
|
+
cmap = cm.bwr
|
|
170
|
+
vmin = cor_tensor.min()
|
|
171
|
+
vmax = cor_tensor.max()
|
|
172
|
+
|
|
173
|
+
if vmin == vmax:
|
|
174
|
+
vmax += 1e-10
|
|
175
|
+
|
|
176
|
+
norm = mcolors.TwoSlopeNorm(vmin=vmin, vcenter=0.0, vmax=vmax)
|
|
177
|
+
|
|
178
|
+
# Edge normalization
|
|
179
|
+
if C.cors.dtype == np.int64:
|
|
180
|
+
edge_ratio *= 0.01
|
|
181
|
+
else:
|
|
182
|
+
edge_ratio *= 10
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
# Update function
|
|
186
|
+
def update(frame):
|
|
187
|
+
adjacency_matrix = cor_tensor[:, :, frame]
|
|
188
|
+
_draw_network_frame(
|
|
189
|
+
ax=ax,
|
|
190
|
+
G=G,
|
|
191
|
+
pos=pos,
|
|
192
|
+
genes=genes,
|
|
193
|
+
adjacency_matrix=adjacency_matrix,
|
|
194
|
+
node_sizes=node_sizes[frame],
|
|
195
|
+
cmap=cmap,
|
|
196
|
+
norm=norm,
|
|
197
|
+
edge_ratio=edge_ratio,
|
|
198
|
+
font_size=font_size,
|
|
199
|
+
min_target_margin=min_target_margin,
|
|
200
|
+
min_source_margin=min_source_margin,
|
|
201
|
+
)
|
|
202
|
+
|
|
203
|
+
ax.set_title(f"Time step {frame * C.cors.step}")
|
|
204
|
+
|
|
205
|
+
return ax
|
|
206
|
+
|
|
207
|
+
# Create animation
|
|
208
|
+
if plt.rcParams["animation.embed_limit"] < 250:
|
|
209
|
+
plt.rcParams["animation.embed_limit"] = 250
|
|
210
|
+
|
|
211
|
+
ani = FuncAnimation(fig, update, frames=C.cors.n_steps, interval=1000, repeat=True, blit=False)
|
|
212
|
+
plt.close(fig)
|
|
213
|
+
# HTML(ani.to_jshtml())
|
|
214
|
+
return ani
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
## PLOTS ##
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import matplotlib.pyplot as plt
|
|
7
|
+
import matplotlib.patches as patches
|
|
8
|
+
|
|
9
|
+
from coregraph.graphics.colors import ColorBlind
|
|
10
|
+
from coregraph.analysis.queries import (get_correlation)
|
|
11
|
+
from coregraph.utils.decorators import (requires_cors, requires_bins)
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@requires_bins
|
|
15
|
+
def plot_expression(C: CoReGraph, genes: str | list[str], scaled: bool=False, std: bool=False, color=None, linewidth: float=1, figsize: tuple=(8,4)):
|
|
16
|
+
'''Plot gene expression along pseudotime.'''
|
|
17
|
+
|
|
18
|
+
### TODO : make std parameter functional ###
|
|
19
|
+
|
|
20
|
+
if C.graphics.n_bins is None:
|
|
21
|
+
raise ValueError("Bins have not been calculated.")
|
|
22
|
+
|
|
23
|
+
if isinstance(genes, str):
|
|
24
|
+
genes = [genes]
|
|
25
|
+
|
|
26
|
+
# Handle colors — one per gene if not specified
|
|
27
|
+
if color is None:
|
|
28
|
+
colors = plt.cm.tab10(np.linspace(0, 1, len(genes)))
|
|
29
|
+
elif isinstance(color, str):
|
|
30
|
+
colors = [color] * len(genes)
|
|
31
|
+
else:
|
|
32
|
+
colors = color
|
|
33
|
+
|
|
34
|
+
values = (C.graphics.scaled if scaled else C.graphics.mean)
|
|
35
|
+
fig, ax = plt.subplots(figsize=figsize)
|
|
36
|
+
x = np.unique(C.graphics.bin_id)
|
|
37
|
+
|
|
38
|
+
for gene, colo in zip(genes, colors):
|
|
39
|
+
if (gene not in C.gene_to_idx):
|
|
40
|
+
raise ValueError(f"Unknown gene: {gene}")
|
|
41
|
+
|
|
42
|
+
idx = C.gene_to_idx[gene]
|
|
43
|
+
ax.plot(x, values[idx, :], label=gene, color=colo, linewidth=linewidth)
|
|
44
|
+
|
|
45
|
+
ax.set_xlabel("Pseudotime")
|
|
46
|
+
ax.set_ylabel("Scaled Expression" if scaled else "Expression")
|
|
47
|
+
ax.legend()
|
|
48
|
+
|
|
49
|
+
fig.tight_layout()
|
|
50
|
+
return ax
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
@requires_bins
|
|
54
|
+
@requires_cors
|
|
55
|
+
def plot_correlations(C: CoReGraph, regulator: str, target: str, ref_lag: tuple[int, int], scaled: bool=False, decimals: int=15, color: list=(ColorBlind[0], ColorBlind[3]), linewidth: float=1, figsize: tuple=(8,4)) :
|
|
56
|
+
'''Plot lagged correlations between genes.'''
|
|
57
|
+
if C.graphics.n_bins is None:
|
|
58
|
+
raise ValueError("Bins have not been calculated.")
|
|
59
|
+
|
|
60
|
+
values = (C.graphics.scaled if scaled else C.graphics.mean)
|
|
61
|
+
reg_idx = C.gene_to_idx[regulator]
|
|
62
|
+
tar_idx = C.gene_to_idx[target]
|
|
63
|
+
y_max = np.max([values[reg_idx, :], values[tar_idx, :]]) * 1.1
|
|
64
|
+
|
|
65
|
+
fig, ax = plt.subplots(figsize=figsize)
|
|
66
|
+
x = np.unique(C.graphics.bin_id)
|
|
67
|
+
|
|
68
|
+
ax.plot(x, values[reg_idx, :], color=color[0], linewidth=linewidth)
|
|
69
|
+
ax.plot(x, values[tar_idx, :], color=color[1], linewidth=linewidth)
|
|
70
|
+
|
|
71
|
+
ref, lag = ref_lag
|
|
72
|
+
ax.add_patch(patches.Rectangle((ref, 0), C.cors.window, y_max, facecolor=color[0], alpha=0.1,))
|
|
73
|
+
ax.add_patch(patches.Rectangle((lag, 0), C.cors.window, y_max, facecolor=color[1], alpha=0.1,))
|
|
74
|
+
|
|
75
|
+
cor = get_correlation(C, regulator, target, ref_lag)
|
|
76
|
+
ax.annotate(f"cor = {np.round(cor, decimals)}", (C.n_cells * 0.01, y_max * 0.85,))
|
|
77
|
+
ax.set_xlim(0, C.n_cells)
|
|
78
|
+
ax.set_ylim(0, y_max)
|
|
79
|
+
ax.set_xlabel("Pseudotime")
|
|
80
|
+
ax.set_ylabel("Scaled Expression" if scaled else "Expression")
|
|
81
|
+
ax.set_title(f'{regulator} [{ref_lag[0]}] -> {target} [{ref_lag[1]}]')
|
|
82
|
+
|
|
83
|
+
ax.legend([regulator, target])
|
|
84
|
+
|
|
85
|
+
return ax
|
|
@@ -0,0 +1,221 @@
|
|
|
1
|
+
## COREGRAPH ##
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import Optional
|
|
6
|
+
|
|
7
|
+
import warnings
|
|
8
|
+
import numpy as np
|
|
9
|
+
import pandas as pd
|
|
10
|
+
|
|
11
|
+
from tqdm import tqdm
|
|
12
|
+
|
|
13
|
+
from coregraph.models.cors import (CoReGraph_Cors)
|
|
14
|
+
from coregraph.models.graphics import (CoReGraph_Graphics)
|
|
15
|
+
from coregraph.analysis.correlations import (_calculate_all_cors)
|
|
16
|
+
from coregraph.analysis.bins import (_calculate_bins)
|
|
17
|
+
from coregraph.analysis.fdr import (_calculate_FDR)
|
|
18
|
+
|
|
19
|
+
from coregraph.utils.validation import (_solve_step, _solve_total_dim3)
|
|
20
|
+
from coregraph.utils.formatting import (_set_numerical_format)
|
|
21
|
+
from coregraph.utils.decorators import (requires_cors)
|
|
22
|
+
|
|
23
|
+
class CoReGraph:
|
|
24
|
+
'''
|
|
25
|
+
Main CoreGraph object.
|
|
26
|
+
|
|
27
|
+
Attributes
|
|
28
|
+
----------
|
|
29
|
+
data : ndarray
|
|
30
|
+
Original matrix with genes in rows and pseudotime-sorted cells in columns
|
|
31
|
+
n_genes : int
|
|
32
|
+
Total number of genes
|
|
33
|
+
n_cells : int
|
|
34
|
+
Total number of cells
|
|
35
|
+
pseudotime : ndarray
|
|
36
|
+
Pseudotime associated to each cells
|
|
37
|
+
gene_id : list
|
|
38
|
+
Vector of gene names
|
|
39
|
+
cell_id : list
|
|
40
|
+
Vector of cell names
|
|
41
|
+
gene_to_idx : dict
|
|
42
|
+
Dictionnary of {gene: gene_id}
|
|
43
|
+
cell_to_idx : dict
|
|
44
|
+
Dictionnary of {cell: cell_id}
|
|
45
|
+
cors : CoReGraph_Cors
|
|
46
|
+
CoReGraph_Cors object containing correlations once calculated
|
|
47
|
+
graphics : CoReGraph_Graphics
|
|
48
|
+
CoReGraph_Graphics object containing binned epxressions once calculated to draw plots
|
|
49
|
+
'''
|
|
50
|
+
|
|
51
|
+
## INITIALIZATION AND CLASS METHODS ##
|
|
52
|
+
|
|
53
|
+
def __init__(self, matrix: np.ndarray, pseudotime: Optional[np.ndarray]=None, genes=None, cells=None):
|
|
54
|
+
self.data = np.asarray(matrix)
|
|
55
|
+
self.n_genes = self.data.shape[0]
|
|
56
|
+
self.n_cells = self.data.shape[1]
|
|
57
|
+
|
|
58
|
+
self.pseudotime = (np.asarray(pseudotime) if pseudotime is not None else np.arange(self.n_cells))
|
|
59
|
+
self.gene_id = (list(genes) if genes is not None else list(range(self.n_genes)))
|
|
60
|
+
self.cell_id = (list(cells) if cells is not None else list(range(self.n_cells)))
|
|
61
|
+
self.gene_to_idx = {str(gene): int(idx) for idx, gene in enumerate(self.gene_id)}
|
|
62
|
+
self.cell_to_idx = {str(cell): int(idx) for idx, cell in enumerate(self.cell_id)}
|
|
63
|
+
|
|
64
|
+
self.cors = CoReGraph_Cors()
|
|
65
|
+
self.graphics = CoReGraph_Graphics()
|
|
66
|
+
|
|
67
|
+
@classmethod
|
|
68
|
+
def from_pandas(cls, dataframe: pd.DataFrame, with_pseudotime: bool=False):
|
|
69
|
+
if with_pseudotime:
|
|
70
|
+
return cls(matrix=dataframe.iloc[1:], pseudotime=dataframe.iloc[0], genes=dataframe.index[1:], cells=dataframe.columns)
|
|
71
|
+
|
|
72
|
+
return cls(matrix=dataframe, genes=dataframe.index, cells=dataframe.columns)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
## __DUNDERS__ ##
|
|
76
|
+
|
|
77
|
+
def __repr__(self):
|
|
78
|
+
msg = f'CoReGraph({self.n_genes} genes, {self.n_cells} cells)'
|
|
79
|
+
if self.cors.window is not None:
|
|
80
|
+
msg += f'\n└ cors(window={self.cors.window})'
|
|
81
|
+
if self.graphics.n_bins is not None:
|
|
82
|
+
msg += f'\n└ graphics(n_bins={self.graphics.n_bins})'
|
|
83
|
+
return msg
|
|
84
|
+
|
|
85
|
+
def __str__(self):
|
|
86
|
+
msg = f'CoReGraph object with {self.n_genes} genes and {self.n_cells} cells.'
|
|
87
|
+
if self.cors.window is not None:
|
|
88
|
+
msg += f'\n-- Correlations calculated with a window of {self.cors.window} cells.'
|
|
89
|
+
if self.graphics.n_bins is not None:
|
|
90
|
+
msg += f'\n-- Discretized with {self.graphics.n_bins} bins.'
|
|
91
|
+
return msg
|
|
92
|
+
|
|
93
|
+
def __getattr__(self, name):
|
|
94
|
+
for target in [self.cors, self.graphics]:
|
|
95
|
+
attr = getattr(target, name, None)
|
|
96
|
+
if attr is not None:
|
|
97
|
+
return attr
|
|
98
|
+
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
## CORRELATIONS CALCULATION METHOD ##
|
|
102
|
+
|
|
103
|
+
def calculate_all_cors(self, window: int, step: int | None = None, n_steps: int | None = None, verbose: bool=False, dtype: str="float"):
|
|
104
|
+
'''
|
|
105
|
+
Calculate all lagged correlations.
|
|
106
|
+
|
|
107
|
+
Parameters
|
|
108
|
+
----------
|
|
109
|
+
window : int
|
|
110
|
+
Number of cells in the window for applying lags (usually n_cells*2/3)
|
|
111
|
+
step : int | None
|
|
112
|
+
Step to apply while shifting window
|
|
113
|
+
n_steps : int | None
|
|
114
|
+
Total number of steps to consider while shifting window
|
|
115
|
+
verbose : bool
|
|
116
|
+
Wether show progression bar (slightly slow down the process)
|
|
117
|
+
dtype : type
|
|
118
|
+
Data type used to store results. Should be 'float' or 'int' (recommanded to reduce memory use).
|
|
119
|
+
|
|
120
|
+
Return
|
|
121
|
+
------
|
|
122
|
+
Update tensor, window, reflag, MAC, LAG_id, LAG, LAG_means, LAG_scaled attributes, as well as step, n_steps, dtype and num_format
|
|
123
|
+
'''
|
|
124
|
+
|
|
125
|
+
n_genes, n_cells = self.data.shape
|
|
126
|
+
resolved_dtype, num_format = _set_numerical_format(dtype)
|
|
127
|
+
step, n_steps = _solve_step(step, n_steps, window, n_cells)
|
|
128
|
+
|
|
129
|
+
# Initialize a new instance of CoReGraph_Cors object only if necessary
|
|
130
|
+
if (self.cors.window == window) and (self.cors.dtype == resolved_dtype[0]) and (self.cors.step == step):
|
|
131
|
+
raise ValueError(f'Correlation with window={window} and step={step} have already been calculated.')
|
|
132
|
+
|
|
133
|
+
# Warning if some cells are excluded
|
|
134
|
+
excluded_cells = n_cells - (window + ((n_steps-1) * step))
|
|
135
|
+
if excluded_cells != 0:
|
|
136
|
+
warnings.warn(f'Provided step and window exclude the last {excluded_cells} cells.')
|
|
137
|
+
|
|
138
|
+
# Remove warning while dividing by zero and managing defined num type
|
|
139
|
+
np.seterr(divide='ignore', invalid='ignore')
|
|
140
|
+
|
|
141
|
+
# Initialize loops and empty objects
|
|
142
|
+
total_dim3 = _solve_total_dim3(step, window, n_cells)
|
|
143
|
+
tensor = np.empty((n_genes, n_genes, total_dim3), dtype=resolved_dtype)
|
|
144
|
+
MAC = np.zeros((n_genes, n_genes), dtype=resolved_dtype)
|
|
145
|
+
LAG_id = np.empty((2, n_genes, n_genes), dtype=np.int64)
|
|
146
|
+
LAG = np.zeros((n_genes, n_genes), dtype=np.int64)
|
|
147
|
+
reflag = np.empty((total_dim3, 2))
|
|
148
|
+
dim3 = 0
|
|
149
|
+
LAG_means = np.empty((total_dim3, n_genes), dtype=np.float64)
|
|
150
|
+
|
|
151
|
+
# Iterate through refs (iterator) and lags (inner function)
|
|
152
|
+
iterator = tqdm(range(0, n_cells-window+1, step), desc="Calculating Correlations", disable=not verbose)
|
|
153
|
+
for i in iterator:
|
|
154
|
+
dim3, MAC, LAG_id, LAG, LAG_means = _calculate_all_cors(self.data, window, i, tensor, reflag, MAC, LAG_id, LAG, LAG_means, dim3, step, num_format, resolved_dtype)
|
|
155
|
+
|
|
156
|
+
# Fill in attributes
|
|
157
|
+
self.cors.MAC = MAC
|
|
158
|
+
self.cors.LAG_id = LAG_id
|
|
159
|
+
self.cors.LAG = LAG
|
|
160
|
+
self.cors.LAG_means = LAG_means
|
|
161
|
+
self.cors.tensor = tensor
|
|
162
|
+
self.cors.window = window
|
|
163
|
+
self.cors.step = step
|
|
164
|
+
self.cors.n_steps = n_steps
|
|
165
|
+
self.cors.reflag = reflag
|
|
166
|
+
max = np.max(self.cors.LAG_means.T, axis=1).T
|
|
167
|
+
min = np.min(self.cors.LAG_means.T, axis=1).T
|
|
168
|
+
self.cors.LAG_scaled = (self.cors.LAG_means - min) / (max - min)
|
|
169
|
+
self.cors.dtype = resolved_dtype
|
|
170
|
+
self.cors.num_format = num_format
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
## FDR CALCULATION METHOD ##
|
|
174
|
+
|
|
175
|
+
@requires_cors
|
|
176
|
+
def calculate_FDR(self, n_perms = 100, FDR_cutoffs = 101):
|
|
177
|
+
'''
|
|
178
|
+
Calculate permutations to estimate False Discovery Rate (FDR).
|
|
179
|
+
|
|
180
|
+
Parameters
|
|
181
|
+
----------
|
|
182
|
+
n_perms : int
|
|
183
|
+
Number of permutation to perform for estimation
|
|
184
|
+
FDR_cutoffs : int
|
|
185
|
+
Resolution of FDR estimate, number of bins between 0 to 1 that are displayed in results.
|
|
186
|
+
|
|
187
|
+
Return
|
|
188
|
+
------
|
|
189
|
+
Update FDR attribute
|
|
190
|
+
|
|
191
|
+
'''
|
|
192
|
+
|
|
193
|
+
FDR = _calculate_FDR(self.data, self.cors.MAC/self.cors.num_format[0], self.cors.window, n_perms, FDR_cutoffs, self.cors.step)
|
|
194
|
+
FDR = pd.DataFrame(FDR, columns=["cors", "MACs_observed", "MACs_ave_perm", "FDR"])
|
|
195
|
+
FDR[["MACs_observed"]] = FDR[["MACs_observed"]].astype(int)
|
|
196
|
+
self.cors.FDR = FDR
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
## BIN DISCRETIZATION METHOD ##
|
|
200
|
+
|
|
201
|
+
def calculate_bins(self, n_bins: int):
|
|
202
|
+
'''
|
|
203
|
+
Calculate pseudotime bins.
|
|
204
|
+
|
|
205
|
+
Parameters
|
|
206
|
+
----------
|
|
207
|
+
n_bins : int
|
|
208
|
+
Number of bins to consider for pseudotime discretization
|
|
209
|
+
'''
|
|
210
|
+
# Initialize a new instance of CoReGraph_Graphics object only if necessary
|
|
211
|
+
if self.graphics.n_bins == n_bins:
|
|
212
|
+
raise ValueError(f'Bins for n_bins={n_bins} have already been calculated.')
|
|
213
|
+
|
|
214
|
+
self.graphics.bin_id, self.graphics.mean, self.graphics.std, self.graphics.scaled = _calculate_bins(self.data, self.pseudotime, n_bins)
|
|
215
|
+
self.graphics.n_bins = n_bins
|
|
216
|
+
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
|
coregraph/models/cors.py
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
## CORS ##
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import pandas as pd
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@dataclass
|
|
10
|
+
class CoReGraph_Cors:
|
|
11
|
+
'''
|
|
12
|
+
Container for correlation results.
|
|
13
|
+
|
|
14
|
+
Attributes
|
|
15
|
+
----------
|
|
16
|
+
window : int
|
|
17
|
+
Window used for the current correlations
|
|
18
|
+
step : int
|
|
19
|
+
Step used while looping over pseudotime
|
|
20
|
+
n_steps : int
|
|
21
|
+
Total number of steps considered along pseudotime
|
|
22
|
+
tensor : (3D)ndarray
|
|
23
|
+
Result of all lagged or not correlations
|
|
24
|
+
reflag : ndarray
|
|
25
|
+
Vector of (ref,lag) corresponding to tensor dimensions
|
|
26
|
+
MAC : ndarray
|
|
27
|
+
Maximum Absolute Correlation values
|
|
28
|
+
LAG_id : ndarray
|
|
29
|
+
Reflag corresponding to the dimension used for MAC
|
|
30
|
+
LAG : ndarray
|
|
31
|
+
Number of cells corresponding to the lag used for MAC
|
|
32
|
+
LAG_means : (3D)ndarray
|
|
33
|
+
Gene mean expressions calculated for each window
|
|
34
|
+
LAG_scaled : (3D)ndarray
|
|
35
|
+
Gene scaled expressions calculated for each window
|
|
36
|
+
FDR : DataFrame
|
|
37
|
+
Pandas DataFrame recapitulating metrics for False Dicovery Rate estimation
|
|
38
|
+
dtype : str
|
|
39
|
+
Numpy type of object to store correlation values (should be 'int' or 'float')
|
|
40
|
+
num_format : ndarray
|
|
41
|
+
Numpy array of two elements used to define the numerical format [scale_factor, n_decimals]
|
|
42
|
+
'''
|
|
43
|
+
|
|
44
|
+
window: int | None = None
|
|
45
|
+
step: int | None = None
|
|
46
|
+
n_steps: int | None = None
|
|
47
|
+
tensor: np.ndarray | None = None
|
|
48
|
+
reflag: np.ndarray | None = None
|
|
49
|
+
MAC: np.ndarray | None = None
|
|
50
|
+
LAG_id: np.ndarray | None = None
|
|
51
|
+
LAG: np.ndarray | None = None
|
|
52
|
+
LAG_means: np.ndarray | None = None
|
|
53
|
+
LAG_scaled: np.ndarray | None = None
|
|
54
|
+
FDR: pd.DataFrame | None = None
|
|
55
|
+
dtype: np.dtype | None = None
|
|
56
|
+
num_format: np.ndarray | None = None
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def __repr__(self):
|
|
60
|
+
return f'CoReGraph_Cors(window: {self.window}, step: {self.step})'
|
|
61
|
+
|
|
62
|
+
def __str__(self):
|
|
63
|
+
return f'CoReGraph_Cors object'
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
## GRAPHICS ##
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass
|
|
9
|
+
class CoReGraph_Graphics:
|
|
10
|
+
'''
|
|
11
|
+
Container for graphical data.
|
|
12
|
+
|
|
13
|
+
Attributes
|
|
14
|
+
----------
|
|
15
|
+
n_bins : int
|
|
16
|
+
Bin number used to discretize pseudotime
|
|
17
|
+
bin_id : ndarray
|
|
18
|
+
Vector of bin attributed to each cell
|
|
19
|
+
mean : ndarray
|
|
20
|
+
Binned mean expression
|
|
21
|
+
std : ndarray
|
|
22
|
+
Binned standard deviation
|
|
23
|
+
scaled : ndarray
|
|
24
|
+
Binned scaled expression
|
|
25
|
+
'''
|
|
26
|
+
|
|
27
|
+
n_bins: int | None = None
|
|
28
|
+
bin_id: np.ndarray | None = None
|
|
29
|
+
mean: np.ndarray | None = None
|
|
30
|
+
std: np.ndarray | None = None
|
|
31
|
+
scaled: np.ndarray | None = None
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def __repr__(self):
|
|
35
|
+
return f'CoReGraph_Graphics(n_bins: {self.n_bins})'
|
|
36
|
+
|
|
37
|
+
def __str__(self):
|
|
38
|
+
return f'CoReGraph_Graphics object'
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
## STARGRAPH ##
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import pandas as pd
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@dataclass
|
|
10
|
+
class StarGraph:
|
|
11
|
+
'''
|
|
12
|
+
Local star-shaped regulatory network.
|
|
13
|
+
'''
|
|
14
|
+
|
|
15
|
+
central_gene: str | None = None
|
|
16
|
+
window: int | None = None
|
|
17
|
+
step: int | None = None
|
|
18
|
+
n_steps: int | None = None
|
|
19
|
+
FDR_thr: float | None = None
|
|
20
|
+
cors_thr: float | None = None
|
|
21
|
+
MAC: pd.DataFrame | None = None
|
|
22
|
+
FDR: pd.DataFrame | None = None
|
|
23
|
+
gene_list: list[str] | None = None
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def __repr__(self):
|
|
27
|
+
return f'StarGraph(central_gene: {self.central_gene}, window: {self.window})'
|
|
28
|
+
|
|
29
|
+
def __str__(self):
|
|
30
|
+
return f'StarGraph object'
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def set_correlation_threshold(self, threshold: float):
|
|
34
|
+
self.cors_thr = threshold
|
|
35
|
+
self.FDR_thr = np.max(self.FDR.loc[self.FDR['cors'] > threshold, 'FDR'])
|
|
36
|
+
self.gene_list = (self.MAC.loc[np.abs(self.MAC['MAC']) >= threshold].index.tolist())
|
|
37
|
+
|
|
38
|
+
def set_FDR_threshold(self, threshold: float):
|
|
39
|
+
self.fdr_threshold = threshold
|
|
40
|
+
self.correlation_threshold = np.min(self.FDR.loc[self.FDR['FDR'] < threshold, 'cors'])
|
|
41
|
+
self.gene_list = (self.MAC.loc[np.abs(self.MAC['MAC']) >= self.correlation_threshold].index.tolist())
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
from .decorators import (
|
|
2
|
+
requires_cors,
|
|
3
|
+
requires_FDR,
|
|
4
|
+
requires_bins,
|
|
5
|
+
)
|
|
6
|
+
|
|
7
|
+
from .formatting import (
|
|
8
|
+
_set_numerical_format,
|
|
9
|
+
)
|
|
10
|
+
|
|
11
|
+
from .validation import (
|
|
12
|
+
_solve_step,
|
|
13
|
+
_solve_total_dim3,
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
__all__ = [
|
|
17
|
+
"requires_cors",
|
|
18
|
+
"requires_FDR",
|
|
19
|
+
"requires_bins",
|
|
20
|
+
"_set_numerical_format",
|
|
21
|
+
"_solve_step",
|
|
22
|
+
"_solve_total_dim3",
|
|
23
|
+
]
|