Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 100 additions & 8 deletions src/tiledbsoma_ml/scvi.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

import os
from enum import Enum
from typing import Any, Sequence

import pandas as pd
Expand All @@ -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(),
Expand All @@ -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.

Expand All @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down