Setting the file. One moment. Train Model · 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/train_model.py
Python·370 lines·11 KB
16
17
18def train_scvi(adata, batch_key=None, n_latent=30, n_layers=2, max_epochs=200):
19 """Train scVI model."""
20 import scvi
21
22 scvi.model.SCVI.setup_anndata(
23 adata,
24 layer="counts",
25 batch_key=batch_key
26 )
27
28 model = scvi.model.SCVI(
29 adata,
30 n_latent=n_latent,
31 n_layers=n_layers
32 )
33
34 model.train(
35 max_epochs=max_epochs,
36 early_stopping=True,
37 early_stopping_patience=10
38 )
39
40 adata.obsm["X_scVI"] = model.get_latent_representation()
41 return model, "X_scVI"
42
43
44def train_scanvi(adata, batch_key=None, labels_key=None, n_latent=30, n_layers=2, max_epochs=200):
45 """Train scANVI model (scVI + labels)."""
46 import scvi
47
48 # First train scVI
49 scvi.model.SCVI.setup_anndata(
50 adata,
51 layer="counts",
52 batch_key=batch_key
53 )
54
55 scvi_model = scvi.model.SCVI(
56 adata,
57 n_latent=n_latent,
58 n_layers=n_layers
59 )
60 scvi_model.train(max_epochs=max_epochs, early_stopping=True)
61
62 # Initialize scANVI from scVI
63 model = scvi.model.SCANVI.from_scvi_model(
64 scvi_model,
65 labels_key=labels_key,
66 unlabeled_category="Unknown"
67 )
68
69 # Fine-tune scANVI
70 model.train(max_epochs=max_epochs // 4)
71
72 adata.obsm["X_scANVI"] = model.get_latent_representation()
73 return model, "X_scANVI"
74
75
76def train_totalvi(adata, batch_key=None, protein_key="protein_expression", n_latent=20, max_epochs=200):
77 """Train totalVI model for CITE-seq."""
78 import scvi
79 import numpy as np
80
81 scvi.model.TOTALVI.setup_anndata(
82 adata,
83 layer="counts",
84 batch_key=batch_key,
85 protein_expression_obsm_key=protein_key
86 )
87
88 model = scvi.model.TOTALVI(
89 adata,
90 n_latent=n_latent
91 )
92
93 model.train(max_epochs=max_epochs, early_stopping=True)
94
95 adata.obsm["X_totalVI"] = model.get_latent_representation()
96
97 # Also get denoised protein - convert to numpy array for h5ad compatibility
98 _, protein_denoised = model.get_normalized_expression(return_mean=True)
99 if hasattr(protein_denoised, 'values'):
100 adata.obsm["protein_denoised"] = protein_denoised.values
101 else:
102 adata.obsm["protein_denoised"] = np.array(protein_denoised)
103
104 return model, "X_totalVI"
105
106
107def train_peakvi(adata, batch_key=None, n_latent=20, max_epochs=200):
108 """Train PeakVI model for scATAC-seq."""
109 import scvi
110 import numpy as np
111
112 # Binarize if not already
113 if adata.X.max() > 1:
114 print("Binarizing ATAC data...")
115 adata.X = (adata.X > 0).astype(np.float32)
116
117 scvi.model.PEAKVI.setup_anndata(
118 adata,
119 batch_key=batch_key
120 )
121
122 model = scvi.model.PEAKVI(
123 adata,
124 n_latent=n_latent
125 )
126
127 model.train(max_epochs=max_epochs, early_stopping=True)
128
129 adata.obsm["X_PeakVI"] = model.get_latent_representation()
130 return model, "X_PeakVI"
131
132
133def train_velovi(adata, max_epochs=500):
134 """Train veloVI model for RNA velocity.
135
136 Note: Requires scvelo preprocessing. If Ms/Mu layers don't exist,
137 will run preprocessing automatically.
138 """
139 import scvi
140 import scvelo as scv
141
142 # Check if data needs preprocessing
143 if "Ms" not in adata.layers or "Mu" not in adata.layers:
144 print("Preprocessing data for veloVI (scvelo moments)...")
145
146 # Filter and normalize
147 scv.pp.filter_and_normalize(adata, min_shared_counts=30, n_top_genes=2000)
148
149 # Calculate moments (creates Ms, Mu layers)
150 scv.pp.moments(adata, n_pcs=30, n_neighbors=30)
151
152 print(f"After preprocessing: {adata.shape}")
153
154 # VELOVI is in scvi.external, not scvi.model
155 scvi.external.VELOVI.setup_anndata(
156 adata,
157 spliced_layer="Ms",
158 unspliced_layer="Mu"
159 )
160
161 model = scvi.external.VELOVI(adata)
162 model.train(max_epochs=max_epochs, early_stopping=True)
163
164 # Get latent representation (cells x latent_dim)
165 adata.obsm["X_veloVI"] = model.get_latent_representation()
166
167 # Get velocity (cells x genes)
168 adata.layers["velocity"] = model.get_velocity()
169
170 # Get latent time per gene (cells x genes) - store mean across genes as summary
171 latent_time_df = model.get_latent_time()
172 adata.obs["latent_time_mean"] = latent_time_df.mean(axis=1).values
173
174 return model, "X_veloVI"
175
176
177def train_multivi(adata, batch_key=None, n_latent=20, max_epochs=300):
178 """Train MultiVI model for multiome (RNA + ATAC).
179
180 Note: Expects MuData or AnnData with both RNA and ATAC data.
181 For AnnData, ATAC peaks should be concatenated with genes,
182 or use MuData format.
183 """
184 import scvi
185 import numpy as np
186
187 # Check if this is MuData
188 try:
189 import mudata as md
190 if isinstance(adata, md.MuData):
191 # Setup for MuData
192 scvi.model.MULTIVI.setup_mudata(
193 adata,
194 rna_layer="counts",
195 atac_layer="counts",
196 batch_key=batch_key,
197 modalities={
198 "rna_layer": "rna",
199 "batch_key": "rna",
200 "atac_layer": "atac"
201 }
202 )
203 else:
204 raise ValueError("MultiVI requires MuData format with 'rna' and 'atac' modalities")
205 except ImportError:
206 raise ImportError("MultiVI requires mudata. Install with: pip install mudata")
207
208 model = scvi.model.MULTIVI(
209 adata,
210 n_latent=n_latent
211 )
212
213 model.train(max_epochs=max_epochs, early_stopping=True)
214
215 # Get latent representation
216 latent = model.get_latent_representation()
217 adata.obsm["X_MultiVI"] = latent
218
219 return model, "X_MultiVI"
220
221
222MODELS = {
223 "scvi": train_scvi,
224 "scanvi": train_scanvi,
225 "totalvi": train_totalvi,
226 "peakvi": train_peakvi,
227 "velovi": train_velovi,
228 "multivi": train_multivi,
229}
230
231
232def main():
233 parser = argparse.ArgumentParser(
234 description="Train scvi-tools models",
235 formatter_class=argparse.RawDescriptionHelpFormatter,
236 epilog="""
237Examples:
238 # Train scVI for batch correction
239 python train_model.py prepared.h5ad results/ --model scvi --batch-key batch
240
241 # Train scANVI with cell type labels
242 python train_model.py prepared.h5ad results/ --model scanvi --batch-key batch --labels-key cell_type
243
244 # Train totalVI for CITE-seq
245 python train_model.py citeseq.h5ad results/ --model totalvi --batch-key batch
246
247 # Train PeakVI for ATAC-seq
248 python train_model.py atac.h5ad results/ --model peakvi
249
250 # Train veloVI for RNA velocity
251 python train_model.py velocity.h5ad results/ --model velovi
252
253 # Train MultiVI for multiome (RNA + ATAC) - requires MuData format
254 python train_model.py multiome.h5mu results/ --model multivi --batch-key batch
255 """
256 )
257 parser.add_argument("input", help="Input h5ad file (prepared)")
258 parser.add_argument("output_dir", help="Output directory for model and results")
259 parser.add_argument("--model", choices=list(MODELS.keys()), default="scvi",
260 help="Model type (default: scvi)")
261 parser.add_argument("--batch-key", help="Batch column in obs")
262 parser.add_argument("--labels-key", help="Labels column (required for scanvi)")
263 parser.add_argument("--protein-key", default="protein_expression",
264 help="Protein obsm key for totalvi")
265 parser.add_argument("--n-latent", type=int, default=30, help="Latent dimensions (default: 30)")
266 parser.add_argument("--n-layers", type=int, default=2, help="Encoder/decoder layers (default: 2)")
267 parser.add_argument("--max-epochs", type=int, default=200, help="Max training epochs (default: 200)")
268
269 args = parser.parse_args()
270
271 # Validate
272 if args.model == "scanvi" and args.labels_key is None:
273 print("Error: --labels-key required for scanvi model")
274 sys.exit(1)
275
276 try:
277 import scvi
278 import scanpy as sc
279 except ImportError:
280 print("Error: scvi-tools and scanpy required")
281 sys.exit(1)
282
283 # Create output directory
284 os.makedirs(args.output_dir, exist_ok=True)
285
286 # Load data
287 print(f"Loading {args.input}...")
288 if args.input.endswith('.h5mu') or args.model == "multivi":
289 try:
290 import mudata as md
291 adata = md.read(args.input)
292 print(f"MuData: {adata.n_obs} cells")
293 for mod_name, mod in adata.mod.items():
294 print(f" {mod_name}: {mod.shape}")
295 except ImportError:
296 print("Error: mudata required for .h5mu files. Install with: pip install mudata")
297 sys.exit(1)
298 else:
299 adata = sc.read_h5ad(args.input)
300 print(f"Data: {adata.shape}")
301
302 # Check for counts layer
303 if "counts" not in adata.layers:
304 print("Warning: 'counts' layer not found, using X")
305 adata.layers["counts"] = adata.X.copy()
306
307 # Train model
308 print(f"\nTraining {args.model.upper()}...")
309
310 if args.model == "scvi":
311 model, rep_key = train_scvi(
312 adata, args.batch_key, args.n_latent, args.n_layers, args.max_epochs
313 )
314 elif args.model == "scanvi":
315 model, rep_key = train_scanvi(
316 adata, args.batch_key, args.labels_key, args.n_latent, args.n_layers, args.max_epochs
317 )
318 elif args.model == "totalvi":
319 model, rep_key = train_totalvi(
320 adata, args.batch_key, args.protein_key, args.n_latent, args.max_epochs
321 )
322 elif args.model == "peakvi":
323 model, rep_key = train_peakvi(
324 adata, args.batch_key, args.n_latent, args.max_epochs
325 )
326 elif args.model == "velovi":
327 model, rep_key = train_velovi(adata, args.max_epochs)
328 elif args.model == "multivi":
329 model, rep_key = train_multivi(adata, args.batch_key, args.n_latent, args.max_epochs)
330
331 print("Training complete!")
332
333 # Save model
334 model_path = os.path.join(args.output_dir, "model")
335 model.save(model_path)
336 print(f"Model saved to {model_path}")
337
338 # Save adata with latent representation
339 adata_path = os.path.join(args.output_dir, "adata_trained.h5ad")
340 adata.write_h5ad(adata_path)
341 print(f"AnnData saved to {adata_path}")
342
343 # Save training history plot
344 try:
345 import matplotlib.pyplot as plt
346
347 fig, ax = plt.subplots(figsize=(8, 4))
348 if "elbo_train" in model.history:
349 ax.plot(model.history["elbo_train"], label="Train")
350 if "elbo_validation" in model.history:
351 ax.plot(model.history["elbo_validation"], label="Validation")
352 ax.set_xlabel("Epoch")
353 ax.set_ylabel("ELBO")
354 ax.legend()
355 ax.set_title(f"{args.model.upper()} Training History")
356
357 plot_path = os.path.join(args.output_dir, "training_history.png")
358 plt.savefig(plot_path, dpi=150, bbox_inches="tight")
359 plt.close()
360 print(f"Training plot saved to {plot_path}")
361 except Exception as e:
362 print(f"Could not save training plot: {e}")
363
364 print("\nDone! Next steps:")
365 print(f" - Run clustering: python cluster_embed.py {adata_path} {args.output_dir}")
366 print(f" - Load model: scvi.model.{args.model.upper()}.load('{model_path}')")
367
368
369if __name__ == "__main__":
370 main()