diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 105b8cc..d889a64 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,7 +50,7 @@ jobs: run: | uv run ruff format --check - pytest: + typing: runs-on: ubuntu-latest needs: [install-job] @@ -69,6 +69,33 @@ jobs: run: | uv sync --all-extras --dev + - name: run type checking + run: | + uv run ty check + + pytest: + runs-on: ubuntu-latest + + needs: [install-job] + + strategy: + matrix: + python-version: ["3.12", "3.13", "3.14"] + + 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: "${{ matrix.python-version }}" + + - name: install dependencies + run: | + uv sync --all-extras --dev + - name: run pytest run: | uv run pytest -v 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/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`) diff --git a/pyproject.toml b/pyproject.toml index b0f6ca1..17a79bd 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 = [ @@ -8,15 +8,16 @@ 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" 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", "pyarrow>=18.0.0", "tqdm>=4.67.1", + "anndata>=0.12.10", ] [build-system] @@ -24,11 +25,12 @@ 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", + "ty>=0.0.19", +] [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 diff --git a/src/cell_eval/_baseline.py b/src/cell_eval/_baseline.py index 7ad1107..0ad610d 100644 --- a/src/cell_eval/_baseline.py +++ b/src/cell_eval/_baseline.py @@ -1,11 +1,12 @@ 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 parallel_differential_expression +from pdex import pdex from scipy.sparse import issparse from ._evaluator import _build_pdex_kwargs, _convert_to_normlog @@ -23,9 +24,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): @@ -67,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, ) @@ -78,21 +77,20 @@ 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") 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}") @@ -137,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 ( @@ -161,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/_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 0e86ce8..30e4dd1 100644 --- a/src/cell_eval/_evaluator.py +++ b/src/cell_eval/_evaluator.py @@ -7,12 +7,13 @@ 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 from ._pipeline import MetricPipeline from ._types import PerturbationAnndataPair, initialize_de_comparison +from .utils import _cast_float16_to_float32 logger = logging.getLogger(__name__) @@ -38,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 @@ -63,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, @@ -96,9 +91,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, @@ -170,6 +163,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) @@ -220,9 +217,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, @@ -233,9 +228,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 +238,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 +249,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 +273,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 +285,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/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/_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] = {} 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 f3f5ad1..8c6ae1e 100644 --- a/src/cell_eval/utils.py +++ b/src/cell_eval/utils.py @@ -1,7 +1,10 @@ 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 logger = logging.getLogger(__name__) @@ -37,8 +40,8 @@ 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: - frac, _ = np.modf(adata.X.toarray()) + 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") else: @@ -57,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: @@ -103,5 +106,23 @@ 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() } + + +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 = 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( + f"Casting {which} anndata from float16 to float32 (NUMBA does not support float16)." + ) + adata.X = x.astype(np.float32) diff --git a/tests/test_eval.py b/tests/test_eval.py index 968a701..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 @@ -265,7 +266,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 +283,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( @@ -390,20 +391,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) 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