Skip to content
Draft
Show file tree
Hide file tree
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
47 changes: 41 additions & 6 deletions scdataloader/collator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.
Expand Down Expand Up @@ -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 = {}
Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -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"][
Expand Down Expand Up @@ -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])
Expand Down
83 changes: 77 additions & 6 deletions scdataloader/datamodule.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import hashlib
import math
import multiprocessing as mp
import os
Expand All @@ -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,
Expand All @@ -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__))
Expand All @@ -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"],
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down
92 changes: 92 additions & 0 deletions tests/test_gene_axis_mapping.py
Original file line number Diff line number Diff line change
@@ -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}
Loading