scvi.dataloaders.SemiSupervisedDataSplitter#

class scvi.dataloaders.SemiSupervisedDataSplitter(adata_manager=None, datamodule=None, train_size=None, validation_size=None, shuffle_set_split=True, n_samples_per_label=None, pin_memory=False, external_indexing=None, **kwargs)[source]#

Bases: LightningDataModule

Creates data loaders train_set, validation_set, test_set.

If train_size + validation_set < 1 ``, then ``test_set is non-empty. The ratio between labeled and unlabeled data in adata will be preserved in the train/test/val sets.

Parameters:
  • adata_manager (AnnDataManager | None (default: None)) – AnnDataManager object that has been created via setup_anndata.

  • train_size (float | None (default: None)) – float, or None (default is None, which is practically 0.9 and potentially adding a small last batch to validation cells)

  • validation_size (float | None (default: None)) – float, or None (default is None)

  • shuffle_set_split (bool (default: True)) – Whether to shuffle indices before splitting. If False, the val, train, and test set are split in the sequential order of the data according to validation_size and train_size percentages.

  • n_samples_per_label (int | None (default: None)) – Number of subsamples for each label class to sample per epoch

  • pin_memory (bool (default: False)) – Whether to copy tensors into device-pinned memory before returning them. Passed into AnnDataLoader.

  • external_indexing (list[ndarray] | None (default: None)) – A list of data split indices in the order of training, validation, and test sets. Validation and test set are not required and can be left empty. Note that per group (train,valid,test) it will cover both the labeled and unlebeled parts

  • **kwargs – Keyword args for data loader. If adata has labeled data, the data loader class is SemiSupervisedDataLoader, else the data loader class is AnnDataLoader.

Examples

>>> adata = scvi.data.synthetic_iid()
>>> scvi.model.SCVI.setup_anndata(adata, labels_key="labels")
>>> adata_manager = scvi.model.SCVI(adata).adata_manager
>>> unknown_label = "label_0"
>>> splitter = SemiSupervisedDataSplitter(adata, unknown_label)
>>> splitter.setup()
>>> train_dl = splitter.train_dataloader()

Attributes table#

Methods table#

setup([stage])

Split indices in train/test/val sets.

test_dataloader()

Create the test data loader.

train_dataloader()

Create the train data loader.

val_dataloader()

Create the validation data loader.

Attributes#

Methods#

SemiSupervisedDataSplitter.setup(stage=None)[source]#

Split indices in train/test/val sets.

SemiSupervisedDataSplitter.test_dataloader()[source]#

Create the test data loader.

SemiSupervisedDataSplitter.train_dataloader()[source]#

Create the train data loader.

SemiSupervisedDataSplitter.val_dataloader()[source]#

Create the validation data loader.