Skip to content
Merged
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
19 changes: 10 additions & 9 deletions src/cellink/_core/donordata.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import copy as _copy
import logging
from collections.abc import Callable

Expand Down Expand Up @@ -167,11 +168,13 @@ def _match_donors(self, G: AnnData | MuData, C: AnnData | MuData) -> None:
self._G = G

def copy(self) -> DonorData:
if self._G.is_view:
self._G = self._G.copy()
if self._C.is_view:
self._C = self._C.copy()
return self
new = DonorData.__new__(DonorData)
new._var_dims_to_sync = list(self._var_dims_to_sync)
new.donor_id = self.donor_id
new._G = self._G.copy()
new._C = self._C.copy()
new.uns = _copy.deepcopy(self.uns)
return new

def _write_dd(self, f: h5py.File, zarr_path: str | None = None, x_chunks=None):
is_zarr = isinstance(f, zarr.Group)
Expand Down Expand Up @@ -312,10 +315,8 @@ def sel(
C_obs: slice = slice(None),
C_var: slice = slice(None),
):
_G = self.G[G_obs]
_G = _G[:, G_var]
_C = self.C[C_obs]
_C = _C[:, C_var]
_G = self.G[G_obs, G_var]
_C = self.C[C_obs, C_var]

_G = self._sync_var_dims(_G, _C)
return DonorData(G=_G, C=_C, donor_id=self.donor_id, var_dims_to_sync=self._var_dims_to_sync)
Expand Down
7 changes: 7 additions & 0 deletions src/cellink/io/_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import sys

import numpy as np
import pandas as pd
import xarray as xr
from anndata import AnnData
from pandas_plink import write_plink1_bin
Expand Down Expand Up @@ -72,6 +73,12 @@ def to_plink(
if not output_prefix.endswith(".bed"):
output_prefix += ".bed"

categorical_cols = [c for c in (chrom, a0, a1) if isinstance(gdata.var[c].dtype, pd.CategoricalDtype)]
if categorical_cols:
gdata = gdata.copy()
for c in categorical_cols:
gdata.var[c] = gdata.var[c].astype(str)

xarr = xr.DataArray(
gdata.X.astype("float32", copy=False),
dims=("sample", "variant"),
Expand Down
64 changes: 41 additions & 23 deletions src/cellink/io/_pgen.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,16 +29,24 @@ def _is_matrix_elem(elem_name: str) -> bool:
return elem_name.endswith("/X") or elem_name.rsplit("/", 1)[0].endswith("/layers")


def lazy_anndata_zarr_callback(func, elem_name: str, elem, iospec):


def lazy_anndata_zarr_callback(func, elem_name: str, elem, iospec, backend: Literal["dask", "zarr"] = "dask"):
"""``read_dispatched`` callback that reconstructs an AnnData (at any
nesting depth, e.g. as ``G``/``C`` inside a larger DonorData zarr store)
while keeping a dense ``X``/``layers`` entry Dask-backed instead of
while keeping a dense ``X``/``layers`` entry lazily backed instead of
materializing it.

``backend="dask"`` (default, unchanged behavior) wraps the dense array in
a Dask array. ``backend="zarr"`` instead returns the raw Zarr array
directly.
"""
if iospec.encoding_type == "anndata" or elem_name.endswith("/"):
return ad.AnnData(
**{
k: read_dispatched(v, lazy_anndata_zarr_callback)
k: read_dispatched(
v, lambda f, n, e, iospec: lazy_anndata_zarr_callback(f, n, e, iospec, backend=backend)
)
for k, v in dict(elem).items()
if not k.startswith("raw.")
}
Expand All @@ -51,22 +59,23 @@ def lazy_anndata_zarr_callback(func, elem_name: str, elem, iospec):
):
return read_elem(elem)
elif _is_matrix_elem(elem_name) and iospec.encoding_type == "array":
return da.from_zarr(elem)
return elem if backend == "zarr" else da.from_zarr(elem)
else:
return func(elem)


def read_pgen_zarr(store: str | Path) -> ad.AnnData:
def read_pgen_zarr(store: str | Path, backend: Literal["dask", "zarr"] = "dask") -> ad.AnnData:
"""
Lazily read an AnnData Zarr v3 store written by `stream_pgen_to_zarr`.

This function reconstructs an :class:`anndata.AnnData` object from a Zarr
store while keeping the primary data matrix (`X`) backed by Dask arrays.
It is designed for large genotype matrices that cannot be loaded fully
into memory.
store while keeping the primary data matrix (`X`) lazily backed rather
than materializing it. It is designed for large genotype matrices that
cannot be loaded fully into memory.

The reader preserves:
- Dense X stored as a Zarr array (returned as a Dask-backed array)
- Dense X stored as a Zarr array (backed by Dask, or the raw Zarr
array directly; see `backend`)
- Sparse matrices (CSR/CSC)
- DataFrames (obs, var)
- Awkward arrays
Expand All @@ -77,22 +86,28 @@ def read_pgen_zarr(store: str | Path) -> ad.AnnData:
store : str or pathlib.Path
Path to a Zarr directory created by `stream_pgen_to_zarr`
or a compatible AnnData Zarr v3 store.
backend : {"dask", "zarr"}
How to back a dense `X`/`layers` entry. ``"dask"`` (default) wraps it in a Dask array, useful if
you need Dask's own lazy-graph chaining on X. ``"zarr"`` returns the
raw Zarr array directly instead.

Returns
-------
anndata.AnnData
AnnData object with:
- `X` as a Dask-backed array (for dense storage)
- `X` as a Dask-backed or raw-Zarr-backed array (for dense storage),
per `backend`
- `obs` and `var` as pandas DataFrames
- empty container groups (`uns`, `obsm`, `varm`, `layers`, etc.)
if present in the store

Notes
-----
- The returned object is **lazy** when X is dense. Computation is triggered
only when `.compute()` or in-memory materialization is requested.
only when `.compute()` (Dask backend) or ordinary indexing (Zarr
backend) is requested.
- For sparse X written via `stream_pgen_to_zarr(..., sparse=True)`,
the matrix is loaded as a SciPy sparse matrix.
the matrix is loaded as a SciPy sparse matrix regardless of `backend`.
- This function relies on AnnData's experimental dispatched I/O API.

Examples
Expand All @@ -101,12 +116,16 @@ def read_pgen_zarr(store: str | Path) -> ad.AnnData:
>>> adata = cellink.io.read_pgen_zarr("genotypes.zarr")
>>> adata
AnnData object with n_obs x n_vars = ...
>>> fast = cellink.io.read_pgen_zarr("genotypes.zarr", backend="zarr")
>>> fast.X[donor_idx, :] # direct Zarr read, no Dask task overhead

>>> # Trigger computation
>>> X = adata.X.compute()
"""
f = zarr.open(str(store), mode="r")
return read_dispatched(f, callback=lazy_anndata_zarr_callback)
return read_dispatched(
f, callback=lambda func, name, elem, iospec: lazy_anndata_zarr_callback(func, name, elem, iospec, backend=backend)
)


def _read_pvar(pvar_file: Path) -> pd.DataFrame:
Expand Down Expand Up @@ -141,12 +160,15 @@ def _read_pvar(pvar_file: Path) -> pd.DataFrame:
pv = pv.rename(columns={k: v for k, v in rename_map.items() if k in pv.columns})
if VAnn.index in pv.columns:
pv[VAnn.index] = pv[VAnn.index].astype(str)

if VAnn.chrom in pv.columns:
pv[VAnn.chrom] = pv[VAnn.chrom].astype(str)
pv[VAnn.chrom] = pv[VAnn.chrom].astype(str).astype("category")
if "FILTER" in pv.columns:
pv["FILTER"] = pv["FILTER"].astype(str).astype("category")
if VAnn.a0 in pv.columns:
pv[VAnn.a0] = pv[VAnn.a0].astype(str)
pv[VAnn.a0] = pv[VAnn.a0].astype(str).astype("category")
if VAnn.a1 in pv.columns:
pv[VAnn.a1] = pv[VAnn.a1].astype(str)
pv[VAnn.a1] = pv[VAnn.a1].astype(str).astype("category")
return pv


Expand All @@ -159,7 +181,7 @@ def stream_pgen_to_zarr(
chunk_samples: int = 4096,
chunk_variants: int = 2048,
memory_limit_gb: float = 10.0,
compressor: str = "zstd",
compressor: str | None = "zstd",
compression_level: int = 7,
sparse: bool = False,
sparse_format: Literal["csc", "csr"] = "csc",
Expand Down Expand Up @@ -284,11 +306,7 @@ def _base(p: str) -> str:
pvar.index = pvar.index.astype(str)

output_path = Path(output_path)
blosc = BloscCodec(
cname=compressor,
clevel=compression_level,
shuffle=BloscShuffle.bitshuffle,
)
codecs = () if compressor is None else (BloscCodec(cname=compressor, clevel=compression_level, shuffle=BloscShuffle.bitshuffle),)

if sparse:
if sparse_format not in ("csr", "csc"):
Expand Down Expand Up @@ -356,7 +374,7 @@ def _base(p: str) -> str:
shape=(n_samples, n_variants_total),
chunks=(chunk_samples, chunk_variants),
dtype="i1",
compressors=(blosc,),
compressors=codecs,
)
Xz.attrs["encoding-type"] = "array"
Xz.attrs["encoding-version"] = "0.2.0"
Expand Down
2 changes: 1 addition & 1 deletion src/cellink/tl/external/_ld.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ def calculate_ld(
plink_export_kwargs = {}

if run and shutil.which("plink") is None:
raise ImportError("plink is required for `calculate_pcs`. Please install it.")
raise ImportError("plink is required for `calculate_ld`. Please install it.")

if out is None:
out = f"{prefix}_ld"
Expand Down
6 changes: 1 addition & 5 deletions src/cellink/tl/external/_seismic_torch.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,11 +189,7 @@ def forward(self, G: torch.Tensor, verbose: bool = False, return_all: bool = Fal
if return_all:
pval_two_sided = torch.tensor(st.chi2(1).sf(self.lrt.cpu().data.numpy()), device=self.F.device)
pval_one_sided = torch.where(self.beta_g > 0, pval_two_sided / 2.0, 1.0 - (pval_two_sided / 2.0))
z = np.sign(self.beta_g.cpu().data.numpy()) * np.sqrt(
st.chi2.ppf(1.0 - pval_two_sided.cpu().data.numpy(), df=1)
)
z = torch.tensor(z, device=self.F.device)
ste = self.beta_g / z
ste = torch.sqrt(self.s2 * n[:, None])
return nll, pval_one_sided, self.beta_g, ste

return nll
Expand Down
2 changes: 1 addition & 1 deletion src/cellink/tl/external/_sldsc_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -549,7 +549,7 @@ def _pick_var_col(adata: AnnData, candidates: list[str], default: str | None) ->

def _normalize_chromosome(chr_series: pd.Series) -> pd.Series:
"""Normalize chromosome labels to standard format."""
normalized = chr_series.astype(str).str.replace("^chr", "", regex=True).str.upper()
normalized = chr_series.astype(str).str.replace("^chr", "", regex=True, case=False).str.upper()
return normalized.str.extract(r"^([0-9XYM]+)", expand=False)


Expand Down
122 changes: 122 additions & 0 deletions tests/test_categorical_var_dtype.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
import numpy as np
import pandas as pd
import pytest

from cellink._core.dummy_data import sim_gdata
from cellink.io._export import write_variants_to_vcf
from cellink.tl._subset_region import subset_genomic_region


def _str_and_categorical_gdata():
"""Two AnnDatas with byte-identical underlying data, differing only in
whether chrom/a0/a1 are plain string or category dtype -- sim_gdata()
draws random alleles internally, so calling it twice independently (as
an earlier version of this fixture did) compares two different random
genotypes, not the same data under two dtypes.
"""
gdata_str = sim_gdata(n_donors=20, n_snps=30)
gdata_str.var["chrom"] = gdata_str.var["chrom"].astype(str)
gdata_cat = gdata_str.copy()
gdata_cat.var["chrom"] = gdata_cat.var["chrom"].astype("category")
gdata_cat.var["a0"] = gdata_cat.var["a0"].astype("category")
gdata_cat.var["a1"] = gdata_cat.var["a1"].astype("category")
return gdata_str, gdata_cat


def test_subset_genomic_region_matches_plain_string_dtype():
gdata_str, gdata_cat = _str_and_categorical_gdata()

start, end = int(gdata_str.var["pos"].min()), int(gdata_str.var["pos"].max()) + 1
sub_str = subset_genomic_region(gdata_str, chrom="1", start=start, end=end)
sub_cat = subset_genomic_region(gdata_cat, chrom="1", start=start, end=end)

assert sub_str.shape == sub_cat.shape
assert list(sub_str.var.index) == list(sub_cat.var.index)


def test_np_unique_and_equality_on_categorical_chrom():
_, gdata = _str_and_categorical_gdata()
uniq = np.unique(gdata.var["chrom"])
assert list(uniq) == ["1"]
mask = gdata.var["chrom"] == "1"
assert mask.all()
assert type(gdata.var["chrom"].iloc[0]) is str


def test_write_variants_to_vcf_identical_with_categorical(tmp_path):
gdata_str, gdata_cat = _str_and_categorical_gdata()

out_str, out_cat = tmp_path / "str.vcf", tmp_path / "cat.vcf"
write_variants_to_vcf(gdata_str, out_file=str(out_str))
write_variants_to_vcf(gdata_cat, out_file=str(out_cat))
assert out_str.read_text() == out_cat.read_text()


def test_to_plink_roundtrip_identical_with_categorical(tmp_path):
bed_reader = pytest.importorskip("bed_reader")
from cellink.io._export import to_plink

gdata_str, gdata_cat = _str_and_categorical_gdata()
gdata_str.obs["donor_id"] = gdata_str.obs.index
gdata_str.obs["sex"] = 0
gdata_cat.obs["donor_id"] = gdata_cat.obs.index
gdata_cat.obs["sex"] = 0

prefix_str, prefix_cat = str(tmp_path / "str"), str(tmp_path / "cat")
to_plink(gdata_str, output_prefix=prefix_str)
to_plink(gdata_cat, output_prefix=prefix_cat)

b_str = bed_reader.open_bed(prefix_str + ".bed")
b_cat = bed_reader.open_bed(prefix_cat + ".bed")
np.testing.assert_array_equal(b_str.read(), b_cat.read())
assert list(b_str.chromosome) == list(b_cat.chromosome)
assert list(b_str.allele_1) == list(b_cat.allele_1)
assert list(b_str.allele_2) == list(b_cat.allele_2)


def test_tensorqtl_input_generator_cis_matches_plain_string_dtype():
"""The real integration risk flagged for this change: cellink's
run_tensorqtl(use_python_api=True) hands variant_df straight to
tensorqtl's own genotypeio.InputGeneratorCis, which does
variant_df['chrom'].unique() / .groupby('chrom') / membership checks --
code cellink does not control. Verify categorical dtype produces
byte-identical cis-window results there directly, not just in cellink's
own call sites.
"""
tensorqtl_genotypeio = pytest.importorskip("tensorqtl.genotypeio")

rng = np.random.default_rng(0)
n_var, n_genes, n_samples = 60, 6, 12
chrom_str = np.array(["1"] * 30 + ["2"] * 30)
pos = np.concatenate([np.sort(rng.choice(1_000_000, 30, replace=False)),
np.sort(rng.choice(1_000_000, 30, replace=False))])
variant_df_str = pd.DataFrame({"chrom": chrom_str, "pos": pos},
index=[f"var{i}" for i in range(n_var)])
variant_df_cat = variant_df_str.copy()
variant_df_cat["chrom"] = variant_df_cat["chrom"].astype("category")

phenotype_pos_df = pd.DataFrame({
"chr": np.array(["1"] * 3 + ["2"] * 3),
"start": np.sort(rng.choice(1_000_000, n_genes, replace=False)),
}, index=[f"gene{i}" for i in range(n_genes)])
phenotype_pos_df["end"] = phenotype_pos_df["start"] + 1000

genotype_df = pd.DataFrame(rng.integers(0, 3, size=(n_var, n_samples)),
index=variant_df_str.index, columns=[f"s{i}" for i in range(n_samples)])
phenotype_df = pd.DataFrame(rng.normal(size=(n_genes, n_samples)),
index=phenotype_pos_df.index, columns=[f"s{i}" for i in range(n_samples)])

def cis_ranges_for(variant_df):
gen = tensorqtl_genotypeio.InputGeneratorCis(
genotype_df, variant_df, phenotype_df, phenotype_pos_df, window=1_000_000
)
return dict(gen.cis_ranges), gen.chrs, gen.phenotype_df.index.tolist()

ranges_str, chrs_str, kept_str = cis_ranges_for(variant_df_str)
ranges_cat, chrs_cat, kept_cat = cis_ranges_for(variant_df_cat)

assert chrs_str == chrs_cat
assert kept_str == kept_cat
assert set(ranges_str) == set(ranges_cat)
for k in ranges_str:
np.testing.assert_array_equal(ranges_str[k], ranges_cat[k])
Loading
Loading