diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index e713779..1a9c20a 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -1,6 +1,7 @@ from __future__ import annotations import os +from enum import Enum from typing import Any, Sequence import pandas as pd @@ -12,6 +13,8 @@ from tiledbsoma_ml import ExperimentDataset, experiment_dataloader from tiledbsoma_ml._common import MiniBatch +from tiledbsoma_ml._query_ids import QueryIDs +from tiledbsoma_ml.x_locator import XLocator DEFAULT_DATALOADER_KWARGS: dict[str, Any] = { "pin_memory": torch.cuda.is_available(), @@ -20,6 +23,13 @@ } +class DatasetSplit(Enum): + """Enum for dataset splits.""" + + TRAIN = "train" + VAL = "val" + + class SCVIDataModule(LightningDataModule): # type: ignore[misc] """PyTorch Lightning DataModule for training scVI models from SOMA data. @@ -38,6 +48,8 @@ def __init__( batch_column_names: Sequence[str] | None = None, batch_labels: Sequence[str] | None = None, dataloader_kwargs: dict[str, Any] | None = None, + train_size: float = 1.0, + seed: int = 42, **kwargs: Any, ): """Args: @@ -63,6 +75,13 @@ def __init__( dataloader_kwargs: dict, optional Keyword arguments passed to `tiledbsoma_ml.experiment_dataloader()`, e.g. `num_workers`. + + train_size: float, optional + Fraction of data to use for training (between 0 and 1). Default is 1.0 (use all data for training). + If less than 1.0, the remaining data will be used for validation. + + seed: int, optional + Random seed for deterministic train/validation split. Default is 42. """ super().__init__() self.query = query @@ -93,22 +112,95 @@ def __init__( batch_labels = obs_df[self.batch_colname].unique() self.batch_labels = batch_labels self.batch_encoder = LabelEncoder().fit(self.batch_labels) + self.train_size = train_size + self.seed = seed + self.train_query_ids: QueryIDs | None = None + self.val_query_ids: QueryIDs | None = None + self.x_locator: XLocator | None = None + self.layer_name = kwargs.get("layer_name", "raw") def setup(self, stage: str | None = None) -> None: - # Instantiate the ExperimentDataset with the provided args and kwargs. - self.train_dataset = ExperimentDataset( - self.query, - *self.dataset_args, - obs_column_names=self.batch_column_names, # type: ignore[arg-type] - **self.dataset_kwargs, # type: ignore[misc] + # Create QueryIDs and XLocator from the query + query_ids = QueryIDs.create(self.query) + self.x_locator = XLocator.create( + self.query.experiment, + measurement_name=self.query.measurement_name, + layer_name=self.layer_name, ) - def train_dataloader(self) -> DataLoader: + # Split data into train and validation sets if train_size < 1.0 + if self.train_size < 1.0: + # Use QueryIDs.random_split() for efficient splitting + val_size = 1.0 - self.train_size + train_ids, val_ids = query_ids.random_split( + self.train_size, val_size, seed=self.seed + ) + self.train_query_ids = train_ids + self.val_query_ids = val_ids + else: + # Use all data for training + self.train_query_ids = query_ids + self.val_query_ids = None + + def _create_dataloader(self, split: DatasetSplit) -> DataLoader | None: + """Create a dataloader for the specified dataset split. + + Args: + split: The dataset split (TRAIN or VAL) + + Returns: + DataLoader for the specified split, or None if the split doesn't exist + """ + # Get the appropriate query_ids based on split + query_ids_map = { + DatasetSplit.TRAIN: self.train_query_ids, + DatasetSplit.VAL: self.val_query_ids, + } + + query_ids = query_ids_map.get(split) + if query_ids is None or self.x_locator is None: + return None + + # Filter out query and layer_name from dataset_kwargs since we're using x_locator and query_ids + filtered_kwargs = { + k: v + for k, v in self.dataset_kwargs.items() + if k not in ("query", "layer_name") + } + + # Create dataset with appropriate query_ids + dataset = ExperimentDataset( + x_locator=self.x_locator, + query_ids=query_ids, + obs_column_names=list(self.batch_column_names), + **filtered_kwargs, + ) return experiment_dataloader( - self.train_dataset, + dataset, **self.dataloader_kwargs, ) + def train_dataloader(self) -> DataLoader: + """Create the training dataloader. + + Returns: + DataLoader for training data + + Raises: + AssertionError: If setup() hasn't been called + """ + loader = self._create_dataloader(DatasetSplit.TRAIN) + assert loader is not None, "setup() must be called before train_dataloader()" + return loader + + def val_dataloader(self) -> DataLoader | None: + """Create the validation dataloader. + + Returns: + DataLoader for validation data, or None if no validation split exists + """ + return self._create_dataloader(DatasetSplit.VAL) + def _add_batch_col( self, obs_df: pd.DataFrame, inplace: bool = False ) -> pd.DataFrame: