Training#
Training utilities and callbacks for spatial models.
Training Plans#
Training plans for spatial models.
- class spatialvi.train._training_plans.SpatialTrainingPlan(module, lr=0.001, weight_decay=1e-06, spatial_weight=1.0, spatial_warmup_epochs=10, **kwargs)[source]#
Bases:
TrainingPlanTraining plan with spatial-specific features.
This training plan extends the base TrainingPlan with: - Spatial loss weighting - Neighbor-aware loss computation - Multi-scale training
- Parameters:
module¶ (
BaseModuleClass) – Neural network module.spatial_weight¶ (
float) – Weight for spatial regularization term.spatial_warmup_epochs¶ (
int) – Number of epochs to warmup spatial regularization.**kwargs¶ – Additional keyword arguments for TrainingPlan.
module (BaseModuleClass)
lr (float)
weight_decay (float)
spatial_weight (float)
spatial_warmup_epochs (int)
- prepare_data_per_node: bool#
- allow_zero_length_dataloader_with_multiple_devices: bool#
- training: bool#
- class spatialvi.train._training_plans.NicheTrainingPlan(module, lr=0.001, classification_weight=1.0, niche_weight=1.0, **kwargs)[source]#
Bases:
SpatialTrainingPlanTraining plan for niche-aware models.
Extends SpatialTrainingPlan with: - Niche composition loss - Classification loss for semi-supervised training - Balanced sampling for cell types
- Parameters:
- prepare_data_per_node: bool#
- allow_zero_length_dataloader_with_multiple_devices: bool#
- training: bool#
- class spatialvi.train._training_plans.DeconvolutionTrainingPlan(module, lr=0.001, sparsity_weight=0.1, reference_weight=0.0, **kwargs)[source]#
Bases:
SpatialTrainingPlanTraining plan for spatial deconvolution models.
Extends SpatialTrainingPlan with: - Proportion sparsity regularization - Reference-guided constraints - Multi-stage training
- Parameters:
- prepare_data_per_node: bool#
- allow_zero_length_dataloader_with_multiple_devices: bool#
- training: bool#
Callbacks#
Training callbacks for spatial models.
- class spatialvi.train._callbacks.SpatialMetricsCallback(compute_every_n_epochs=10, metrics=None)[source]#
Bases:
CallbackCallback for computing spatial metrics during training.
- Parameters:
- class spatialvi.train._callbacks.NeighborSamplingCallback(neighbor_key='nn_index', include_neighbor_expr=True)[source]#
Bases:
CallbackCallback for neighbor-aware batch sampling.
This callback modifies the data loader to include neighbor information in each batch.
- Parameters:
- class spatialvi.train._callbacks.EarlyStoppingOnSpatialLoss(patience=20, min_delta=0.0001, monitor='spatial_loss')[source]#
Bases:
CallbackEarly stopping based on spatial loss convergence.
- Parameters: