Setting the file. One moment. Integrate Datasets · 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/integrate_datasets.py
scripts/integrate_datasets.py
Python·237 lines·7 KB
17
def
integrate_datasets
(
18 adatas,
19 batch_names=None,
20 labels_key=None,
21 n_top_genes=2000,
22 n_latent=30,
23 max_epochs=200
24):
25 """
26 Integrate multiple datasets.
27
28 Parameters
29 ----------
30 adatas : list of AnnData
31 Datasets to integrate
32 batch_names : list of str, optional
33 Names for each dataset (default: dataset_0, dataset_1, ...)
34 labels_key : str, optional
35 Cell type column (uses scANVI if provided)
36 n_top_genes : int
37 Number of HVGs
38 n_latent : int
39 Latent dimensions
40 max_epochs : int
41 Training epochs
42
43 Returns
44 -------
45 Integrated AnnData and trained model
46 """
47 import scvi
48 import scanpy as sc
49 import numpy as np
50
51 # Assign batch names
52 if batch_names is None:
53 batch_names = [f"dataset_{i}" for i in range(len(adatas))]
54
55 if len(batch_names) != len(adatas):
56 raise ValueError(f"Number of batch names ({len(batch_names)}) must match datasets ({len(adatas)})")
57
58 # Add batch labels
59 for adata, name in zip(adatas, batch_names):
60 adata.obs["batch"] = name
61 print(f"{name}: {adata.shape}")
62
63 # Find common genes
64 common_genes = set(adatas[0].var_names)
65 for adata in adatas[1:]:
66 common_genes = common_genes.intersection(adata.var_names)
67 common_genes = list(common_genes)
68 print(f"\nCommon genes: {len(common_genes)}")
69
70 # Subset to common genes
71 adatas = [adata[:, common_genes].copy() for adata in adatas]
72
73 # Concatenate
74 print("Concatenating datasets...")
75 adata = sc.concat(adatas, label="batch", keys=batch_names)
76 print(f"Combined: {adata.shape}")
77
78 # Store counts
79 adata.layers["counts"] = adata.X.copy()
80
81 # HVG selection
82 print(f"Selecting {n_top_genes} HVGs...")
83 sc.pp.highly_variable_genes(
84 adata,
85 n_top_genes=n_top_genes,
86 flavor="seurat_v3",
87 batch_key="batch",
88 layer="counts"
89 )
90 adata = adata[:, adata.var["highly_variable"]].copy()
91
92 # Train model
93 if labels_key is not None and labels_key in adata.obs.columns:
94 print(f"\nTraining scANVI with labels ({labels_key})...")
95
96 # First train scVI
97 scvi.model.SCVI.setup_anndata(adata, layer="counts", batch_key="batch")
98 scvi_model = scvi.model.SCVI(adata, n_latent=n_latent)
99 scvi_model.train(max_epochs=max_epochs, early_stopping=True)
100
101 # Then scANVI
102 model = scvi.model.SCANVI.from_scvi_model(
103 scvi_model,
104 labels_key=labels_key,
105 unlabeled_category="Unknown"
106 )
107 model.train(max_epochs=max_epochs // 4)
108
109 adata.obsm["X_scANVI"] = model.get_latent_representation()
110 rep_key = "X_scANVI"
111
112 else:
113 print("\nTraining scVI...")
114 scvi.model.SCVI.setup_anndata(adata, layer="counts", batch_key="batch")
115 model = scvi.model.SCVI(adata, n_latent=n_latent)
116 model.train(max_epochs=max_epochs, early_stopping=True)
117
118 adata.obsm["X_scVI"] = model.get_latent_representation()
119 rep_key = "X_scVI"
120
121 # Cluster
122 print("\nClustering...")
123 sc.pp.neighbors(adata, use_rep=rep_key)
124 sc.tl.umap(adata)
125 sc.tl.leiden(adata)
126
127 print(f"Found {adata.obs['leiden'].nunique()} clusters")
128
129 return adata, model
130
131
132def plot_integration(adata, output_dir, labels_key=None):
133 """Plot integration results."""
134 import scanpy as sc
135 import matplotlib.pyplot as plt
136
137 plots = [
138 ("batch", "By Batch"),
139 ("leiden", "Clusters")
140 ]
141
142 if labels_key is not None and labels_key in adata.obs.columns:
143 plots.append((labels_key, f"Cell Types ({labels_key})"))
144
145 if "predicted_cell_type" in adata.obs.columns:
146 plots.append(("predicted_cell_type", "Predicted Types"))
147
148 n_plots = len(plots)
149 fig, axes = plt.subplots(1, n_plots, figsize=(5 * n_plots, 4))
150 if n_plots == 1:
151 axes = [axes]
152
153 for ax, (color, title) in zip(axes, plots):
154 sc.pl.umap(adata, color=color, ax=ax, show=False, title=title)
155
156 plt.tight_layout()
157 plot_path = os.path.join(output_dir, "integration.png")
158 plt.savefig(plot_path, dpi=150, bbox_inches="tight")
159 plt.close()
160 print(f"Integration plot saved to {plot_path}")
161
162
163def main():
164 parser = argparse.ArgumentParser(
165 description="Integrate multiple datasets with scvi-tools",
166 formatter_class=argparse.RawDescriptionHelpFormatter,
167 epilog="""
168Examples:
169 # Integrate multiple files
170 python integrate_datasets.py results/ data1.h5ad data2.h5ad data3.h5ad
171
172 # With custom batch names
173 python integrate_datasets.py results/ *.h5ad --batch-names ctrl,treat1,treat2
174
175 # With cell type labels (uses scANVI)
176 python integrate_datasets.py results/ *.h5ad --labels-key cell_type
177 """
178 )
179 parser.add_argument("output_dir", help="Output directory")
180 parser.add_argument("inputs", nargs="+", help="Input h5ad files")
181 parser.add_argument("--batch-names", help="Comma-separated batch names")
182 parser.add_argument("--labels-key", help="Cell type column (uses scANVI)")
183 parser.add_argument("--n-hvgs", type=int, default=2000, help="Number of HVGs (default: 2000)")
184 parser.add_argument("--n-latent", type=int, default=30, help="Latent dimensions (default: 30)")
185 parser.add_argument("--max-epochs", type=int, default=200, help="Max epochs (default: 200)")
186
187 args = parser.parse_args()
188
189 try:
190 import scvi
191 import scanpy as sc
192 except ImportError:
193 print("Error: scvi-tools and scanpy required")
194 sys.exit(1)
195
196 # Create output directory
197 os.makedirs(args.output_dir, exist_ok=True)
198
199 # Parse batch names
200 batch_names = None
201 if args.batch_names:
202 batch_names = args.batch_names.split(",")
203
204 # Load datasets
205 print("Loading datasets...")
206 adatas = []
207 for path in args.inputs:
208 print(f" Loading {path}...")
209 adatas.append(sc.read_h5ad(path))
210
211 # Integrate
212 adata, model = integrate_datasets(
213 adatas,
214 batch_names=batch_names,
215 labels_key=args.labels_key,
216 n_top_genes=args.n_hvgs,
217 n_latent=args.n_latent,
218 max_epochs=args.max_epochs
219 )
220
221 # Save results
222 adata_path = os.path.join(args.output_dir, "integrated.h5ad")
223 adata.write_h5ad(adata_path)
224 print(f"\nIntegrated data saved to {adata_path}")
225
226 model_path = os.path.join(args.output_dir, "model")
227 model.save(model_path)
228 print(f"Model saved to {model_path}")
229
230 # Plot
231 plot_integration(adata, args.output_dir, args.labels_key)
232
233 print("\nDone!")
234
235
236if __name__ == "__main__":
237 main()