From 4d39b85120aae9740d5a19e8c436f5df663ec383 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 10:55:48 -0800 Subject: [PATCH 01/18] refactor: update to new pdex --- pyproject.toml | 10 +++++++--- src/cell_eval/_baseline.py | 11 +++++------ src/cell_eval/_evaluator.py | 38 +++++++++++-------------------------- tests/test_eval.py | 6 +++--- 4 files changed, 26 insertions(+), 39 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index b0f6ca1..2c5e5a1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,10 +8,10 @@ authors = [ { name = "Abhinav Adduri", email = "abhinav.adduri@arcinstitute.org" }, { name = "Yusuf Roohani", email = "yusuf.roohani@arcinstitute.org" }, ] -requires-python = ">=3.10,<3.13" +requires-python = ">=3.11,<3.13" dependencies = [ "igraph>=0.11.8", - "pdex>=0.1.26", + "pdex>=0.2.0", "polars>=1.30.0", "pyyaml>=6.0.2", "scanpy>=1.10.3", @@ -24,7 +24,11 @@ requires = ["hatchling"] build-backend = "hatchling.build" [dependency-groups] -dev = ["ipykernel>=6.29.5", "pytest>=8.3.5", "ruff>=0.11.8"] +dev = [ + "ipykernel>=6.29.5", + "pytest>=8.3.5", + "ruff>=0.11.8", +] [project.scripts] cell-eval = "cell_eval.__main__:main" diff --git a/src/cell_eval/_baseline.py b/src/cell_eval/_baseline.py index 7ad1107..e92874b 100644 --- a/src/cell_eval/_baseline.py +++ b/src/cell_eval/_baseline.py @@ -5,7 +5,7 @@ import numpy as np import polars as pl from numpy.typing import NDArray -from pdex import parallel_differential_expression +from pdex import pdex from scipy.sparse import issparse from ._evaluator import _build_pdex_kwargs, _convert_to_normlog @@ -83,16 +83,15 @@ def build_base_mean_adata( if output_de_path is not None: logger.info("Calculating differential expression") pdex_kwargs = _build_pdex_kwargs( - groupby_key=pert_col, + groupby=pert_col, reference=control_pert, - num_workers=num_threads, - metric=de_method, - batch_size=batch_size, + threads=num_threads, allow_discrete=allow_discrete, pdex_kwargs=pdex_kwargs, ) - frame = parallel_differential_expression( + frame = pdex( adata=baseline_adata, + mode="ref", **pdex_kwargs, ) logger.info(f"Saving differential expression results to {output_de_path}") diff --git a/src/cell_eval/_evaluator.py b/src/cell_eval/_evaluator.py index 0e86ce8..f4cb0ab 100644 --- a/src/cell_eval/_evaluator.py +++ b/src/cell_eval/_evaluator.py @@ -7,7 +7,7 @@ import pandas as pd import polars as pl import scanpy as sc -from pdex import parallel_differential_expression +from pdex import pdex from cell_eval.utils import guess_is_lognorm @@ -233,9 +233,7 @@ def _build_de_comparison( mode="real", de_path=de_real, anndata_pair=anndata_pair, - de_method=de_method, num_threads=num_threads, - batch_size=batch_size, allow_discrete=allow_discrete, outdir=outdir, prefix=prefix, @@ -245,9 +243,7 @@ def _build_de_comparison( mode="pred", de_path=de_pred, anndata_pair=anndata_pair, - de_method=de_method, num_threads=num_threads, - batch_size=batch_size, allow_discrete=allow_discrete, outdir=outdir, prefix=prefix, @@ -258,32 +254,23 @@ def _build_de_comparison( def _build_pdex_kwargs( reference: str, - groupby_key: str, - num_workers: int, - batch_size: int, - metric: str, + groupby: str, + threads: int, allow_discrete: bool, pdex_kwargs: dict[str, Any] | None = None, ) -> dict[str, Any]: pdex_kwargs = pdex_kwargs or {} if "reference" not in pdex_kwargs: pdex_kwargs["reference"] = reference - if "groupby_key" not in pdex_kwargs: - pdex_kwargs["groupby_key"] = groupby_key - if "num_workers" not in pdex_kwargs: - pdex_kwargs["num_workers"] = num_workers - if "batch_size" not in pdex_kwargs: - pdex_kwargs["batch_size"] = batch_size - if "metric" not in pdex_kwargs: - pdex_kwargs["metric"] = metric + if "groupby" not in pdex_kwargs: + pdex_kwargs["groupby"] = groupby + if "threads" not in pdex_kwargs: + pdex_kwargs["threads"] = threads if "is_log1p" not in pdex_kwargs: if allow_discrete: pdex_kwargs["is_log1p"] = False else: pdex_kwargs["is_log1p"] = True - - # always return polars DataFrames - pdex_kwargs["as_polars"] = True return pdex_kwargs @@ -291,9 +278,7 @@ def _load_or_build_de( mode: Literal["pred", "real"], de_path: pl.DataFrame | str | None = None, anndata_pair: PerturbationAnndataPair | None = None, - de_method: str = "wilcoxon", num_threads: int = 1, - batch_size: int = 100, outdir: str | None = None, prefix: str | None = None, allow_discrete: bool = False, @@ -305,16 +290,15 @@ def _load_or_build_de( logger.info(f"Computing DE for {mode} data") pdex_kwargs = _build_pdex_kwargs( reference=anndata_pair.control_pert, - groupby_key=anndata_pair.pert_col, - num_workers=num_threads, - metric=de_method, - batch_size=batch_size, + groupby=anndata_pair.pert_col, + threads=num_threads, allow_discrete=allow_discrete, pdex_kwargs=pdex_kwargs or {}, ) logger.info(f"Using the following pdex kwargs: {pdex_kwargs}") - frame = parallel_differential_expression( + frame = pdex( adata=anndata_pair.real if mode == "real" else anndata_pair.pred, + mode="ref", **pdex_kwargs, ) if outdir is not None: diff --git a/tests/test_eval.py b/tests/test_eval.py index 968a701..541929c 100644 --- a/tests/test_eval.py +++ b/tests/test_eval.py @@ -265,7 +265,7 @@ def test_eval_pdex_kwargs(): control_pert="control", pert_col="perturbation", pdex_kwargs={ - "exp_post_agg": True, + "geometric_mean": False, }, ) evaluator.compute( @@ -282,8 +282,8 @@ def test_eval_pdex_kwargs_duplicated(): control_pert="control", pert_col="perturbation", pdex_kwargs={ - "exp_post_agg": True, - "num_workers": 4, + "geometric_mean": False, + "threads": 4, }, ) evaluator.compute( From 13bfa6cebe5f8dee63b184c3f6753bac08ccd567 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:08:14 -0800 Subject: [PATCH 02/18] refactor: remove unused pdex arguments for batch_size and metric --- src/cell_eval/_baseline.py | 2 -- src/cell_eval/_cli/_run.py | 16 ---------------- src/cell_eval/_evaluator.py | 4 ---- 3 files changed, 22 deletions(-) diff --git a/src/cell_eval/_baseline.py b/src/cell_eval/_baseline.py index e92874b..122edec 100644 --- a/src/cell_eval/_baseline.py +++ b/src/cell_eval/_baseline.py @@ -23,9 +23,7 @@ def build_base_mean_adata( allow_discrete: bool = False, output_path: str | None = None, output_de_path: str | None = None, - batch_size: int = 1000, num_threads: int = 1, - de_method: str = "wilcoxon", pdex_kwargs: dict[str, Any] = {}, ) -> ad.AnnData: if isinstance(adata, str): diff --git a/src/cell_eval/_cli/_run.py b/src/cell_eval/_cli/_run.py index d40f541..d166045 100644 --- a/src/cell_eval/_cli/_run.py +++ b/src/cell_eval/_cli/_run.py @@ -79,18 +79,6 @@ def parse_args_run(parser: ap.ArgumentParser): default=1, help="Number of threads to use for parallel processing [default: %(default)s]", ) - parser.add_argument( - "--batch-size", - type=int, - default=100, - help="Batch size for parallel processing [default: %(default)s]", - ) - parser.add_argument( - "--de-method", - type=str, - default="wilcoxon", - help="Method to use for differential expression analysis [default: %(default)s]", - ) parser.add_argument( "--allow-discrete", action="store_true", @@ -166,9 +154,7 @@ def run_evaluation(args: ap.Namespace): de_real=args.de_real, control_pert=args.control_pert, pert_col=args.pert_col, - de_method=args.de_method, num_threads=args.num_threads, - batch_size=args.batch_size, outdir=args.outdir, allow_discrete=args.allow_discrete, prefix=ct, @@ -189,9 +175,7 @@ def run_evaluation(args: ap.Namespace): de_real=args.de_real, control_pert=args.control_pert, pert_col=args.pert_col, - de_method=args.de_method, num_threads=args.num_threads, - batch_size=args.batch_size, outdir=args.outdir, allow_discrete=args.allow_discrete, skip_de=args.profile == "pds", diff --git a/src/cell_eval/_evaluator.py b/src/cell_eval/_evaluator.py index f4cb0ab..2494434 100644 --- a/src/cell_eval/_evaluator.py +++ b/src/cell_eval/_evaluator.py @@ -96,9 +96,7 @@ def __init__( anndata_pair=self.anndata_pair, de_pred=de_pred, de_real=de_real, - de_method=de_method, num_threads=num_threads if num_threads != -1 else mp.cpu_count(), - batch_size=batch_size, allow_discrete=allow_discrete, outdir=outdir, prefix=prefix, @@ -220,9 +218,7 @@ def _build_de_comparison( anndata_pair: PerturbationAnndataPair | None = None, de_pred: pl.DataFrame | str | None = None, de_real: pl.DataFrame | str | None = None, - de_method: str = "wilcoxon", num_threads: int = 1, - batch_size: int = 100, allow_discrete: bool = False, outdir: str | None = None, prefix: str | None = None, From e27697085c6d5a7816fc0a2d3574e9ae579df7ff Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:11:51 -0800 Subject: [PATCH 03/18] fix: properly recase float16 to float32 at minimum --- src/cell_eval/_evaluator.py | 25 +++++++++++++++++++++++++ 1 file changed, 25 insertions(+) diff --git a/src/cell_eval/_evaluator.py b/src/cell_eval/_evaluator.py index 2494434..8aff393 100644 --- a/src/cell_eval/_evaluator.py +++ b/src/cell_eval/_evaluator.py @@ -168,6 +168,10 @@ def _build_anndata_pair( logger.info(f"Reading pred anndata from {pred}") pred = ad.read_h5ad(pred) + # Cast float16 to float32 since NUMBA (used by pdex) does not support float16 + _cast_float16_to_float32(real, which="real") + _cast_float16_to_float32(pred, which="pred") + # Validate that the input is normalized and log-transformed _convert_to_normlog(real, which="real", allow_discrete=allow_discrete) _convert_to_normlog(pred, which="pred", allow_discrete=allow_discrete) @@ -178,6 +182,27 @@ def _build_anndata_pair( ) +def _cast_float16_to_float32(adata: ad.AnnData, which: str | None = None): + """Cast float16 expression matrix to float32 (inplace). + + NUMBA (used by pdex) does not support float16 operations. + """ + import numpy as np + import scipy.sparse as sp + + x = adata.X + dtype = x.dtype if not sp.issparse(x) else x.data.dtype + if dtype == np.float16: + if which: + logger.info( + f"Casting {which} anndata from float16 to float32 (NUMBA does not support float16)." + ) + if sp.issparse(x): + adata.X = x.astype(np.float32) + else: + adata.X = x.astype(np.float32) + + def _convert_to_normlog( adata: ad.AnnData, which: str | None = None, From e5407e42a84ef17fd106d0c2eeb95fba1e9a50f2 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:13:16 -0800 Subject: [PATCH 04/18] refactor: move recast to utils --- src/cell_eval/_evaluator.py | 22 +--------------------- src/cell_eval/utils.py | 20 ++++++++++++++++++++ 2 files changed, 21 insertions(+), 21 deletions(-) diff --git a/src/cell_eval/_evaluator.py b/src/cell_eval/_evaluator.py index 8aff393..a6cc767 100644 --- a/src/cell_eval/_evaluator.py +++ b/src/cell_eval/_evaluator.py @@ -13,6 +13,7 @@ from ._pipeline import MetricPipeline from ._types import PerturbationAnndataPair, initialize_de_comparison +from .utils import _cast_float16_to_float32 logger = logging.getLogger(__name__) @@ -182,27 +183,6 @@ def _build_anndata_pair( ) -def _cast_float16_to_float32(adata: ad.AnnData, which: str | None = None): - """Cast float16 expression matrix to float32 (inplace). - - NUMBA (used by pdex) does not support float16 operations. - """ - import numpy as np - import scipy.sparse as sp - - x = adata.X - dtype = x.dtype if not sp.issparse(x) else x.data.dtype - if dtype == np.float16: - if which: - logger.info( - f"Casting {which} anndata from float16 to float32 (NUMBA does not support float16)." - ) - if sp.issparse(x): - adata.X = x.astype(np.float32) - else: - adata.X = x.astype(np.float32) - - def _convert_to_normlog( adata: ad.AnnData, which: str | None = None, diff --git a/src/cell_eval/utils.py b/src/cell_eval/utils.py index f3f5ad1..4c6fbb9 100644 --- a/src/cell_eval/utils.py +++ b/src/cell_eval/utils.py @@ -2,6 +2,7 @@ import anndata as ad import numpy as np +import scipy.sparse as sp from scipy.sparse import csc_matrix, csr_matrix logger = logging.getLogger(__name__) @@ -105,3 +106,22 @@ def split_anndata_on_celltype( ct: adata[adata.obs[celltype_col] == ct] for ct in adata.obs[celltype_col].unique() } + + +def _cast_float16_to_float32(adata: ad.AnnData, which: str | None = None): + """Cast float16 expression matrix to float32 (inplace). + + NUMBA (used by pdex) does not support float16 operations. + """ + + x = adata.X + dtype = x.dtype if not sp.issparse(x) else x.data.dtype + if dtype == np.float16: + if which: + logger.info( + f"Casting {which} anndata from float16 to float32 (NUMBA does not support float16)." + ) + if sp.issparse(x): + adata.X = x.astype(np.float32) + else: + adata.X = x.astype(np.float32) From f3ea3abead4e3ad0e68de16339e9fa953cde9030 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:14:51 -0800 Subject: [PATCH 05/18] test: remove unused alt metric --- tests/test_eval.py | 17 ----------------- 1 file changed, 17 deletions(-) diff --git a/tests/test_eval.py b/tests/test_eval.py index 541929c..79ee17a 100644 --- a/tests/test_eval.py +++ b/tests/test_eval.py @@ -390,20 +390,3 @@ def test_eval_downsampled_cells(): break_on_error=True, ) validate_expected_files(OUTDIR) - - -def test_eval_alt_metric(): - adata_real = build_random_anndata() - adata_pred = downsample_cells(adata_real, fraction=0.5) - evaluator = MetricsEvaluator( - adata_pred=adata_pred, - adata_real=adata_real, - control_pert=CONTROL_VAR, - pert_col=PERT_COL, - outdir=OUTDIR, - de_method="anderson", - ) - evaluator.compute( - break_on_error=True, - ) - validate_expected_files(OUTDIR) From 824e63aa22b87870baea864bfdaa0e226c707d5f Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:15:03 -0800 Subject: [PATCH 06/18] chore: finish removing all outdated pdex arguments --- src/cell_eval/_evaluator.py | 6 ------ 1 file changed, 6 deletions(-) diff --git a/src/cell_eval/_evaluator.py b/src/cell_eval/_evaluator.py index a6cc767..30e4dd1 100644 --- a/src/cell_eval/_evaluator.py +++ b/src/cell_eval/_evaluator.py @@ -39,12 +39,8 @@ class MetricsEvaluator: Control perturbation name. pert_col: str = "target" Perturbation column name. - de_method: str = "wilcoxon" - Differential expression method. num_threads: int = -1 Number of threads for parallel differential expression. - batch_size: int = 100 - Batch size for parallel differential expression. outdir: str = "./cell-eval-outdir" Output directory. allow_discrete: bool = False @@ -64,9 +60,7 @@ def __init__( de_real: pl.DataFrame | str | None = None, control_pert: str = "non-targeting", pert_col: str = "target", - de_method: str = "wilcoxon", num_threads: int = -1, - batch_size: int = 100, outdir: str = "./cell-eval-outdir", allow_discrete: bool = False, prefix: str | None = None, From 76b3f934e8c99ca0d49ffc8d0e4d2ce2d0a1429c Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:19:01 -0800 Subject: [PATCH 07/18] dep: added ty to project --- pyproject.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/pyproject.toml b/pyproject.toml index 2c5e5a1..c776b55 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,6 +28,7 @@ dev = [ "ipykernel>=6.29.5", "pytest>=8.3.5", "ruff>=0.11.8", + "ty>=0.0.19", ] [project.scripts] From e25e1e1285640ce8148e30d2f3ae003f490aab70 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:27:31 -0800 Subject: [PATCH 08/18] style: fix all typing errors or ambiguities --- src/cell_eval/_baseline.py | 13 ++++++------ src/cell_eval/_cli/_prep.py | 11 ++++++----- src/cell_eval/_types/_anndata.py | 15 +++++++++----- src/cell_eval/metrics/_anndata.py | 30 +++++++++++++++++----------- src/cell_eval/metrics/base.py | 4 ++-- src/cell_eval/utils.py | 16 +++++++++------ tests/test_eval.py | 13 ++++++------ tutorials/vcc/vcc.ipynb | 33 +++++-------------------------- 8 files changed, 66 insertions(+), 69 deletions(-) diff --git a/src/cell_eval/_baseline.py b/src/cell_eval/_baseline.py index 122edec..0ad610d 100644 --- a/src/cell_eval/_baseline.py +++ b/src/cell_eval/_baseline.py @@ -1,8 +1,9 @@ import logging -from typing import Any +from typing import Any, cast import anndata as ad import numpy as np +import pandas as pd import polars as pl from numpy.typing import NDArray from pdex import pdex @@ -65,7 +66,7 @@ def build_base_mean_adata( (int(counts[counts_col].sum()), baseline.size), baseline, ), - var=adata.var, + var=cast(pd.DataFrame, adata.var), obs=obs, ) @@ -76,7 +77,7 @@ def build_base_mean_adata( if output_path is not None: logger.info(f"Saving baseline data to {output_path}") - baseline_adata.write_h5ad(output_path) + baseline_adata.write_h5ad(output_path) # type: ignore[invalid-argument-type] if output_de_path is not None: logger.info("Calculating differential expression") @@ -134,9 +135,9 @@ def _build_counts_df_from_adata( raise ValueError( f"Column '{pert_col}' not found in adata.obs: {adata.obs.columns}" ) - if control_pert not in adata.obs[pert_col].unique(): + if control_pert not in cast(pd.Series, adata.obs[pert_col]).unique(): raise ValueError( - f"Control pert '{control_pert}' not found in adata.obs[{pert_col}]: {adata.obs[pert_col].unique()}" + f"Control pert '{control_pert}' not found in adata.obs[{pert_col}]: {cast(pd.Series, adata.obs[pert_col]).unique()}" ) logger.info("Building counts DataFrame from adata") return ( @@ -158,7 +159,7 @@ def _build_pert_baseline( raise ValueError( f"Column '{pert_col}' not found in adata.obs: {adata.obs.columns}" ) - unique_perts = adata.obs[pert_col].unique() + unique_perts = cast(pd.Series, adata.obs[pert_col]).unique() if control_pert not in unique_perts: raise ValueError( f"Control pert '{control_pert}' not found in unique_perts: {unique_perts}" diff --git a/src/cell_eval/_cli/_prep.py b/src/cell_eval/_cli/_prep.py index aba2e98..80e5d67 100644 --- a/src/cell_eval/_cli/_prep.py +++ b/src/cell_eval/_cli/_prep.py @@ -5,6 +5,7 @@ import shutil import subprocess from tempfile import TemporaryDirectory +from typing import cast import anndata as ad import numpy as np @@ -136,9 +137,9 @@ def strip_anndata( raise ValueError( f"Provided celltype column: '{celltype_col}' missing from anndata: {adata.obs.columns}" ) - if ntc_name not in adata.obs[pert_col].unique(): + if ntc_name not in cast(pd.Series, adata.obs[pert_col]).unique(): raise ValueError( - f"Provided negative control name: '{ntc_name}' missing from anndata: {adata.obs[pert_col].unique()}" + f"Provided negative control name: '{ntc_name}' missing from anndata: {cast(pd.Series, adata.obs[pert_col]).unique()}" ) # Check if expected dimension is provided and matches the length of the genelist @@ -196,11 +197,11 @@ def strip_anndata( logger.info("Simplifying obs dataframe") new_obs = pd.DataFrame( - {output_pert_col: adata.obs[pert_col].values}, + {output_pert_col: cast(pd.Series, adata.obs[pert_col]).values}, index=np.arange(adata.shape[0]).astype(str), ) if celltype_col: - new_obs[output_celltype_col] = adata.obs[celltype_col].values + new_obs[output_celltype_col] = cast(pd.Series, adata.obs[celltype_col]).values logger.info("Simplifying var dataframe") new_var = pd.DataFrame( @@ -225,7 +226,7 @@ def strip_anndata( # Write the h5ad file logger.info(f"Writing h5ad output to {tmp_h5ad}") - minimal.write_h5ad(tmp_h5ad) + minimal.write_h5ad(tmp_h5ad) # type: ignore[invalid-argument-type] # Zstd compress the h5ad file (will create pred.h5ad.zst) logger.info(f"Zstd compressing {tmp_h5ad}") diff --git a/src/cell_eval/_types/_anndata.py b/src/cell_eval/_types/_anndata.py index 2364ae1..3459d2e 100644 --- a/src/cell_eval/_types/_anndata.py +++ b/src/cell_eval/_types/_anndata.py @@ -1,9 +1,10 @@ import logging from dataclasses import dataclass, field -from typing import Iterator, Literal +from typing import Iterator, Literal, cast import anndata as ad import numpy as np +import pandas as pd import polars as pl from numpy.typing import NDArray from scipy.sparse import issparse @@ -70,8 +71,12 @@ def __post_init__(self) -> None: f"Perturbation column ({self.pert_col}) not found in pred AnnData: {self.pred.obs.columns}" ) - perts_real = np.unique(self.real.obs[self.pert_col].to_numpy(str)) - perts_pred = np.unique(self.pred.obs[self.pert_col].to_numpy(str)) + perts_real = np.unique( + cast(pd.Series, self.real.obs[self.pert_col]).to_numpy(str) + ) + perts_pred = np.unique( + cast(pd.Series, self.pred.obs[self.pert_col]).to_numpy(str) + ) if not np.array_equal(perts_real, perts_pred): raise ValueError( f"Perturbation mismatch: real {perts_real} != pred {perts_pred}" @@ -90,10 +95,10 @@ def __post_init__(self) -> None: perts = np.array([p for p in perts if p != self.control_pert]) pert_mask_real = self.pert_mask( - self.real.obs[self.pert_col].to_numpy(str), + cast(pd.Series, self.real.obs[self.pert_col]).to_numpy(str), ) pert_mask_pred = self.pert_mask( - self.pred.obs[self.pert_col].to_numpy(str), + cast(pd.Series, self.pred.obs[self.pert_col]).to_numpy(str), ) object.__setattr__(self, "perts", perts) diff --git a/src/cell_eval/metrics/_anndata.py b/src/cell_eval/metrics/_anndata.py index 8bcdf7d..458de62 100644 --- a/src/cell_eval/metrics/_anndata.py +++ b/src/cell_eval/metrics/_anndata.py @@ -1,7 +1,7 @@ """Array metrics module.""" from logging import getLogger -from typing import Callable, Literal, Sequence +from typing import Callable, Literal, Sequence, cast import anndata as ad import numpy as np @@ -27,7 +27,7 @@ def pearson_delta( """Compute Pearson correlation between mean differences from control.""" return _generic_evaluation( data, - pearsonr, # type: ignore + pearsonr, use_delta=True, embed_key=embed_key, ) @@ -287,21 +287,21 @@ def _centroid_ann( feats = adata.obsm.get(embed_key, adata.X) # type: ignore # Convert to float if not already - if feats.dtype != np.dtype("float64"): # type: ignore - feats = feats.astype(np.float64) # type: ignore + if feats.dtype != np.dtype("float64"): + feats = feats.astype(np.float64) # Densify if required if issparse(feats): - feats = feats.toarray() # type: ignore + feats = feats.toarray() - cats = adata.obs[category_key].values - uniq, inv = np.unique(cats, return_inverse=True) # type: ignore - centroids = np.zeros((uniq.size, feats.shape[1]), dtype=feats.dtype) # type: ignore + cats = cast(pd.Series, adata.obs[category_key]).values + uniq, inv = np.unique(cats, return_inverse=True) + centroids = np.zeros((uniq.size, feats.shape[1]), dtype=feats.dtype) for i, cat in enumerate(uniq): mask = cats == cat if np.any(mask): - centroids[i] = feats[mask].mean(axis=0) # type: ignore + centroids[i] = feats[mask].mean(axis=0) adc = ad.AnnData(X=centroids) adc.obs[category_key] = uniq @@ -329,12 +329,20 @@ def __call__(self, data: PerturbationAnndataPair) -> float: self._cluster_leiden( ad_real_cent, self.real_resolution, real_key, self.n_neighbors ) - ad_real_cent.obs = ad_real_cent.obs.set_index(data.pert_col).loc[cats_sorted] + ad_real_cent.obs = ( + cast(pd.DataFrame, ad_real_cent.obs) + .set_index(data.pert_col) + .loc[cats_sorted] + ) real_labels = pd.Categorical(ad_real_cent.obs[real_key]) # 4. sweep predicted resolutions best_score = 0.0 - ad_pred_cent.obs = ad_pred_cent.obs.set_index(data.pert_col).loc[cats_sorted] + ad_pred_cent.obs = ( + cast(pd.DataFrame, ad_pred_cent.obs) + .set_index(data.pert_col) + .loc[cats_sorted] + ) for r in self.pred_resolutions: pred_key = f"pred_clusters_{r}" self._cluster_leiden(ad_pred_cent, r, pred_key, self.n_neighbors) diff --git a/src/cell_eval/metrics/base.py b/src/cell_eval/metrics/base.py index e8e77c8..85a5a7b 100644 --- a/src/cell_eval/metrics/base.py +++ b/src/cell_eval/metrics/base.py @@ -22,10 +22,10 @@ class MetricResult: value: float | str perturbation: str | None = None - def to_dict(self) -> dict[str, float | str]: + def to_dict(self) -> dict[str, float | str | None]: """Convert result to dictionary.""" return { - "perturbation": self.perturbation, # type: ignore + "perturbation": self.perturbation, "metric": self.name, "value": self.value, } diff --git a/src/cell_eval/utils.py b/src/cell_eval/utils.py index 4c6fbb9..2b9f431 100644 --- a/src/cell_eval/utils.py +++ b/src/cell_eval/utils.py @@ -1,7 +1,9 @@ import logging +from typing import cast import anndata as ad import numpy as np +import pandas as pd import scipy.sparse as sp from scipy.sparse import csc_matrix, csr_matrix @@ -39,7 +41,7 @@ def guess_is_lognorm( if isinstance(adata.X, csr_matrix) or isinstance(adata.X, csc_matrix): frac, _ = np.modf(adata.X.data) elif adata.isview: - frac, _ = np.modf(adata.X.toarray()) + frac, _ = np.modf(adata.X.toarray()) # type: ignore[unresolved-attribute] elif adata.X is None: raise ValueError("adata.X is None") else: @@ -58,8 +60,8 @@ def guess_is_lognorm( max_val = adata.X.max() min_val = adata.X.min() else: - max_val = float(np.max(adata.X)) - min_val = float(np.min(adata.X)) + max_val = float(np.max(adata.X)) # type: ignore[no-matching-overload] + min_val = float(np.min(adata.X)) # type: ignore[no-matching-overload] # Validate range if min_val < 0: @@ -104,7 +106,7 @@ def split_anndata_on_celltype( return { ct: adata[adata.obs[celltype_col] == ct] - for ct in adata.obs[celltype_col].unique() + for ct in cast(pd.Series, adata.obs[celltype_col]).unique() } @@ -114,8 +116,10 @@ def _cast_float16_to_float32(adata: ad.AnnData, which: str | None = None): NUMBA (used by pdex) does not support float16 operations. """ - x = adata.X - dtype = x.dtype if not sp.issparse(x) else x.data.dtype + x = cast(np.ndarray | csr_matrix | csc_matrix, adata.X) + dtype = ( + x.dtype if not sp.issparse(x) else cast(csr_matrix | csc_matrix, x).data.dtype + ) if dtype == np.float16: if which: logger.info( diff --git a/tests/test_eval.py b/tests/test_eval.py index 79ee17a..48beb4f 100644 --- a/tests/test_eval.py +++ b/tests/test_eval.py @@ -1,8 +1,9 @@ import os import shutil -from typing import Literal +from typing import Literal, cast import numpy as np +import pandas as pd import pytest from cell_eval import MetricsEvaluator @@ -118,7 +119,7 @@ def test_broken_adata_missing_pertcol_in_real(): adata_pred = adata_real.copy() # Remove pert_col from adata_real - adata_real.obs.drop(columns=[PERT_COL], inplace=True) + cast(pd.DataFrame, adata_real.obs).drop(columns=[PERT_COL], inplace=True) with pytest.raises(Exception): MetricsEvaluator( @@ -135,7 +136,7 @@ def test_broken_adata_missing_pertcol_in_pred(): adata_pred = adata_real.copy() # Remove pert_col from adata_pred - adata_pred.obs.drop(columns=[PERT_COL], inplace=True) + cast(pd.DataFrame, adata_pred.obs).drop(columns=[PERT_COL], inplace=True) with pytest.raises(Exception): MetricsEvaluator( @@ -195,7 +196,7 @@ def test_unknown_alternative_de_metric(): control_pert=CONTROL_VAR, pert_col=PERT_COL, outdir=OUTDIR, - de_method="unknown", + de_method="unknown", # type: ignore[unknown-argument] ).compute() @@ -239,8 +240,8 @@ def test_eval_missing_celltype_col(): adata_real = build_random_anndata() adata_pred = downsample_cells(adata_real, fraction=0.5) - adata_real.obs.drop(columns="celltype", inplace=True) - adata_pred.obs.drop(columns="celltype", inplace=True) + cast(pd.DataFrame, adata_real.obs).drop(columns="celltype", inplace=True) + cast(pd.DataFrame, adata_pred.obs).drop(columns="celltype", inplace=True) assert "celltype" not in adata_real.obs.columns assert "celltype" not in adata_pred.obs.columns diff --git a/tutorials/vcc/vcc.ipynb b/tutorials/vcc/vcc.ipynb index 5f023c2..e2dc17f 100644 --- a/tutorials/vcc/vcc.ipynb +++ b/tutorials/vcc/vcc.ipynb @@ -239,32 +239,11 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": null, "id": "3879267c", "metadata": {}, "outputs": [], - "source": [ - "# Define our path to the training anndata\n", - "tr_adata_path = \"./adata_Training.h5ad\"\n", - "\n", - "# Read in the anndata\n", - "tr_adata = ad.read_h5ad(tr_adata_path)\n", - "\n", - "# Filter for non-targeting\n", - "ntc_adata = tr_adata[tr_adata.obs[\"target_gene\"] == \"non-targeting\"]\n", - "\n", - "# Append the non-targeting controls to the example anndata if they're missing\n", - "if \"non-targeting\" not in adata.obs[\"target_gene\"].unique():\n", - " assert np.all(adata.var_names.values == ntc_adata.var_names.values), (\n", - " \"Gene-Names are out of order or unequal\"\n", - " )\n", - " adata = ad.concat(\n", - " [\n", - " adata,\n", - " ntc_adata,\n", - " ]\n", - " )" - ] + "source": "import pandas as pd\n\n# Define our path to the training anndata\ntr_adata_path = \"./adata_Training.h5ad\"\n\n# Read in the anndata\ntr_adata = ad.read_h5ad(tr_adata_path)\n\n# Filter for non-targeting\nntc_adata = tr_adata[tr_adata.obs[\"target_gene\"] == \"non-targeting\"]\n\n# Append the non-targeting controls to the example anndata if they're missing\nif \"non-targeting\" not in pd.Series(adata.obs[\"target_gene\"]).unique():\n assert np.array_equal(adata.var_names, ntc_adata.var_names), (\n \"Gene-Names are out of order or unequal\"\n )\n adata = ad.concat(\n [\n adata,\n ntc_adata,\n ]\n )" }, { "cell_type": "markdown", @@ -276,13 +255,11 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "id": "386ee994", "metadata": {}, "outputs": [], - "source": [ - "adata.write_h5ad(\"./example.h5ad\")" - ] + "source": "adata.write_h5ad(\"./example.h5ad\") # type: ignore[invalid-argument-type]" }, { "cell_type": "markdown", @@ -336,4 +313,4 @@ }, "nbformat": 4, "nbformat_minor": 5 -} +} \ No newline at end of file From 4c529f9fa9f99caf3269aaa9abfa75eee0ab15ae Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:28:42 -0800 Subject: [PATCH 09/18] ci: added typing to ci --- .github/workflows/CI.yml | 25 ++++++++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 105b8cc..ffa890d 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -5,7 +5,7 @@ on: [push, pull_request] jobs: all_jobs: runs-on: ubuntu-latest - needs: [formatting, pytest, cli-test] + needs: [formatting, typing, pytest, cli-test] steps: - name: Complete run: echo "Complete" @@ -50,6 +50,29 @@ jobs: run: | uv run ruff format --check + typing: + runs-on: ubuntu-latest + + needs: [install-job] + + steps: + - uses: actions/checkout@v4 + + - name: install uv + uses: astral-sh/setup-uv@v5 + with: + enable-cache: true + cache-dependency-glob: "pyproject.toml" + python-version: "3.12" + + - name: install dependencies + run: | + uv sync --all-extras --dev + + - name: run type checking + run: | + uv run ty check + pytest: runs-on: ubuntu-latest From a3609d14097555bd521e20fe23b50df24371f38a Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:36:16 -0800 Subject: [PATCH 10/18] chore: remove unused configuration --- pyproject.toml | 4 ---- ruff.toml | 5 ----- 2 files changed, 9 deletions(-) delete mode 100644 ruff.toml diff --git a/pyproject.toml b/pyproject.toml index c776b55..d1bba78 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,7 +33,3 @@ dev = [ [project.scripts] cell-eval = "cell_eval.__main__:main" - -[tool.pyright] -venvPath = "." -venv = ".venv" diff --git a/ruff.toml b/ruff.toml deleted file mode 100644 index 235fc08..0000000 --- a/ruff.toml +++ /dev/null @@ -1,5 +0,0 @@ -[lint] -select = ["E", "F", "ERA"] - -[lint.pycodestyle] -max-line-length = 120 From fe81c31c6b91e84034fa267c1719efdae029c641 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:36:22 -0800 Subject: [PATCH 11/18] feat: added a claude md --- CLAUDE.md | 84 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 84 insertions(+) create mode 100644 CLAUDE.md diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 0000000..cda10e2 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,84 @@ +# CLAUDE.md + +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. + +## Project Overview + +**cell-eval** is a Python package and CLI tool for evaluating the performance of models that predict cellular responses to perturbations at the single-cell level. Developed by the Arc Research Institute. + +It generally revolves around a *real* anndata and a *predicted* anndata where it measures the general differences between the two across a variety of metrics. + +- Python 3.11–3.12, managed with **UV** and built with **hatchling** +- CLI entry point: `cell-eval` (defined in `src/cell_eval/__main__.py`) + +## Common Commands + +```bash +# Install dependencies +uv sync --all-extras --dev + +# Run all tests +uv run pytest -v + +# Run a single test +uv run pytest tests/test_eval.py::test_broken_adata_not_normlog -v + +# Formatting (check / fix) +uv run ruff format --check +uv run ruff format + +# Type checking +uv run ty check + +# Verify CLI works +uv run cell-eval --help +``` + +CI runs: formatting, typing, pytest, and cli-test (see `.github/workflows/CI.yml`). + +## Architecture + +### Core Data Flow + +``` +AnnData inputs (predicted + real) + → MetricsEvaluator (validation, normalization, DE computation) + → MetricPipeline (profile-based metric selection + execution) + → metrics_registry (global MetricRegistry instance) + → individual metric functions + → polars DataFrames (per-perturbation + aggregated results) +``` + +### Key Abstractions + +- **`MetricsEvaluator`** (`src/cell_eval/_evaluator.py`) — Main programmatic entry point. Validates input AnnData objects, computes differential expression via `pdex`, and orchestrates the metric pipeline. + +- **`MetricRegistry`** (`src/cell_eval/metrics/_registry.py`) — Global singleton `metrics_registry`. Metrics are registered with a name, type (`DE` or `ANNDATA_PAIR`), compute function, and best-value indicator. Supports both plain functions and class-based metrics requiring instantiation. + +- **`MetricPipeline`** (`src/cell_eval/_pipeline/_runner.py`) — Selects and runs metrics based on a profile (`full`, `minimal`, `vcc`, `de`, `anndata`, `pds`). Collects per-perturbation results and aggregates them. + +- **`Metric` protocol** (`src/cell_eval/metrics/base.py`) — All metric functions take either a `PerturbationAnndataPair` or `DEComparison` and return `float | dict[str, float]`. + +- **Type system** (`src/cell_eval/_types/`) — Immutable dataclasses: `PerturbationAnndataPair`, `DEComparison`, plus enums `MetricType`, `MetricBestValue`, `DESortBy`. + +### Metrics + +Metrics are split into two categories registered in `src/cell_eval/metrics/_impl.py`: + +- **AnnData metrics** (`_anndata.py`): pearson_delta, mse, mae, mse_delta, mae_delta, discrimination_score, clustering_agreement, edistance +- **DE metrics** (`_de.py`): overlap/precision at N, spearman correlations, direction match, significant gene recall, ROC/PR AUC + +### CLI + +Subcommands in `src/cell_eval/_cli/`: `prep` (data preparation for VCC), `run` (evaluation), `baseline` (create baseline), `score` (normalize against baseline). CLI defaults are in `_cli/_const.py`. + +### Test Data Utilities + +`cell_eval.data` provides `build_random_anndata()` and `downsample_cells()` for generating synthetic AnnData objects in tests. + +## Conventions + +- Uses `polars` (not pandas) for DataFrames +- Uses `match`/`case` statements (Python 3.10+ syntax) +- Type hints throughout; PEP 561 `py.typed` marker present +- Private modules prefixed with `_` (public API is re-exported from `__init__.py`) From 8901383f1d5b154c3ed9226d5b4c65e93c1c2612 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 11:43:56 -0800 Subject: [PATCH 12/18] chore(semver): bump - breaking changes --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index d1bba78..aba1108 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "cell-eval" -version = "0.6.8" +version = "0.7.0" description = "Evaluation metrics for single-cell perturbation predictions" readme = "README.md" authors = [ From 4984a711b1fffd02060c94bbda661bde2e542a4d Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 14:41:57 -0800 Subject: [PATCH 13/18] refactor: remove redundant code --- src/cell_eval/utils.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/src/cell_eval/utils.py b/src/cell_eval/utils.py index 2b9f431..7828e74 100644 --- a/src/cell_eval/utils.py +++ b/src/cell_eval/utils.py @@ -125,7 +125,4 @@ def _cast_float16_to_float32(adata: ad.AnnData, which: str | None = None): logger.info( f"Casting {which} anndata from float16 to float32 (NUMBA does not support float16)." ) - if sp.issparse(x): - adata.X = x.astype(np.float32) - else: - adata.X = x.astype(np.float32) + adata.X = x.astype(np.float32) From e7047a2a5fc97de82470067a54e26ba0de2094be Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 14:43:31 -0800 Subject: [PATCH 14/18] fix: deprecation warning on is_view --- src/cell_eval/utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/cell_eval/utils.py b/src/cell_eval/utils.py index 7828e74..8c6ae1e 100644 --- a/src/cell_eval/utils.py +++ b/src/cell_eval/utils.py @@ -40,7 +40,7 @@ def guess_is_lognorm( # Check for fractional values if isinstance(adata.X, csr_matrix) or isinstance(adata.X, csc_matrix): frac, _ = np.modf(adata.X.data) - elif adata.isview: + elif adata.is_view: frac, _ = np.modf(adata.X.toarray()) # type: ignore[unresolved-attribute] elif adata.X is None: raise ValueError("adata.X is None") From 3e36c0ea26c6176454b135a8336a1218d03b924a Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Thu, 26 Feb 2026 14:45:18 -0800 Subject: [PATCH 15/18] dep: enforce anndata version --- pyproject.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/pyproject.toml b/pyproject.toml index aba1108..0e6c69b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,6 +17,7 @@ dependencies = [ "scanpy>=1.10.3", "pyarrow>=18.0.0", "tqdm>=4.67.1", + "anndata>=0.12.10", ] [build-system] From 2d314f5064afa55f476ff1e6d6e5f16b741593e9 Mon Sep 17 00:00:00 2001 From: Beatrice Bevilacqua Date: Thu, 5 Mar 2026 18:25:17 +0000 Subject: [PATCH 16/18] Use left join for DE Spearman/AUC Fixes #227 --- src/cell_eval/metrics/_de.py | 28 +++++++++++++++++----------- 1 file changed, 17 insertions(+), 11 deletions(-) diff --git a/src/cell_eval/metrics/_de.py b/src/cell_eval/metrics/_de.py index 1f2ca00..bdae102 100644 --- a/src/cell_eval/metrics/_de.py +++ b/src/cell_eval/metrics/_de.py @@ -116,11 +116,15 @@ def __call__(self, data: DEComparison) -> dict[str, float]: """Compute correlation between log fold changes of significant genes.""" correlations = {} - merged = data.real.filter_to_significant(fdr_threshold=self.fdr_threshold).join( - data.pred.data, - on=[data.real.target_col, data.real.feature_col], - suffix="_pred", - how="inner", + merged = ( + data.real.filter_to_significant(fdr_threshold=self.fdr_threshold) + .join( + data.pred.data, + on=[data.real.target_col, data.real.feature_col], + suffix="_pred", + how="left", + ) + .with_columns(pl.col(f"{data.real.fold_change_col}_pred").fill_null(0.0)) ) for row in ( @@ -224,17 +228,19 @@ def compute_generic_auc( (pl.col(real_fdr_col) < 0.05).cast(pl.Float32).alias("label") ).select([target_col, feature_col, "label"]) + pred_q = pl.col(pred_fdr_col).fill_null(1.0).clip(1e-10, 1.0) merged = ( - data.pred.data.select([target_col, feature_col, pred_fdr_col]) - .join( - labeled_real, + labeled_real.join( + data.pred.data.select([target_col, feature_col, pred_fdr_col]), on=[target_col, feature_col], - how="inner", + how="left", coalesce=True, ) .drop_nulls(["label"]) - .with_columns((-pl.col(pred_fdr_col).replace(0, 1e-10).log10()).alias("nlp")) - .drop_nulls(["nlp"]) + .with_columns( + pred_q.alias(pred_fdr_col), + (-pred_q.log10()).alias("nlp"), + ) ) results: dict[str, float] = {} From ff5acaeea50b10eefb5ddf430b9733903a4609f8 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Tue, 24 Mar 2026 11:06:58 -0700 Subject: [PATCH 17/18] dep: remove upper limit on python version --- .python-version | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.python-version b/.python-version index e4fba21..6324d40 100644 --- a/.python-version +++ b/.python-version @@ -1 +1 @@ -3.12 +3.14 diff --git a/pyproject.toml b/pyproject.toml index 0e6c69b..17a79bd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,7 @@ authors = [ { name = "Abhinav Adduri", email = "abhinav.adduri@arcinstitute.org" }, { name = "Yusuf Roohani", email = "yusuf.roohani@arcinstitute.org" }, ] -requires-python = ">=3.11,<3.13" +requires-python = ">=3.11" dependencies = [ "igraph>=0.11.8", "pdex>=0.2.0", From e629a20f9e783a527f8783730fd9993d9fc2bc48 Mon Sep 17 00:00:00 2001 From: noam teyssier <22600644+noamteyssier@users.noreply.github.com> Date: Tue, 24 Mar 2026 11:08:36 -0700 Subject: [PATCH 18/18] ci: added more python versions to pytest --- .github/workflows/CI.yml | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index ffa890d..d889a64 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -78,6 +78,10 @@ jobs: needs: [install-job] + strategy: + matrix: + python-version: ["3.12", "3.13", "3.14"] + steps: - uses: actions/checkout@v4 @@ -86,7 +90,7 @@ jobs: with: enable-cache: true cache-dependency-glob: "pyproject.toml" - python-version: "3.12" + python-version: "${{ matrix.python-version }}" - name: install dependencies run: |