Setting the file. One moment. Differential Expression · Scvi Tools · anthropics/knowledge-work-plugins · Skills DocsSpatial Deconvolution
22
Validate Data
Tech Debt
62
Recruiting Pipeline
71
Vendor Check
125
Zoom Meeting SDK Web
88
Vendor Review
181
Create An Asset
Video Sdk/web
scripts/differential_expression.py
scripts/differential_expression.py
Python·220 lines·7 KB
16
17
18def run_de_analysis(
19 model,
20 adata,
21 groupby,
22 group1=None,
23 group2=None,
24 n_genes=None
25):
26 """
27 Run differential expression analysis.
28
29 Parameters
30 ----------
31 model : scvi model
32 Trained model with differential_expression method
33 adata : AnnData
34 Data used for training
35 groupby : str
36 Column in obs to group by
37 group1 : str, optional
38 First group (if None, computes for all groups)
39 group2 : str, optional
40 Second group (rest if None)
41 n_genes : int, optional
42 Limit to top N genes per group
43
44 Returns
45 -------
46 DataFrame with DE results
47 """
48 import pandas as pd
49
50 if group1 is not None:
51 # Specific comparison
52 print(f"Comparing {group1} vs {group2 or 'rest'}...")
53 de_results = model.differential_expression(
54 groupby=groupby,
55 group1=group1,
56 group2=group2
57 )
58
59 # Add comparison info
60 de_results["comparison"] = f"{group1}_vs_{group2 or 'rest'}"
61
62 else:
63 # All pairwise or one-vs-rest
64 groups = adata.obs[groupby].unique()
65 print(f"Computing DE for {len(groups)} groups...")
66
67 all_results = []
68 for group in groups:
69 print(f" Processing {group}...")
70 try:
71 de = model.differential_expression(
72 groupby=groupby,
73 group1=group
74 )
75 de["group"] = group
76 all_results.append(de)
77 except Exception as e:
78 print(f" Warning: Failed for {group}: {e}")
79
80 de_results = pd.concat(all_results, ignore_index=False)
81
82 # Filter to significant
83 if "is_de_fdr_0.05" in de_results.columns:
84 n_sig = de_results["is_de_fdr_0.05"].sum()
85 print(f"Found {n_sig} significant DE genes (FDR < 0.05)")
86
87 # Limit to top genes if requested
88 if n_genes is not None and "lfc_mean" in de_results.columns:
89 if "group" in de_results.columns:
90 # Top N per group
91 de_results = de_results.groupby("group").apply(
92 lambda x: x.nlargest(n_genes, "lfc_mean")
93 ).reset_index(drop=True)
94 else:
95 de_results = de_results.nlargest(n_genes, "lfc_mean")
96
97 return de_results
98
99
100def plot_volcano(de_results, output_path, group_name=None):
101 """Create volcano plot of DE results."""
102 import matplotlib.pyplot as plt
103 import numpy as np
104
105 if "lfc_mean" not in de_results.columns:
106 print("Cannot create volcano plot: missing lfc_mean column")
107 return
108
109 fig, ax = plt.subplots(figsize=(8, 6))
110
111 # Get values
112 lfc = de_results["lfc_mean"].values
113 if "bayes_factor" in de_results.columns:
114 y_val = de_results["bayes_factor"].values
115 y_label = "Bayes Factor"
116 elif "proba_de" in de_results.columns:
117 y_val = -np.log10(1 - de_results["proba_de"].values + 1e-10)
118 y_label = "-log10(1 - P(DE))"
119 else:
120 y_val = np.ones(len(lfc))
121 y_label = ""
122
123 # Color by significance
124 if "is_de_fdr_0.05" in de_results.columns:
125 sig = de_results["is_de_fdr_0.05"].values
126 colors = ["red" if s else "gray" for s in sig]
127 else:
128 colors = "gray"
129
130 ax.scatter(lfc, y_val, c=colors, alpha=0.5, s=10)
131 ax.axvline(0, color="black", linestyle="--", alpha=0.5)
132 ax.set_xlabel("Log Fold Change")
133 ax.set_ylabel(y_label)
134
135 title = "Differential Expression"
136 if group_name:
137 title += f": {group_name}"
138 ax.set_title(title)
139
140 plt.tight_layout()
141 plt.savefig(output_path, dpi=150, bbox_inches="tight")
142 plt.close()
143 print(f"Volcano plot saved to {output_path}")
144
145
146def main():
147 parser = argparse.ArgumentParser(
148 description="Differential expression with scvi-tools",
149 formatter_class=argparse.RawDescriptionHelpFormatter,
150 epilog="""
151Examples:
152 # DE for all clusters (one-vs-rest)
153 python differential_expression.py model/ adata.h5ad de_results.csv --groupby leiden
154
155 # Specific comparison
156 python differential_expression.py model/ adata.h5ad de_results.csv \\
157 --groupby cell_type --group1 "T cells" --group2 "B cells"
158
159 # Top 50 genes per cluster
160 python differential_expression.py model/ adata.h5ad de_results.csv \\
161 --groupby leiden --n-genes 50
162 """
163 )
164 parser.add_argument("model_dir", help="Directory containing saved model")
165 parser.add_argument("input", help="Input h5ad file (same as training)")
166 parser.add_argument("output", help="Output CSV file for DE results")
167 parser.add_argument("--groupby", required=True, help="Column to group by")
168 parser.add_argument("--group1", help="First group for comparison")
169 parser.add_argument("--group2", help="Second group (default: rest)")
170 parser.add_argument("--n-genes", type=int, help="Limit to top N genes per group")
171 parser.add_argument("--model-type", choices=["scvi", "scanvi", "totalvi"],
172 default="scvi", help="Model type (default: scvi)")
173 parser.add_argument("--plot", action="store_true", help="Generate volcano plot")
174
175 args = parser.parse_args()
176
177 try:
178 import scvi
179 import scanpy as sc
180 except ImportError:
181 print("Error: scvi-tools and scanpy required")
182 sys.exit(1)
183
184 # Load data
185 print(f"Loading {args.input}...")
186 adata = sc.read_h5ad(args.input)
187
188 # Load model
189 print(f"Loading model from {args.model_dir}...")
190 if args.model_type == "scvi":
191 model = scvi.model.SCVI.load(args.model_dir, adata=adata)
192 elif args.model_type == "scanvi":
193 model = scvi.model.SCANVI.load(args.model_dir, adata=adata)
194 elif args.model_type == "totalvi":
195 model = scvi.model.TOTALVI.load(args.model_dir, adata=adata)
196
197 # Run DE
198 de_results = run_de_analysis(
199 model,
200 adata,
201 groupby=args.groupby,
202 group1=args.group1,
203 group2=args.group2,
204 n_genes=args.n_genes
205 )
206
207 # Save results
208 de_results.to_csv(args.output)
209 print(f"DE results saved to {args.output}")
210
211 # Plot
212 if args.plot:
213 plot_path = args.output.replace(".csv", "_volcano.png")
214 plot_volcano(de_results, plot_path, args.group1)
215
216 print("\nDone!")
217
218
219if __name__ == "__main__":
220 main()