diff --git a/scdataloader/collator.py b/scdataloader/collator.py index b9d7c41..4342545 100644 --- a/scdataloader/collator.py +++ b/scdataloader/collator.py @@ -8,6 +8,20 @@ from .utils import load_genes +def build_axis_mapping(native_axis: list[str], global_gene_index: dict[str, int]): + """Map one artifact's native gene axis into the model's global gene axis.""" + positions = np.fromiter( + (global_gene_index.get(gene, -1) for gene in native_axis), + dtype=np.int64, + count=len(native_axis), + ) + accepted = positions >= 0 + model_positions = positions[accepted] + if len(model_positions) == 0: + raise ValueError("native var axis has no model genes") + return accepted, model_positions + + class Collator: def __init__( self, @@ -25,6 +39,7 @@ def __init__( class_names: List[str] = [], genelist: List[str] = [], genedf: Optional[pd.DataFrame] = None, + storage_gene_mappings: list[tuple[np.ndarray, np.ndarray]] | None = None, ): """ Collator for preparing gene expression data batches for the scPRINT model. @@ -94,6 +109,7 @@ def __init__( self.organism_name = organism_name self.tp_name = tp_name self.class_names = class_names + self.storage_gene_mappings = storage_gene_mappings self.start_idx = {} self.accepted_genes = {} self.to_subset = {} @@ -156,7 +172,7 @@ def __call__(self, batch) -> dict[str, Tensor]: - organism_name (any): Organism identifier (column name set by `organism_name`). - tp_name (float, optional): Time point value (column name set by `tp_name`). - class_names... (any, optional): Additional class labels. - - "_storage_idx" (int, optional): Dataset storage index. + - "_store_idx" (int, optional): Dataset storage index. - "is_meta" (int, optional): Metadata flag. - "knn_cells" (array, optional): KNN neighbor expression data. - "knn_cells_info" (array, optional): KNN neighbor metadata. @@ -178,7 +194,7 @@ def __call__(self, batch) -> dict[str, Tensor]: - "knn_cells_info" (Tensor, optional): KNN metadata. Present if input contains "knn_cells_info". - "dataset" (Tensor, optional): Dataset indices as int64. Present if input - contains "_storage_idx". + contains "_store_idx". Note: Batch size in output may be smaller than input if some samples are filtered @@ -198,14 +214,30 @@ def __call__(self, batch) -> dict[str, Tensor]: knn_cells = [] knn_cells_info = [] for elem in batch: + storage_mapping = None + storage_gene_locs = None organism_id = elem[self.organism_name] if organism_id not in self.organism_ids: continue - if "_storage_idx" in elem: - dataset.append(elem["_storage_idx"]) + storage_idx = elem.get("_store_idx", elem.get("_storage_idx")) + if storage_idx is not None: + dataset.append(storage_idx) expr = np.array(elem["X"]) total_count.append(expr.sum()) - if len(self.accepted_genes) > 0: + if self.storage_gene_mappings is not None: + if storage_idx is None: + raise KeyError("storage-aware gene mapping requires _store_idx") + storage_mapping = self.storage_gene_mappings[int(storage_idx)] + accepted, storage_gene_locs = storage_mapping + if len(expr) != len(accepted): + raise ValueError( + "expression/native-axis length mismatch for storage " + f"{storage_idx}: {len(expr)} != {len(accepted)}" + ) + expr = expr[accepted] + if "knn_cells" in elem: + elem["knn_cells"] = elem["knn_cells"][:, accepted] + elif len(self.accepted_genes) > 0: expr = expr[self.accepted_genes[organism_id]] if "knn_cells" in elem: elem["knn_cells"] = elem["knn_cells"][ @@ -287,7 +319,10 @@ def __call__(self, batch) -> dict[str, Tensor]: knn_cells_info.append(elem["knn_cells_info"]) # then we need to add the start_idx to the loc to give it the correct index # according to the model - gene_locs.append(loc + self.start_idx[organism_id]) + if storage_mapping is not None and storage_gene_locs is not None: + gene_locs.append(storage_gene_locs[loc]) + else: + gene_locs.append(loc + self.start_idx[organism_id]) if self.tp_name is not None: tp.append(elem[self.tp_name]) diff --git a/scdataloader/datamodule.py b/scdataloader/datamodule.py index 6c83377..1845dde 100644 --- a/scdataloader/datamodule.py +++ b/scdataloader/datamodule.py @@ -1,3 +1,4 @@ +import hashlib import math import multiprocessing as mp import os @@ -12,6 +13,7 @@ import numpy as np import pandas as pd import torch +from lamindb.core.storage._anndata_accessor import _safer_read_index from torch.utils.data import DataLoader, Sampler, Subset from torch.utils.data.sampler import ( RandomSampler, @@ -21,8 +23,9 @@ ) from tqdm import tqdm -from .collator import Collator +from .collator import Collator, build_axis_mapping from .data import Dataset +from .mapped import _Connect from .utils import fileToList, getBiomartTable, listToFile FILE_DIR = os.path.dirname(os.path.abspath(__file__)) @@ -38,6 +41,7 @@ def __init__( n_samples_per_epoch: int = 2_000_000, validation_split: float = 0.2, test_split: float = 0, + gene_embeddings: str = "", use_default_col: bool = True, # this is for the mappedCollection clss_to_predict: List[str] = ["organism_ontology_term_id"], @@ -94,6 +98,9 @@ def __init__( test_split (float | int, optional): Proportion (float) or absolute number (int) of samples for testing. Uses entire datasets as test sets, rounding to nearest dataset boundary. Defaults to 0. + gene_embeddings (str, optional): Parquet file whose index defines the exact + model gene axis. This legacy CLI source can be paired to the model's + ``precpt_gene_emb`` argument. Defaults to an empty string. use_default_col (bool, optional): Whether to use the default Collator for batch preparation. If False, no collate_fn is applied. Defaults to True. clss_to_predict (List[str], optional): Observation columns to encode as prediction @@ -186,14 +193,64 @@ def __init__( self.metacell_mode = bool(metacell_mode) self.gene_pos = None self.collection_name = collection_name - if gene_subset is not None: - tokeep = set(mdataset.genedf.index.tolist()) - gene_subset = [u for u in gene_subset if u in tokeep] + gene_frame = mdataset.genedf + if gene_frame is None: + raise ValueError("dataset gene registry is unavailable") + if gene_embeddings: + embedding_genes = ( + pd.read_parquet(gene_embeddings, columns=[]).index.astype(str).tolist() + ) + available = set(gene_frame.index.astype(str)) + gene_subset = [gene for gene in embedding_genes if gene in available] + elif gene_subset is not None: + available = set(gene_frame.index.astype(str)) + gene_subset = [str(gene) for gene in gene_subset if str(gene) in available] self.classes = {k: len(v) for k, v in mdataset.class_topred.items()} # we might want not to order the genes by expression (or do it?) # we might want to not introduce zeros and if use_default_col: + model_genes = gene_subset or gene_frame.index.astype(str).tolist() + global_gene_index = {gene: index for index, gene in enumerate(model_genes)} + if len(global_gene_index) != len(model_genes): + raise ValueError("global model gene ids are not unique") + mapping_cache = {} + storage_gene_mappings = [] + overlap_sizes = [] + for storage_index, storage in enumerate( + mdataset.mapped_dataset.storages, start=1 + ): + with _Connect(storage) as store: + native_axis = _safer_read_index(store["var"]).astype(str).tolist() + digest = hashlib.sha256("\0".join(native_axis).encode()).digest() + key = (len(native_axis), digest) + mapping = mapping_cache.get(key) + if mapping is None: + mapping = build_axis_mapping(native_axis, global_gene_index) + if len(mapping[1]) < max_len: + raise ValueError( + "artifact maps fewer genes than max_len: " + f"storage={storage_index - 1}, mapped={len(mapping[1])}, " + f"max_len={max_len}" + ) + mapping_cache[key] = mapping + storage_gene_mappings.append(mapping) + overlap_sizes.append(len(mapping[1])) + if storage_index % 1000 == 0: + print( + "STORAGE_GENE_MAP_PROGRESS " + f"{storage_index}/{len(mdataset.mapped_dataset.storages)} " + f"unique_axes={len(mapping_cache)}", + flush=True, + ) + print( + "STORAGE_GENE_MAP_PASS " + f"storages={len(storage_gene_mappings)} " + f"unique_axes={len(mapping_cache)} " + f"min_overlap={min(overlap_sizes)} " + f"max_overlap={max(overlap_sizes)}", + flush=True, + ) kwargs["collate_fn"] = Collator( organisms=mdataset.organisms if organisms is None else organisms, how=how, @@ -205,6 +262,7 @@ def __init__( class_names=list(self.classes.keys()), genedf=genedf, n_bins=n_bins, + storage_gene_mappings=storage_gene_mappings, ) self.gene_subset = gene_subset self.n_bins = n_bins @@ -305,9 +363,21 @@ def genes(self) -> list: @property def genes_dict(self): + if self.gene_subset is None: + return { + organism: self.dataset.genedf.index[ + self.dataset.genedf.organism == organism + ].tolist() + for organism in self.dataset.organisms + } + organism_by_gene = self.dataset.genedf["organism"].to_dict() return { - i: self.dataset.genedf.index[self.dataset.genedf.organism == i].tolist() - for i in self.dataset.organisms + organism: [ + gene + for gene in self.gene_subset + if organism_by_gene.get(gene) == organism + ] + for organism in self.dataset.organisms } def set_valid_genes_collator(self, genes): @@ -701,6 +771,7 @@ def __init__( super(LabelWeightedSampler, self).__init__(None) self.count = 0 self.curiculum = curiculum + self.restrict_to_subset = restrict_to_subset # Compute label weights (incorporating class frequencies) # Directly use labels as numpy array without conversion diff --git a/tests/test_gene_axis_mapping.py b/tests/test_gene_axis_mapping.py new file mode 100644 index 0000000..9541d5d --- /dev/null +++ b/tests/test_gene_axis_mapping.py @@ -0,0 +1,92 @@ +from __future__ import annotations + +import inspect +from types import SimpleNamespace + +import numpy as np +import pandas as pd + +from scdataloader import DataModule +from scdataloader.collator import Collator, build_axis_mapping +from scdataloader.datamodule import LabelWeightedSampler + + +def test_gene_embeddings_is_an_explicit_datamodule_cli_parameter(): + assert "gene_embeddings" in inspect.signature(DataModule.__init__).parameters + + +def test_build_axis_mapping_preserves_model_gene_order(): + accepted, model_positions = build_axis_mapping( + ["g2", "missing", "g0"], {"g0": 0, "g1": 1, "g2": 2} + ) + + assert accepted.tolist() == [True, False, True] + assert model_positions.tolist() == [2, 0] + + +def test_storage_mapping_reorders_expression_and_emits_storage_index(): + gene_frame = pd.DataFrame( + {"organism": ["org", "org", "org"]}, index=["g0", "g1", "g2"] + ) + mapping = build_axis_mapping( + ["g2", "missing", "g0"], {"g0": 0, "g1": 1, "g2": 2} + ) + collator = Collator( + organisms=["org"], + how="all", + org_to_id={"org": 0}, + valid_genes=["g0", "g1", "g2"], + max_len=3, + organism_name="organism_ontology_term_id", + genedf=gene_frame, + storage_gene_mappings=[mapping], + ) + # Collator's existing genedf path does not initialize this derived set. + collator.organism_ids = {0} + + result = collator( + [ + { + "X": np.array([20.0, 99.0, 10.0]), + "organism_ontology_term_id": 0, + "_store_idx": 0, + } + ] + ) + + assert result["x"].tolist() == [[20.0, 10.0]] + assert result["genes"].tolist() == [[2, 0]] + assert result["dataset"].tolist() == [0] + assert result["depth"].tolist() == [129.0] + + +def test_genes_dict_uses_embedding_subset_order(): + datamodule = object.__new__(DataModule) + datamodule.gene_subset = ["g2", "g0", "g3"] + datamodule.dataset = SimpleNamespace( + organisms=["org_a", "org_b"], + genedf=pd.DataFrame( + {"organism": ["org_a", "org_a", "org_b", "org_b"]}, + index=["g0", "g1", "g2", "g3"], + ), + ) + + assert datamodule.genes_dict == { + "org_a": ["g0"], + "org_b": ["g2", "g3"], + } + + +def test_weighted_sampler_initializes_and_applies_subset_restriction(): + sampler = LabelWeightedSampler( + labels=np.array([0, 0, 1]), + num_samples=8, + weight_scaler=10, + element_weights=np.array([0.1, 0.9, 1.0]), + n_workers=1, + chunk_size=10, + restrict_to_subset=1, + ) + + assert sampler.restrict_to_subset == 1 + assert set(iter(sampler)) == {1}