From 2f9f839369730b5b9a4f5d16a70c047dab302136 Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Wed, 20 Aug 2025 14:22:45 +0000 Subject: [PATCH 1/8] Validation dataloader so that `train_size` gets used --- src/tiledbsoma_ml/scvi.py | 54 +++++++++++++++++++++++++++++++++++---- 1 file changed, 49 insertions(+), 5 deletions(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index e713779..5496a6e 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -12,6 +12,7 @@ from tiledbsoma_ml import ExperimentDataset, experiment_dataloader from tiledbsoma_ml._common import MiniBatch +from tiledbsoma_ml._query_ids import QueryIDs DEFAULT_DATALOADER_KWARGS: dict[str, Any] = { "pin_memory": torch.cuda.is_available(), @@ -38,6 +39,7 @@ 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, **kwargs: Any, ): """Args: @@ -63,6 +65,10 @@ 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. """ super().__init__() self.query = query @@ -93,21 +99,59 @@ 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.train_query_ids = None + self.val_query_ids = None def setup(self, stage: str | None = None) -> None: - # Instantiate the ExperimentDataset with the provided args and kwargs. - self.train_dataset = ExperimentDataset( + # Create QueryIDs from the query + query_ids = QueryIDs.create(self.query) + + # Split data into train and validation sets if train_size < 1.0 + if self.train_size < 1.0: + # Use QueryIDs.split() for efficient splitting + val_size = 1.0 - self.train_size + self.train_query_ids, self.val_query_ids = query_ids.split( + self.train_size, val_size, seed=42 + ) + else: + # Use all data for training + self.train_query_ids = query_ids + self.val_query_ids = None + + def train_dataloader(self) -> DataLoader: + assert self.train_query_ids is not None, "setup() must be called before train_dataloader()" + + # Create dataset with train query_ids + train_dataset = ExperimentDataset( self.query, *self.dataset_args, obs_column_names=self.batch_column_names, # type: ignore[arg-type] + query_ids=self.train_query_ids, **self.dataset_kwargs, # type: ignore[misc] ) - - def train_dataloader(self) -> DataLoader: return experiment_dataloader( - self.train_dataset, + train_dataset, **self.dataloader_kwargs, ) + + def val_dataloader(self) -> DataLoader | None: + if self.val_query_ids is not None: + # Create dataset with validation query_ids + val_dataset = ExperimentDataset( + self.query, + *self.dataset_args, + obs_column_names=self.batch_column_names, # type: ignore[arg-type] + query_ids=self.val_query_ids, + **self.dataset_kwargs, # type: ignore[misc] + ) + return experiment_dataloader( + val_dataset, + **self.dataloader_kwargs, + ) + else: + # No validation data if train_size == 1.0 + return None def _add_batch_col( self, obs_df: pd.DataFrame, inplace: bool = False From 98675c3fa348518587f2c8541d6279a2a532c452 Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Wed, 20 Aug 2025 14:34:58 +0000 Subject: [PATCH 2/8] use XLocator --- src/tiledbsoma_ml/scvi.py | 29 +++++++++++++++++++---------- 1 file changed, 19 insertions(+), 10 deletions(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index 5496a6e..f24c279 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -13,6 +13,7 @@ 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(), @@ -102,10 +103,17 @@ def __init__( self.train_size = train_size self.train_query_ids = None self.val_query_ids = None + self.x_locator = None + self.layer_name = kwargs.get('layer_name', 'raw') def setup(self, stage: str | None = None) -> None: - # Create QueryIDs from the query + # 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, + ) # Split data into train and validation sets if train_size < 1.0 if self.train_size < 1.0: @@ -121,13 +129,14 @@ def setup(self, stage: str | None = None) -> None: def train_dataloader(self) -> DataLoader: assert self.train_query_ids is not None, "setup() must be called before train_dataloader()" + assert self.x_locator is not None, "setup() must be called before train_dataloader()" - # Create dataset with train query_ids + # Create dataset with train query_ids and x_locator train_dataset = ExperimentDataset( - self.query, - *self.dataset_args, - obs_column_names=self.batch_column_names, # type: ignore[arg-type] + x_locator=self.x_locator, query_ids=self.train_query_ids, + obs_column_names=self.batch_column_names, # type: ignore[arg-type] + *self.dataset_args, **self.dataset_kwargs, # type: ignore[misc] ) return experiment_dataloader( @@ -136,13 +145,13 @@ def train_dataloader(self) -> DataLoader: ) def val_dataloader(self) -> DataLoader | None: - if self.val_query_ids is not None: - # Create dataset with validation query_ids + if self.val_query_ids is not None and self.x_locator is not None: + # Create dataset with validation query_ids and x_locator val_dataset = ExperimentDataset( - self.query, - *self.dataset_args, - obs_column_names=self.batch_column_names, # type: ignore[arg-type] + x_locator=self.x_locator, query_ids=self.val_query_ids, + obs_column_names=self.batch_column_names, # type: ignore[arg-type] + *self.dataset_args, **self.dataset_kwargs, # type: ignore[misc] ) return experiment_dataloader( From 22a8bb3baa089cb51eb277062ea52077d766354b Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Wed, 20 Aug 2025 14:40:19 +0000 Subject: [PATCH 3/8] filter query and layer_name --- src/tiledbsoma_ml/scvi.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index f24c279..233cfa7 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -131,13 +131,17 @@ def train_dataloader(self) -> DataLoader: assert self.train_query_ids is not None, "setup() must be called before train_dataloader()" assert self.x_locator is not None, "setup() must be called before train_dataloader()" + # 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 train query_ids and x_locator train_dataset = ExperimentDataset( x_locator=self.x_locator, query_ids=self.train_query_ids, obs_column_names=self.batch_column_names, # type: ignore[arg-type] *self.dataset_args, - **self.dataset_kwargs, # type: ignore[misc] + **filtered_kwargs, # type: ignore[misc] ) return experiment_dataloader( train_dataset, @@ -146,13 +150,17 @@ def train_dataloader(self) -> DataLoader: def val_dataloader(self) -> DataLoader | None: if self.val_query_ids is not None and self.x_locator is not 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 validation query_ids and x_locator val_dataset = ExperimentDataset( x_locator=self.x_locator, query_ids=self.val_query_ids, obs_column_names=self.batch_column_names, # type: ignore[arg-type] *self.dataset_args, - **self.dataset_kwargs, # type: ignore[misc] + **filtered_kwargs, # type: ignore[misc] ) return experiment_dataloader( val_dataset, From abe24369db0e093767669f6cdf53abea529a8095 Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Wed, 20 Aug 2025 15:18:44 +0000 Subject: [PATCH 4/8] use random_split --- src/tiledbsoma_ml/scvi.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index 233cfa7..823ca65 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -117,9 +117,9 @@ def setup(self, stage: str | None = None) -> None: # Split data into train and validation sets if train_size < 1.0 if self.train_size < 1.0: - # Use QueryIDs.split() for efficient splitting + # Use QueryIDs.random_split() for efficient splitting val_size = 1.0 - self.train_size - self.train_query_ids, self.val_query_ids = query_ids.split( + self.train_query_ids, self.val_query_ids = query_ids.random_split( self.train_size, val_size, seed=42 ) else: @@ -167,7 +167,6 @@ def val_dataloader(self) -> DataLoader | None: **self.dataloader_kwargs, ) else: - # No validation data if train_size == 1.0 return None def _add_batch_col( From 36a442722f8aceb9302e28f7c2a2cc2b17c46e02 Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Wed, 20 Aug 2025 15:53:17 +0000 Subject: [PATCH 5/8] set seed for the validation indices --- src/tiledbsoma_ml/scvi.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index 823ca65..66a9de4 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -3,6 +3,7 @@ import os from typing import Any, Sequence +import numpy as np import pandas as pd import torch from lightning import LightningDataModule @@ -41,6 +42,7 @@ def __init__( batch_labels: Sequence[str] | None = None, dataloader_kwargs: dict[str, Any] | None = None, train_size: float = 1.0, + seed: int = 42, **kwargs: Any, ): """Args: @@ -70,6 +72,9 @@ def __init__( 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 @@ -101,6 +106,7 @@ def __init__( 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 = None self.val_query_ids = None self.x_locator = None @@ -120,7 +126,7 @@ def setup(self, stage: str | None = None) -> None: # Use QueryIDs.random_split() for efficient splitting val_size = 1.0 - self.train_size self.train_query_ids, self.val_query_ids = query_ids.random_split( - self.train_size, val_size, seed=42 + self.train_size, val_size, seed=self.seed ) else: # Use all data for training @@ -150,6 +156,10 @@ def train_dataloader(self) -> DataLoader: def val_dataloader(self) -> DataLoader | None: if self.val_query_ids is not None and self.x_locator is not None: + # Print validation indices for manual verification + val_ids = self.val_query_ids.obs_joinids + print(f"📊 Validation indices (seed={self.seed}): first 10: {val_ids[:10]}, last 10: {val_ids[-10:]}") + # 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')} From 6b1d76c22f7b4fde11318e3fa6fb320457974358 Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Wed, 20 Aug 2025 16:14:46 +0000 Subject: [PATCH 6/8] seed works; no need to print validation indices anymore --- src/tiledbsoma_ml/scvi.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index 66a9de4..2562aa1 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -156,10 +156,6 @@ def train_dataloader(self) -> DataLoader: def val_dataloader(self) -> DataLoader | None: if self.val_query_ids is not None and self.x_locator is not None: - # Print validation indices for manual verification - val_ids = self.val_query_ids.obs_joinids - print(f"📊 Validation indices (seed={self.seed}): first 10: {val_ids[:10]}, last 10: {val_ids[-10:]}") - # 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')} From 3f7c390a1bec5c5f0d3fbe47db8d77eee51906a2 Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Fri, 22 Aug 2025 18:02:48 +0000 Subject: [PATCH 7/8] fix linter issues --- src/tiledbsoma_ml/scvi.py | 83 ++++++++++++++++++++++----------------- 1 file changed, 46 insertions(+), 37 deletions(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index 2562aa1..5c403fa 100644 --- a/src/tiledbsoma_ml/scvi.py +++ b/src/tiledbsoma_ml/scvi.py @@ -3,7 +3,6 @@ import os from typing import Any, Sequence -import numpy as np import pandas as pd import torch from lightning import LightningDataModule @@ -68,11 +67,11 @@ 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. """ @@ -107,10 +106,10 @@ def __init__( self.batch_encoder = LabelEncoder().fit(self.batch_labels) self.train_size = train_size self.seed = seed - self.train_query_ids = None - self.val_query_ids = None - self.x_locator = None - self.layer_name = kwargs.get('layer_name', 'raw') + 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: # Create QueryIDs and XLocator from the query @@ -120,61 +119,71 @@ def setup(self, stage: str | None = None) -> None: measurement_name=self.query.measurement_name, layer_name=self.layer_name, ) - + # 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 - self.train_query_ids, self.val_query_ids = query_ids.random_split( + 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 train_dataloader(self) -> DataLoader: - assert self.train_query_ids is not None, "setup() must be called before train_dataloader()" - assert self.x_locator is not None, "setup() must be called before train_dataloader()" - + assert ( + self.train_query_ids is not None + ), "setup() must be called before train_dataloader()" + assert ( + self.x_locator is not None + ), "setup() must be called before train_dataloader()" + # 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')} - + filtered_kwargs = { + k: v + for k, v in self.dataset_kwargs.items() + if k not in ("query", "layer_name") + } + # Create dataset with train query_ids and x_locator train_dataset = ExperimentDataset( x_locator=self.x_locator, query_ids=self.train_query_ids, - obs_column_names=self.batch_column_names, # type: ignore[arg-type] - *self.dataset_args, - **filtered_kwargs, # type: ignore[misc] + obs_column_names=list(self.batch_column_names), + **filtered_kwargs, ) return experiment_dataloader( train_dataset, **self.dataloader_kwargs, ) - + def val_dataloader(self) -> DataLoader | None: - if self.val_query_ids is not None and self.x_locator is not 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 validation query_ids and x_locator - val_dataset = ExperimentDataset( - x_locator=self.x_locator, - query_ids=self.val_query_ids, - obs_column_names=self.batch_column_names, # type: ignore[arg-type] - *self.dataset_args, - **filtered_kwargs, # type: ignore[misc] - ) - return experiment_dataloader( - val_dataset, - **self.dataloader_kwargs, - ) - else: + if self.val_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 validation query_ids and x_locator + val_dataset = ExperimentDataset( + x_locator=self.x_locator, + query_ids=self.val_query_ids, + obs_column_names=list(self.batch_column_names), + **filtered_kwargs, + ) + return experiment_dataloader( + val_dataset, + **self.dataloader_kwargs, + ) + def _add_batch_col( self, obs_df: pd.DataFrame, inplace: bool = False ) -> pd.DataFrame: From 832babdaefe084b344f9f4d7cd4461f62b013059 Mon Sep 17 00:00:00 2001 From: maarten-devries Date: Fri, 22 Aug 2025 18:06:46 +0000 Subject: [PATCH 8/8] abstract dataloader function shared by train and val --- src/tiledbsoma_ml/scvi.py | 77 ++++++++++++++++++++++++--------------- 1 file changed, 47 insertions(+), 30 deletions(-) diff --git a/src/tiledbsoma_ml/scvi.py b/src/tiledbsoma_ml/scvi.py index 5c403fa..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 @@ -22,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. @@ -134,35 +142,23 @@ def setup(self, stage: str | None = None) -> None: self.train_query_ids = query_ids self.val_query_ids = None - def train_dataloader(self) -> DataLoader: - assert ( - self.train_query_ids is not None - ), "setup() must be called before train_dataloader()" - assert ( - self.x_locator is not None - ), "setup() must be called before train_dataloader()" + def _create_dataloader(self, split: DatasetSplit) -> DataLoader | None: + """Create a dataloader for the specified dataset split. - # 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") - } + Args: + split: The dataset split (TRAIN or VAL) - # Create dataset with train query_ids and x_locator - train_dataset = ExperimentDataset( - x_locator=self.x_locator, - query_ids=self.train_query_ids, - obs_column_names=list(self.batch_column_names), - **filtered_kwargs, - ) - return experiment_dataloader( - train_dataset, - **self.dataloader_kwargs, - ) + 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, + } - def val_dataloader(self) -> DataLoader | None: - if self.val_query_ids is None or self.x_locator is None: + 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 @@ -172,18 +168,39 @@ def val_dataloader(self) -> DataLoader | None: if k not in ("query", "layer_name") } - # Create dataset with validation query_ids and x_locator - val_dataset = ExperimentDataset( + # Create dataset with appropriate query_ids + dataset = ExperimentDataset( x_locator=self.x_locator, - query_ids=self.val_query_ids, + query_ids=query_ids, obs_column_names=list(self.batch_column_names), **filtered_kwargs, ) return experiment_dataloader( - val_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: