diff --git a/server/app/app.py b/server/app/app.py index 44b2718e..05a08da6 100644 --- a/server/app/app.py +++ b/server/app/app.py @@ -104,6 +104,8 @@ cellxgene is a local web application for exploring single cell expression. launch_group.add_argument("--debug", action="store_true", help=argparse.SUPPRESS) launch_group.add_argument("--no-open", help="do not launch the webbrowser", action="store_false", dest="open_browser") + launch_group.add_argument("--obs-names", help="Annotation name to use as unique, human-readable observation name") + launch_group.add_argument("--var-names", help="Annotation name to use as unique, human-readable variable name") launch_group.add_argument( "--max-category-items", type=whole_number, @@ -141,8 +143,7 @@ def run_scanpy(args): log.setLevel(logging.ERROR) from .scanpy_engine.scanpy_engine import ScanpyEngine print(f"Loading data from {args.data} (this may take a while)") - app.data = ScanpyEngine(args.data, layout_method=args.layout, diffexp_method=args.diffexp, - max_category_items=args.max_category_items) + app.data = ScanpyEngine(args.data, args) print(f"Launching cellxgene") if args.open_browser: webbrowser.open(cellxgene_url) diff --git a/server/app/driver/driver.py b/server/app/driver/driver.py index 5f0203fd..80f8c9c6 100644 --- a/server/app/driver/driver.py +++ b/server/app/driver/driver.py @@ -12,11 +12,11 @@ Sort order for methods class CXGDriver(metaclass=ABCMeta): - def __init__(self, data, layout_method=None, diffexp_method=None, max_category_items=100): + def __init__(self, data, args): self.data = self._load_data(data) - self.layout_method = layout_method - self.diffexp_method = diffexp_method - self.max_category_items = max_category_items + self.layout_method = args.layout + self.diffexp_method = args.diffexp + self.max_category_items = args.max_category_items self.cluster = None @property @@ -44,14 +44,6 @@ class CXGDriver(metaclass=ABCMeta): def _load_data(data): pass - @abstractmethod - def cells(self): - pass - - @abstractmethod - def genes(self): - pass - @abstractmethod def filter_dataframe(self, filter): """ diff --git a/server/app/scanpy_engine/scanpy_engine.py b/server/app/scanpy_engine/scanpy_engine.py index 8f8f1b5f..b52b0742 100644 --- a/server/app/scanpy_engine/scanpy_engine.py +++ b/server/app/scanpy_engine/scanpy_engine.py @@ -1,9 +1,9 @@ import warnings import numpy as np -from pandas import DataFrame, Series +from pandas import DataFrame import scanpy.api as sc -from scipy import stats +from scipy import stats, sparse from server.app.app import cache from server.app.driver.driver import CXGDriver @@ -22,17 +22,46 @@ Sort order for methods class ScanpyEngine(CXGDriver): - def __init__(self, data, layout_method=None, diffexp_method=None, max_category_items=100): - super().__init__(data, layout_method=layout_method, diffexp_method=diffexp_method, - max_category_items=max_category_items) - self._validatate_data_types() - self._add_mandatory_annotations() + def __init__(self, data, args): + super().__init__(data, args) + self._alias_annotation_names(Axis.OBS, args.obs_names) + self._alias_annotation_names(Axis.VAR, args.var_names) + self._validate_data_types() self.cell_count = self.data.shape[0] self.gene_count = self.data.shape[1] self.layout_options = ["umap", "tsne"] self.diffexp_options = ["ttest"] self._create_schema() + def _alias_annotation_names(self, axis, name): + """ + Do all user-specified annotation aliasing. + + As a *critical* side-effect, ensure the indices are simple number ranges + (accomplished by calling pandas.DataFrame.reset_index()) + """ + if name == 'name': + # a noop, so skip it + return + + ax_name = str(axis) + df_axis = getattr(self.data, ax_name) + if name is None: + # reset index to simple range; alias 'name' to point at the + # previously specified index. + df_axis = df_axis.reset_index().rename(columns={'index': 'name'}) + elif name in df_axis.columns: + if name not in df_axis.columns: + raise KeyError(f"Annotation name {name}, specified in --{ax_name}-name does not exist.") + if not df_axis[name].is_unique: + raise KeyError(f"Values in -{ax_name}-name must be unique. " + "Please prepare data to contain unique values.") + # reset index to simple range; alias user-specified annotation to 'name' + df_axis = df_axis.reset_index(drop=True).rename(columns={name: 'name'}) + else: + raise KeyError(f"Annotation name {name}, specified in --{ax_name}_name does not exist.") + setattr(self.data, ax_name, df_axis) + def _create_schema(self): self.schema = { "dataframe": { @@ -93,15 +122,16 @@ class ScanpyEngine(CXGDriver): """ return np.where(np.isnan(values), 1, values) - def _add_mandatory_annotations(self): - # ensure gene - self.data.var["name"] = Series(list(self.data.var.index), dtype="unicode_", index=self.data.var.index) - self.data.var.index = Series(list(range(self.data.var.shape[0])), dtype="category") - # ensure cell name - self.data.obs["name"] = Series(list(self.data.obs.index), dtype="unicode_", index=self.data.obs.index) - self.data.obs.index = Series(list(range(self.data.obs.shape[0])), dtype="category") + @staticmethod + def _nan_to_zero(values): + """ + Replaces NaN values with 0 + :param values: numpy ndarray + :return: ndarray + """ + return np.where(np.isnan(values), 0, values) - def _validatate_data_types(self): + def _validate_data_types(self): if self.data.X.dtype != "float32": warnings.warn(f"Scanpy data matrix is in {self.data.X.dtype} format not float32. " f"Precision may be truncated.") @@ -118,12 +148,6 @@ class ScanpyEngine(CXGDriver): warnings.warn(f"Scanpy annotation {ax}:{ann} is in unsupported format: {datatype}. " f"Data will be downcast to {downcast_map[datatype]}.") - def cells(self): - return self.data.obs.index.tolist() - - def genes(self): - return self.data.var.index.tolist() - def filter_dataframe(self, filter, include_uns=False): """ Filter cells from data and return a subset of the data. They can operate on both obs and var dimension with @@ -150,12 +174,8 @@ class ScanpyEngine(CXGDriver): genes_idx = self._filter_index(filter["var"]["index"], genes_idx, Axis.VAR) if "annotation_value" in filter["var"]: genes_idx = self._filter_annotation(filter["var"]["annotation_value"], genes_idx, Axis.VAR) - # Due to anndata issues we can't index into cells and genes at the same time - cell_data = self.data[cells_idx, :] - data = cell_data[:, genes_idx] - # TODO: tmp hack to avoid problems with filter that is limited to single gene - if include_uns: - data.uns = cell_data.uns + + data = self._slice(self.data, cells_idx, genes_idx) return data def _filter_index(self, filter, index, axis): @@ -202,6 +222,35 @@ class ScanpyEngine(CXGDriver): index = np.logical_and(index, key_idx) return index + @staticmethod + def _slice(data, obs_selector=None, vars_selector=None): + """ + Slice date using any selector that the AnnData object + supprots for slicing. If selector is None, will not slice + on that axis. + + This method exists to optimize filtering/slicing sparse data that has + access patterns which impact slicing performance. + + https://docs.scipy.org/doc/scipy/reference/sparse.html + """ + prefer_row_access = sparse.isspmatrix_csr(data._X) or \ + sparse.isspmatrix_lil(data._X) or sparse.isspmatrix_bsr(data._X) + if prefer_row_access: + # Row-major slicing + if obs_selector is not None: + data = data[obs_selector, :] + if vars_selector is not None: + data = data[:, vars_selector] + else: + # Col-major slicing + if vars_selector is not None: + data = data[:, vars_selector] + if obs_selector is not None: + data = data[obs_selector, :] + + return data + @cache.memoize() def annotation(self, filter, axis, fields=None): """ @@ -219,11 +268,11 @@ class ScanpyEngine(CXGDriver): df_axis = getattr(df, axis) if not fields: fields = df_axis.columns.tolist() - annotations = DataFrame(df_axis[fields], index=df_axis.index) - return { + result = { "names": fields, - "data": annotations.reset_index().values.tolist() + "data": DataFrame(df_axis[fields]).to_records(index=True).tolist() } + return result @cache.memoize() def data_frame(self, filter, axis): @@ -237,28 +286,20 @@ class ScanpyEngine(CXGDriver): } """ try: - df = self.filter_dataframe(filter) + slice = self.filter_dataframe(filter) except KeyError as e: raise FilterError(f"Error parsing filter: {e}") from e - var_idx = df.var.index.tolist() - obs_idx = df.obs.index.tolist() - values = df.X - df_shape = df.shape - if df_shape[0] == 1: - values = values[None, :] - elif df_shape[1] == 1: - values = values[:, None] + # convert sparse slice to dense + X = slice._X.toarray() if sparse.issparse(slice._X) else slice._X if axis == Axis.OBS: - expression = DataFrame(values, index=obs_idx) result = { - "var": var_idx, - "obs": expression.reset_index().values.tolist() + "var": slice.var.index.tolist(), + "obs": DataFrame(X, index=slice.obs.index).to_records(index=True).tolist() } else: - expression = DataFrame(values.T, index=var_idx) result = { - "obs": obs_idx, - "var": expression.reset_index().values.tolist(), + "obs": slice.obs.index.tolist(), + "var": DataFrame(X.T, index=slice.var.index).to_records(index=True).tolist() } return result @@ -298,14 +339,18 @@ class ScanpyEngine(CXGDriver): top_n = DEFAULT_TOP_N genes_idx = df1.var.index - diffexp_result = stats.ttest_ind(df1.X, df2.X) + # ensure we are using a dense ndarray + X1 = df1._X.toarray() if sparse.issparse(df1._X) else df1._X + X2 = df2._X.toarray() if sparse.issparse(df2._X) else df2._X + diffexp_result = stats.ttest_ind(X1, X2) + tstats = self._nan_to_zero(diffexp_result.statistic) pval = self._nan_to_one(diffexp_result.pvalue) bonferroni_pval = 1 - (1 - pval) ** self.gene_count - ave_exp_set1 = np.mean(df1.X, axis=0) - ave_exp_set2 = np.mean(df2.X, axis=0) + ave_exp_set1 = np.mean(X1, axis=0) + ave_exp_set2 = np.mean(X2, axis=0) ave_diff = ave_exp_set1 - ave_exp_set2 if mode == DiffExpMode.TOP_N: - sort_order = np.argsort(np.abs(diffexp_result.statistic))[::-1] + sort_order = np.argsort(np.abs(tstats))[::-1] # If top_n > length it will just return length genes = self._top_sort(genes_idx, sort_order, top_n) pval = self._top_sort(pval, sort_order, top_n) @@ -350,6 +395,5 @@ class ScanpyEngine(CXGDriver): index=df.obs.index) return { "ndims": normalized_layout.shape[1], - # reset_index gets obs' id into output - "coordinates": normalized_layout.reset_index().values.tolist() + "coordinates": normalized_layout.to_records(index=True).tolist() } diff --git a/server/requirements.txt b/server/requirements.txt index a1ef2b13..652e8dbe 100644 --- a/server/requirements.txt +++ b/server/requirements.txt @@ -1,4 +1,4 @@ -anndata==0.6.11 +anndata>=0.6.12 click==6.7 Flask==0.12.4 Flask-Caching==1.4.0 diff --git a/server/test/schema.json b/server/test/schema.json index 6609f7a2..aa28eb62 100644 --- a/server/test/schema.json +++ b/server/test/schema.json @@ -6,6 +6,10 @@ }, "annotations": { "obs": [ + { + "name": "name", + "type": "string" + }, { "name": "n_genes", "type": "int32" @@ -31,20 +35,16 @@ "Dendritic cells", "Megakaryocytes" ] - }, - { - "name": "name", - "type": "string" } ], "var": [ - { - "name": "n_cells", - "type": "int32" - }, { "name": "name", "type": "string" + }, + { + "name": "n_cells", + "type": "int32" } ] } diff --git a/server/test/test_api.py b/server/test/test_api.py index f0677698..b4c69126 100644 --- a/server/test/test_api.py +++ b/server/test/test_api.py @@ -102,7 +102,7 @@ class EndPoints(unittest.TestCase): result = self.session.get(url) self.assertEqual(result.status_code, HTTPStatus.OK) result_data = result.json() - self.assertEqual(result_data["names"], ["n_genes", "percent_mito", "n_counts", "louvain", "name"]) + self.assertEqual(result_data["names"], ["name", "n_genes", "percent_mito", "n_counts", "louvain"]) self.assertEqual(len(result_data["data"]), 2638) self.assertEqual(len(result_data["data"][0]), 6) @@ -140,7 +140,7 @@ class EndPoints(unittest.TestCase): result = self.session.put(url, json=obs_filter) self.assertEqual(result.status_code, HTTPStatus.OK) result_data = result.json() - self.assertEqual(result_data["names"], ["n_genes", "percent_mito", "n_counts", "louvain", "name"]) + self.assertEqual(result_data["names"], ["name", "n_genes", "percent_mito", "n_counts", "louvain"]) self.assertEqual(len(result_data["data"]), 15) def test_filter_put_annotations_obs(self): @@ -224,7 +224,7 @@ class EndPoints(unittest.TestCase): result = self.session.get(url) self.assertEqual(result.status_code, HTTPStatus.OK) result_data = result.json() - self.assertEqual(result_data["names"], ["n_cells", "name"]) + self.assertEqual(result_data["names"], ["name", "n_cells"]) self.assertEqual(len(result_data["data"]), 1838) self.assertEqual(len(result_data["data"][0]), 3) @@ -260,7 +260,7 @@ class EndPoints(unittest.TestCase): result = self.session.put(url, json=var_filter) self.assertEqual(result.status_code, HTTPStatus.OK) result_data = result.json() - self.assertEqual(result_data["names"], ["n_cells", "name"]) + self.assertEqual(result_data["names"], ["name", "n_cells"]) self.assertEqual(len(result_data["data"]), 2) def test_filter_put_annotations_var(self): diff --git a/server/test/test_scanpy_engine.py b/server/test/test_scanpy_engine.py index 2223747c..09bf0118 100644 --- a/server/test/test_scanpy_engine.py +++ b/server/test/test_scanpy_engine.py @@ -3,6 +3,7 @@ from os import path import pytest import time import unittest +import argparse import numpy as np from pandas import Series @@ -12,7 +13,14 @@ from server.app.scanpy_engine.scanpy_engine import ScanpyEngine class UtilTest(unittest.TestCase): def setUp(self): - self.data = ScanpyEngine("example-dataset/pbmc3k.h5ad", layout_method="umap", diffexp_method="ttest") + args = argparse.Namespace() + args.layout = "umap" + args.diffexp = "ttest" + args.max_category_items = 100 + args.obs_names = None + args.var_names = None + + self.data = ScanpyEngine("example-dataset/pbmc3k.h5ad", args) self.data._create_schema() def test_init(self): @@ -30,7 +38,7 @@ class UtilTest(unittest.TestCase): @pytest.mark.filterwarnings("ignore:Scanpy data matrix") def test_data_type(self): self.data.data.X = self.data.data.X.astype("float64") - self.assertWarns(UserWarning, self.data._validatate_data_types()) + self.assertWarns(UserWarning, self.data._validate_data_types()) def test_filter_idx(self): filter_ = { @@ -130,10 +138,10 @@ class UtilTest(unittest.TestCase): def test_annotations(self): annotations = self.data.annotation(None, "obs") - self.assertEqual(annotations["names"], ["n_genes", "percent_mito", "n_counts", "louvain", "name"]) + self.assertEqual(annotations["names"], ["name", "n_genes", "percent_mito", "n_counts", "louvain"]) self.assertEqual(len(annotations["data"]), 2638) annotations = self.data.annotation(None, "var") - self.assertEqual(annotations["names"], ["n_cells", "name"]) + self.assertEqual(annotations["names"], ["name", "n_cells"]) self.assertEqual(len(annotations["data"]), 1838) def test_annotation_fields(self): @@ -160,10 +168,10 @@ class UtilTest(unittest.TestCase): } } annotations = self.data.annotation(filter_["filter"], "obs") - self.assertEqual(annotations["names"], ["n_genes", "percent_mito", "n_counts", "louvain", "name"]) + self.assertEqual(annotations["names"], ["name", "n_genes", "percent_mito", "n_counts", "louvain"]) self.assertEqual(len(annotations["data"]), 497) annotations = self.data.annotation(filter_["filter"], "var") - self.assertEqual(annotations["names"], ["n_cells", "name"]) + self.assertEqual(annotations["names"], ["name", "n_cells"]) self.assertEqual(len(annotations["data"]), 2) def test_filtered_layout(self): @@ -222,12 +230,12 @@ class UtilTest(unittest.TestCase): data_frame_obs = self.data.data_frame(filter_["filter"], "obs") self.assertEqual(len(data_frame_obs["var"]), 1838) self.assertEqual(len(data_frame_obs["obs"]), 497) - self.assertEqual(type(data_frame_obs["obs"][0]), list) + self.assertIsInstance(data_frame_obs["obs"][0], (list, tuple)) self.assertEqual(type(data_frame_obs["var"][0]), int) data_frame_var = self.data.data_frame(filter_["filter"], "var") self.assertEqual(len(data_frame_var["var"]), 1838) self.assertEqual(len(data_frame_var["obs"]), 497) - self.assertEqual(type(data_frame_var["var"][0]), list) + self.assertIsInstance(data_frame_var["var"][0], (list, tuple)) self.assertEqual(type(data_frame_var["obs"][0]), int) def test_data_single_gene(self): @@ -244,10 +252,10 @@ class UtilTest(unittest.TestCase): data_frame_var = self.data.data_frame(filter_["filter"], axis) if axis == "obs": self.assertEqual(type(data_frame_var["var"][0]), int) - self.assertEqual(type(data_frame_var["obs"][0]), list) + self.assertIsInstance(data_frame_var["obs"][0], (list, tuple)) elif axis == "var": self.assertEqual(type(data_frame_var["obs"][0]), int) - self.assertEqual(type(data_frame_var["var"][0]), list) + self.assertIsInstance(data_frame_var["var"][0], (list, tuple)) if __name__ == '__main__': unittest.main()