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.
@@ -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,11 @@
1
+ from .coregraph import CoReGraph
2
+ from .cors import CoReGraph_Cors
3
+ from .graphics import CoReGraph_Graphics
4
+ from .stargraph import StarGraph
5
+
6
+ __all__ = [
7
+ "CoReGraph",
8
+ "CoReGraph_Cors",
9
+ "CoReGraph_Graphics",
10
+ "StarGraph",
11
+ ]
@@ -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
+
@@ -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
+ ]