scdiagnostics 0.0.99.1__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,15 @@
1
+ from .dimred import plot_umap, plot_pca, compare_pca, compare_umap
2
+ from .marginal import (
3
+ compare_boxplot,
4
+ compare_ecdf,
5
+ compare_histogram,
6
+ compare_means,
7
+ compare_moments,
8
+ compare_standard_deviation,
9
+ compare_variances,
10
+ compare_histogram2
11
+ )
12
+ from .spatial import (
13
+ plot_dispersion_surface,
14
+ plot_mean_surface
15
+ )
scdiagnostics/data.py ADDED
@@ -0,0 +1,39 @@
1
+ import numpy as np
2
+ import pandas as pd
3
+
4
+
5
+ def adata_df(adata):
6
+ return (
7
+ pd.DataFrame(check_sparse(adata.X), columns=adata.var_names)
8
+ .melt(id_vars=[], value_vars=adata.var_names)
9
+ .reset_index(drop=True)
10
+ )
11
+
12
+
13
+ def merge_samples(adata, sim):
14
+ source = adata_df(adata)
15
+ simulated = adata_df(sim)
16
+ return pd.concat(
17
+ {"real": source, "simulated": simulated}, names=["source"]
18
+ ).reset_index(level="source")
19
+
20
+
21
+ def check_sparse(X):
22
+ if not isinstance(X, np.ndarray):
23
+ X = X.todense()
24
+ return X
25
+
26
+
27
+ def prepare_dense(real, simulated):
28
+ real_ = real.copy()
29
+ simulated_ = simulated.copy()
30
+ real_.X = check_sparse(real_.X)
31
+ simulated_.X = check_sparse(simulated_.X)
32
+ return real_, simulated_
33
+
34
+
35
+ def concat_real_sim(real, simulated):
36
+ real_, simulated_ = prepare_dense(real, simulated)
37
+ real_.obs["source"] = "real"
38
+ simulated_.obs["source"] = "simulated"
39
+ return real_.concatenate(simulated_, join="outer", batch_key=None)
@@ -0,0 +1,96 @@
1
+ import scanpy as sc
2
+ import numpy as np
3
+ import pandas as pd
4
+ import altair as alt
5
+ from .data import check_sparse, concat_real_sim
6
+
7
+
8
+ def plot_umap(
9
+ adata,
10
+ color=None,
11
+ shape=None,
12
+ facet=None,
13
+ opacity=0.6,
14
+ n_comps=20,
15
+ n_neighbors=15,
16
+ transform=np.log1p,
17
+ **kwargs,
18
+ ):
19
+ mapping = {
20
+ "x": alt.X("UMAP1", scale=alt.Scale(zero=False)),
21
+ "y": alt.Y("UMAP2", scale=alt.Scale(zero=False)),
22
+ "color": color,
23
+ "shape": shape,
24
+ }
25
+ mapping = {k: v for k, v in mapping.items() if v is not None}
26
+
27
+ adata_ = adata.copy()
28
+ adata_.X = check_sparse(adata_.X)
29
+ Z = transform(adata_.X)
30
+ if Z.shape[1] == adata_.X.shape[1]:
31
+ adata_.X = transform(adata_.X)
32
+ else:
33
+ adata_ = adata_[:, : Z.shape[1]]
34
+ adata_.X = Z
35
+ adata_.var_names = [f"transform_{k}" for k in range(Z.shape[1])]
36
+
37
+ # umap on the top PCA dimensions
38
+ sc.pp.pca(adata_, n_comps=n_comps)
39
+ sc.pp.neighbors(adata_, n_neighbors=n_neighbors, n_pcs=n_comps)
40
+ sc.tl.umap(adata_, **kwargs)
41
+
42
+ # get umap embeddings
43
+ umap_df = pd.DataFrame(adata_.obsm["X_umap"], columns=["UMAP1", "UMAP2"])
44
+ umap_df = pd.concat([umap_df, adata_.obs.reset_index(drop=True)], axis=1)
45
+
46
+ # encode and visualize
47
+ alt.data_transformers.enable("vegafusion")
48
+ chart = alt.Chart(umap_df).mark_point(opacity=opacity).encode(**mapping)
49
+ if facet is not None:
50
+ chart = chart.facet(column=alt.Facet(facet))
51
+ return chart
52
+
53
+
54
+ def plot_pca(
55
+ adata,
56
+ color=None,
57
+ shape=None,
58
+ facet=None,
59
+ opacity=0.6,
60
+ plot_dims=[0, 1],
61
+ transform=lambda x: np.log1p(x),
62
+ **kwargs,
63
+ ):
64
+ mapping = {
65
+ "x": alt.X("PCA1", scale=alt.Scale(zero=False)),
66
+ "y": alt.Y("PCA2", scale=alt.Scale(zero=False)),
67
+ "color": color,
68
+ "shape": shape,
69
+ }
70
+ mapping = {k: v for k, v in mapping.items() if v is not None}
71
+
72
+ adata_ = adata.copy()
73
+ adata_.X = check_sparse(adata_.X)
74
+ adata_.X = transform(adata_.X)
75
+
76
+ # get PCA scores
77
+ sc.pp.pca(adata_, **kwargs)
78
+ pca_df = pd.DataFrame(adata_.obsm["X_pca"][:, plot_dims], columns=["PCA1", "PCA2"])
79
+ pca_df = pd.concat([pca_df, adata_.obs.reset_index(drop=True)], axis=1)
80
+
81
+ # plot
82
+ alt.data_transformers.enable("vegafusion")
83
+ chart = alt.Chart(pca_df).mark_point(opacity=opacity).encode(**mapping)
84
+ if facet is not None:
85
+ chart = chart.facet(column=alt.Facet(facet))
86
+ return chart
87
+
88
+
89
+ def compare_umap(real, simulated, transform=lambda x: x, **kwargs):
90
+ adata = concat_real_sim(real, simulated)
91
+ return plot_umap(adata, facet="source", transform=transform, **kwargs)
92
+
93
+
94
+ def compare_pca(real, simulated, transform=lambda x: x, **kwargs):
95
+ adata = concat_real_sim(real, simulated)
96
+ return plot_pca(adata, facet="source", transform=transform, **kwargs)
@@ -0,0 +1,178 @@
1
+ import numpy as np
2
+ import pandas as pd
3
+ import altair as alt
4
+ import matplotlib.pyplot as plt
5
+ from .data import prepare_dense, merge_samples
6
+
7
+
8
+ def compare_summary(real, simulated, summary_fun, labels=None):
9
+ df = pd.DataFrame({"real": summary_fun(real), "simulated": summary_fun(simulated)})
10
+
11
+ identity = pd.DataFrame(
12
+ {
13
+ "real": [df["real"].min(), df["real"].max()],
14
+ "simulated": [df["real"].min(), df["real"].max()],
15
+ }
16
+ )
17
+ chart = alt.Chart(identity).mark_line(color="#dedede").encode(
18
+ x="real", y="simulated"
19
+ ) + alt.Chart(df).mark_circle().encode(x="real", y="simulated")
20
+
21
+ if labels is not None:
22
+ df["label"] = labels
23
+ chart = chart + alt.Chart(df[df["label"] != ""]).mark_text(
24
+ dx=6, dy=-6, align="left"
25
+ ).encode(x="real:Q", y="simulated:Q", text="label:N")
26
+
27
+ return chart
28
+
29
+
30
+ def compare_means(real, simulated, transform=np.log1p, labels=None):
31
+ real_, simulated_ = prepare_dense(real, simulated)
32
+ summary = lambda a: np.asarray(transform(a.X).mean(axis=0)).flatten()
33
+ return compare_summary(real_, simulated_, summary, labels)
34
+
35
+
36
+ def compare_variances(real, simulated, transform=np.log1p, labels=None):
37
+ real_, simulated_ = prepare_dense(real, simulated)
38
+ summary = lambda a: np.asarray(np.var(transform(a.X), axis=0)).flatten()
39
+ return compare_summary(real_, simulated_, summary, labels)
40
+
41
+
42
+ def compare_standard_deviation(real, simulated, transform=np.log1p, labels=None):
43
+ real_, simulated_ = prepare_dense(real, simulated)
44
+ summary = lambda a: np.asarray(np.std(transform(a.X), axis=0)).flatten()
45
+ return compare_summary(real_, simulated_, summary, labels)
46
+
47
+
48
+ def compare_histogram2(sim_data, real_data, idx):
49
+ sim = sim_data[:, idx]
50
+ real = real_data[:, idx]
51
+ b = np.linspace(min(min(sim), min(real)), max(max(sim), max(real)), 50)
52
+
53
+ plt.hist([real, sim], b, label=["Real", "Simulated"], histtype="bar")
54
+ plt.xlabel("x")
55
+ plt.ylabel("Density")
56
+ plt.legend()
57
+ plt.show()
58
+
59
+
60
+ def compare_ecdf(adata, sim, var_names=None, max_plot=10, n_cols=5, transform=np.log1p, **kwargs):
61
+ if var_names is None:
62
+ var_names = adata.var_names[:max_plot]
63
+
64
+ combined = merge_samples(adata[:, var_names], sim[:, var_names])
65
+ combined["value"] = transform(combined["value"])
66
+ alt.data_transformers.enable("vegafusion")
67
+
68
+ plot = (
69
+ alt.Chart(combined)
70
+ .transform_window(
71
+ ecdf="cume_dist()", sort=[{"field": "value"}], groupby=["variable", "source"]
72
+ )
73
+ .mark_line(
74
+ interpolate="step-after",
75
+ )
76
+ .encode(
77
+ x="value:Q",
78
+ y="ecdf:Q",
79
+ color="source:N",
80
+ facet=alt.Facet(
81
+ "variable", sort=alt.EncodingSortField("value"), columns=n_cols
82
+ ),
83
+ )
84
+ .properties(**kwargs)
85
+ )
86
+ plot.show()
87
+ return plot, combined
88
+
89
+
90
+ def compare_boxplot(adata, sim, var_names=None, max_plot=20, transform=np.log1p, **kwargs):
91
+ if var_names is None:
92
+ var_names = adata.var_names[:max_plot]
93
+
94
+ combined = merge_samples(adata[:, var_names], sim[:, var_names])
95
+ combined["value"] = transform(combined["value"])
96
+ alt.data_transformers.enable("vegafusion")
97
+
98
+ plot = (
99
+ alt.Chart(combined)
100
+ .mark_boxplot(extent="min-max")
101
+ .encode(
102
+ x=alt.X("value:Q").scale(zero=False),
103
+ y=alt.Y(
104
+ "variable:N",
105
+ sort=alt.EncodingSortField("mid_box_value", order="descending"),
106
+ ),
107
+ facet="source:N",
108
+ )
109
+ .properties(**kwargs)
110
+ )
111
+ plot.show()
112
+ return plot, combined
113
+
114
+
115
+ def compare_histogram(adata, sim, var_names=None, max_plot=10, transform=np.log1p, ncol=5, **kwargs):
116
+ if var_names is None:
117
+ var_names = adata.var_names[:max_plot]
118
+
119
+ combined = merge_samples(adata[:, var_names], sim[:, var_names])
120
+ combined["value"] = transform(combined["value"])
121
+ alt.data_transformers.enable("vegafusion")
122
+
123
+ plot = (
124
+ alt.Chart(combined)
125
+ .mark_bar(opacity=0.7)
126
+ .encode(
127
+ x=alt.X("value:Q").bin(maxbins=20),
128
+ y=alt.Y("count()").stack(None),
129
+ color="source:N",
130
+ facet=alt.Facet(
131
+ "variable", sort=alt.EncodingSortField(f"bin_maxbins_{ncol}_value"), columns=ncol
132
+ ),
133
+ )
134
+ .properties(**kwargs)
135
+ )
136
+ plot.show()
137
+ return plot, combined
138
+
139
+
140
+ def compare_moments(real, simulated, log_scale=True, labels=None, label_threshold=None):
141
+ """Compare real vs. simulated means and standard deviations gene-wise.
142
+
143
+ Parameters
144
+ ----------
145
+ log_scale : bool
146
+ If True, log-transform the gene-level statistics before plotting
147
+ (useful when counts span several orders of magnitude).
148
+ labels : array-like of str, optional
149
+ Gene labels aligned with the variables. Non-empty strings are
150
+ rendered as text marks next to the corresponding point.
151
+ label_threshold : float, optional
152
+ If provided, genes where the absolute difference between real and
153
+ simulated exceeds this value in either the mean or SD are labeled
154
+ with their gene name in both panels.
155
+ """
156
+ real_, simulated_ = prepare_dense(real, simulated)
157
+ transform = np.log if log_scale else lambda x: x
158
+
159
+ def gene_summary(stat_fn):
160
+ def summary(a):
161
+ return transform(np.asarray(stat_fn(a.X)).flatten())
162
+ return summary
163
+
164
+ means_summary = gene_summary(lambda X: X.mean(axis=0))
165
+ sd_summary = gene_summary(lambda X: np.std(X, axis=0))
166
+
167
+ if label_threshold is not None:
168
+ means_real = means_summary(real_)
169
+ means_sim = means_summary(simulated_)
170
+ sd_real = sd_summary(real_)
171
+ sd_sim = sd_summary(simulated_)
172
+ exceeds = (np.abs(means_real - means_sim) > label_threshold) | (np.abs(sd_real - sd_sim) > label_threshold)
173
+ labels = np.where(exceeds, real_.var_names, "")
174
+
175
+ return (
176
+ compare_summary(real_, simulated_, means_summary, labels).properties(title="Means")
177
+ | compare_summary(real_, simulated_, sd_summary, labels).properties(title="Standard Deviations")
178
+ )
@@ -0,0 +1,190 @@
1
+ from scipy.interpolate import RBFInterpolator
2
+ import scipy.sparse as sp
3
+ import altair as alt
4
+ import numpy as np
5
+ import pandas as pd
6
+
7
+
8
+ def make_grid_obs(adata, basis_cols, n_grid=30):
9
+ s1 = adata.obs["spatial1"].values
10
+ s2 = adata.obs["spatial2"].values
11
+ g1 = np.linspace(s1.min(), s1.max(), n_grid)
12
+ g2 = np.linspace(s2.min(), s2.max(), n_grid)
13
+ G1, G2 = np.meshgrid(g1, g2, indexing="ij")
14
+
15
+ grid_base = pd.DataFrame(
16
+ {
17
+ "spatial1": G1.ravel(),
18
+ "spatial2": G2.ravel(),
19
+ "bin1": pd.cut(G1.ravel(), bins=n_grid, labels=False, include_lowest=True),
20
+ "bin2": pd.cut(G2.ravel(), bins=n_grid, labels=False, include_lowest=True),
21
+ }
22
+ )
23
+ obs_bins = (
24
+ pd.DataFrame(
25
+ {
26
+ "bin1": pd.cut(s1, bins=n_grid, labels=False, include_lowest=True),
27
+ "bin2": pd.cut(s2, bins=n_grid, labels=False, include_lowest=True),
28
+ }
29
+ )
30
+ .dropna()
31
+ .astype(int)
32
+ .drop_duplicates()
33
+ )
34
+ occupied = set(map(tuple, obs_bins[["bin1", "bin2"]].to_numpy()))
35
+ keep = pd.Series(list(zip(grid_base["bin1"], grid_base["bin2"]))).isin(occupied)
36
+ grid_base = grid_base.loc[keep, ["spatial1", "spatial2"]].reset_index(drop=True)
37
+
38
+ cell_types = list(adata.obs["cell_type"].unique())
39
+ n_pts = len(grid_base)
40
+ grid_obs = pd.DataFrame(
41
+ {
42
+ "spatial1": np.tile(grid_base["spatial1"].values, len(cell_types)),
43
+ "spatial2": np.tile(grid_base["spatial2"].values, len(cell_types)),
44
+ "cell_type": np.repeat(cell_types, n_pts),
45
+ }
46
+ )
47
+
48
+ train_coords = adata.obs[["spatial1", "spatial2"]].values
49
+ grid_coords = grid_obs[["spatial1", "spatial2"]].values
50
+ interp = RBFInterpolator(train_coords, adata.obs[basis_cols].values, neighbors=50)
51
+ grid_basis = interp(grid_coords)
52
+ for i, col in enumerate(basis_cols):
53
+ grid_obs[col] = grid_basis[:, i]
54
+ return grid_obs
55
+
56
+
57
+ def dense_gene_counts(adata, gene_idx):
58
+ if sp.issparse(adata.X):
59
+ return adata.X[:, gene_idx].toarray().ravel()
60
+ return np.array(adata.X)[:, gene_idx]
61
+
62
+
63
+ def binned_obs_df(adata, gene_idx, n_obs_bins, value_col):
64
+ s1 = adata.obs["spatial1"].values
65
+ s2 = adata.obs["spatial2"].values
66
+ return pd.DataFrame(
67
+ {
68
+ "spatial1": pd.cut(s1, bins=n_obs_bins)
69
+ .map(lambda b: round(b.mid, 2))
70
+ .astype(float),
71
+ "spatial2": pd.cut(s2, bins=n_obs_bins)
72
+ .map(lambda b: round(b.mid, 2))
73
+ .astype(float),
74
+ value_col: dense_gene_counts(adata, gene_idx),
75
+ }
76
+ )
77
+
78
+
79
+ def fitted_surface_df(sim, adata, basis_cols, gene_idx, pred_key, value_col, n_grid):
80
+ grid_obs = make_grid_obs(adata, basis_cols, n_grid)
81
+ n_ct = adata.obs["cell_type"].nunique()
82
+ n_pts = len(grid_obs) // n_ct
83
+ pred = np.log1p(sim.predict(obs=grid_obs)[pred_key][:, gene_idx])
84
+ pred_avg = pred.reshape(n_ct, n_pts).mean(axis=0)
85
+ ref = grid_obs.iloc[:n_pts]
86
+ return pd.DataFrame(
87
+ {
88
+ "spatial1": ref["spatial1"].round(2).values,
89
+ "spatial2": ref["spatial2"].round(2).values,
90
+ value_col: pred_avg,
91
+ }
92
+ )
93
+
94
+
95
+ def surface_heatmaps(
96
+ fitted_df, obs_df, value_col, color_title, fitted_title, obs_title
97
+ ):
98
+ vmax = max(fitted_df[value_col].max(), obs_df[value_col].max())
99
+ scale = alt.Scale(domain=[0, vmax], scheme="viridis")
100
+
101
+ def heatmap(df, title):
102
+ return (
103
+ alt.Chart(df)
104
+ .mark_rect()
105
+ .encode(
106
+ x=alt.X("spatial1:O", title="", axis=alt.Axis(labels=False, ticks=False, domain=False)),
107
+ y=alt.Y("spatial2:O", title="", axis=alt.Axis(labels=False, ticks=False, domain=False)),
108
+ color=alt.Color(f"{value_col}:Q", scale=scale, title=color_title),
109
+ )
110
+ .properties(width=300, height=300, title=title)
111
+ )
112
+
113
+ return heatmap(fitted_df, fitted_title) | heatmap(obs_df, obs_title)
114
+
115
+
116
+ def plot_mean_surface(sim, adata, basis_cols=["spatial1", "spatial2"], gene=None, n_grid=30, n_obs_bins=30):
117
+ # Handle gene specification - automatically detect if gene is index or name
118
+ if isinstance(gene, str):
119
+ gene_idx = adata.var_names.get_loc(gene)
120
+ gene_name = gene
121
+ elif isinstance(gene, int):
122
+ gene_idx = gene
123
+ gene_name = adata.var_names[gene_idx]
124
+ else:
125
+ raise ValueError("gene must be either integer index or string name")
126
+
127
+ fitted_df = fitted_surface_df(
128
+ sim,
129
+ adata,
130
+ basis_cols,
131
+ gene_idx,
132
+ pred_key="mean",
133
+ value_col="mu",
134
+ n_grid=n_grid,
135
+ )
136
+ obs_df = (
137
+ binned_obs_df(adata, gene_idx, n_obs_bins, value_col="mu")
138
+ .groupby(["spatial1", "spatial2"], as_index=False)["mu"]
139
+ .mean()
140
+ .assign(mu=lambda d: np.log1p(d["mu"]))
141
+ )
142
+ return surface_heatmaps(
143
+ fitted_df,
144
+ obs_df,
145
+ value_col="mu",
146
+ color_title="log(mu+1)",
147
+ fitted_title=f"Fitted mean(x) - Gene {gene_name}",
148
+ obs_title=f"Observed bin mean - Gene {gene_name}",
149
+ )
150
+
151
+
152
+ def plot_dispersion_surface(sim, adata, basis_cols=["spatial1", "spatial2"], gene=None, n_grid=10, n_obs_bins=10):
153
+ if isinstance(gene, str):
154
+ gene_idx = adata.var_names.get_loc(gene)
155
+ gene_name = gene
156
+ elif isinstance(gene, int):
157
+ gene_idx = gene
158
+ gene_name = adata.var_names[gene_idx]
159
+ else:
160
+ raise ValueError("gene must be either integer index or string name")
161
+
162
+ fitted_df = fitted_surface_df(
163
+ sim,
164
+ adata,
165
+ basis_cols,
166
+ gene_idx,
167
+ pred_key="dispersion",
168
+ value_col="dispersion",
169
+ n_grid=n_grid,
170
+ )
171
+ obs_df = (
172
+ binned_obs_df(adata, gene_idx, n_obs_bins, value_col="count")
173
+ .groupby(["spatial1", "spatial2"], as_index=False)["count"]
174
+ .agg(mu="mean", var="var")
175
+ .assign(
176
+ dispersion=lambda d: np.where(
177
+ d["var"] > d["mu"], d["mu"] ** 2 / (d["var"] - d["mu"]), np.nan
178
+ )
179
+ )
180
+ .assign(dispersion=lambda d: np.log1p(d["dispersion"]))
181
+ .dropna(subset=["dispersion"])
182
+ )
183
+ return surface_heatmaps(
184
+ fitted_df,
185
+ obs_df,
186
+ value_col="dispersion",
187
+ color_title="log(dispersion+1)",
188
+ fitted_title=f"Fitted dispersion(x) - Gene {gene_name}",
189
+ obs_title=f"Observed bin dispersion - Gene {gene_name}",
190
+ )
@@ -0,0 +1,14 @@
1
+ Metadata-Version: 2.4
2
+ Name: scdiagnostics
3
+ Version: 0.0.99.1
4
+ Author-email: Kris Sankaran <ksankaran@wisc.edu>
5
+ Requires-Python: >=3.10
6
+ Requires-Dist: altair==6.0.0
7
+ Requires-Dist: pyarrow==23.0.1
8
+ Requires-Dist: scanpy==1.12
9
+ Requires-Dist: scipy==1.15.2
10
+ Requires-Dist: vegafusion==2.0.3
11
+ Requires-Dist: vl-convert-python>=1.8.0
12
+ Provides-Extra: dev
13
+ Requires-Dist: pytest; extra == "dev"
14
+ Requires-Dist: ruff; extra == "dev"
@@ -0,0 +1,9 @@
1
+ scdiagnostics/__init__.py,sha256=OtPZabaI4wXtscFOyULey_9bN6-2bb7oTYf6t7WXAq8,348
2
+ scdiagnostics/data.py,sha256=pmefSH7bmKzP4uWV-XLHLFdgPlRra6j8GkbwCMZGCmQ,998
3
+ scdiagnostics/dimred.py,sha256=g2JXwkYz0H5B6z3-4rSM3sURAyTPnU_PnffHXuQdXuM,2821
4
+ scdiagnostics/marginal.py,sha256=rwPnjJisTpP44bsSqnJdMjQkn9JgASWYlVZhMB_Hak8,6186
5
+ scdiagnostics/spatial.py,sha256=Z-zGrw9lfc3ZPW49PHXMeeNIjCLVyDMO1qXRwS3i_0U,6464
6
+ scdiagnostics-0.0.99.1.dist-info/METADATA,sha256=TEBpxAheLQEomZrSQCAVhnwOpO2Z_5PoOx4jRW-I5GE,417
7
+ scdiagnostics-0.0.99.1.dist-info/WHEEL,sha256=aeYiig01lYGDzBgS8HxWXOg3uV61G9ijOsup-k9o1sk,91
8
+ scdiagnostics-0.0.99.1.dist-info/top_level.txt,sha256=bXGLWwICaXihQ7AVa6iCbiIYWyIzYvuC3y7FvFsuB1I,14
9
+ scdiagnostics-0.0.99.1.dist-info/RECORD,,
@@ -0,0 +1,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (82.0.1)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+
@@ -0,0 +1 @@
1
+ scdiagnostics