add --obs-names and --var-names CLI params (#371)

* add --obs-names and --var-names CLI params

* fix lint

* performance improvements in scanpy engine

* fix lint

* fix typo

* correctly handle sparse formats in diffexp

* fix diffexp and 1d slicing

* diffexp uses t-stat, not pval; clean up arg handling

* make _slice a static method

* revise scanpy tests to match new API
This commit is contained in:
Bruce Martin
2018-10-30 14:07:38 -07:00
committed by GitHub
parent 01369801fc
commit b181751493
7 changed files with 133 additions and 88 deletions
+3 -2
View File
@@ -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("--debug", action="store_true", help=argparse.SUPPRESS)
launch_group.add_argument("--no-open", help="do not launch the webbrowser", action="store_false", launch_group.add_argument("--no-open", help="do not launch the webbrowser", action="store_false",
dest="open_browser") 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( launch_group.add_argument(
"--max-category-items", "--max-category-items",
type=whole_number, type=whole_number,
@@ -141,8 +143,7 @@ def run_scanpy(args):
log.setLevel(logging.ERROR) log.setLevel(logging.ERROR)
from .scanpy_engine.scanpy_engine import ScanpyEngine from .scanpy_engine.scanpy_engine import ScanpyEngine
print(f"Loading data from {args.data} (this may take a while)") 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, app.data = ScanpyEngine(args.data, args)
max_category_items=args.max_category_items)
print(f"Launching cellxgene") print(f"Launching cellxgene")
if args.open_browser: if args.open_browser:
webbrowser.open(cellxgene_url) webbrowser.open(cellxgene_url)
+4 -12
View File
@@ -12,11 +12,11 @@ Sort order for methods
class CXGDriver(metaclass=ABCMeta): 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.data = self._load_data(data)
self.layout_method = layout_method self.layout_method = args.layout
self.diffexp_method = diffexp_method self.diffexp_method = args.diffexp
self.max_category_items = max_category_items self.max_category_items = args.max_category_items
self.cluster = None self.cluster = None
@property @property
@@ -44,14 +44,6 @@ class CXGDriver(metaclass=ABCMeta):
def _load_data(data): def _load_data(data):
pass pass
@abstractmethod
def cells(self):
pass
@abstractmethod
def genes(self):
pass
@abstractmethod @abstractmethod
def filter_dataframe(self, filter): def filter_dataframe(self, filter):
""" """
+95 -51
View File
@@ -1,9 +1,9 @@
import warnings import warnings
import numpy as np import numpy as np
from pandas import DataFrame, Series from pandas import DataFrame
import scanpy.api as sc import scanpy.api as sc
from scipy import stats from scipy import stats, sparse
from server.app.app import cache from server.app.app import cache
from server.app.driver.driver import CXGDriver from server.app.driver.driver import CXGDriver
@@ -22,17 +22,46 @@ Sort order for methods
class ScanpyEngine(CXGDriver): class ScanpyEngine(CXGDriver):
def __init__(self, data, layout_method=None, diffexp_method=None, max_category_items=100): def __init__(self, data, args):
super().__init__(data, layout_method=layout_method, diffexp_method=diffexp_method, super().__init__(data, args)
max_category_items=max_category_items) self._alias_annotation_names(Axis.OBS, args.obs_names)
self._validatate_data_types() self._alias_annotation_names(Axis.VAR, args.var_names)
self._add_mandatory_annotations() self._validate_data_types()
self.cell_count = self.data.shape[0] self.cell_count = self.data.shape[0]
self.gene_count = self.data.shape[1] self.gene_count = self.data.shape[1]
self.layout_options = ["umap", "tsne"] self.layout_options = ["umap", "tsne"]
self.diffexp_options = ["ttest"] self.diffexp_options = ["ttest"]
self._create_schema() 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): def _create_schema(self):
self.schema = { self.schema = {
"dataframe": { "dataframe": {
@@ -93,15 +122,16 @@ class ScanpyEngine(CXGDriver):
""" """
return np.where(np.isnan(values), 1, values) return np.where(np.isnan(values), 1, values)
def _add_mandatory_annotations(self): @staticmethod
# ensure gene def _nan_to_zero(values):
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") Replaces NaN values with 0
# ensure cell name :param values: numpy ndarray
self.data.obs["name"] = Series(list(self.data.obs.index), dtype="unicode_", index=self.data.obs.index) :return: ndarray
self.data.obs.index = Series(list(range(self.data.obs.shape[0])), dtype="category") """
return np.where(np.isnan(values), 0, values)
def _validatate_data_types(self): def _validate_data_types(self):
if self.data.X.dtype != "float32": if self.data.X.dtype != "float32":
warnings.warn(f"Scanpy data matrix is in {self.data.X.dtype} format not float32. " warnings.warn(f"Scanpy data matrix is in {self.data.X.dtype} format not float32. "
f"Precision may be truncated.") f"Precision may be truncated.")
@@ -118,12 +148,6 @@ class ScanpyEngine(CXGDriver):
warnings.warn(f"Scanpy annotation {ax}:{ann} is in unsupported format: {datatype}. " warnings.warn(f"Scanpy annotation {ax}:{ann} is in unsupported format: {datatype}. "
f"Data will be downcast to {downcast_map[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): 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 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) genes_idx = self._filter_index(filter["var"]["index"], genes_idx, Axis.VAR)
if "annotation_value" in filter["var"]: if "annotation_value" in filter["var"]:
genes_idx = self._filter_annotation(filter["var"]["annotation_value"], genes_idx, Axis.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 = self._slice(self.data, cells_idx, genes_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
return data return data
def _filter_index(self, filter, index, axis): def _filter_index(self, filter, index, axis):
@@ -202,6 +222,35 @@ class ScanpyEngine(CXGDriver):
index = np.logical_and(index, key_idx) index = np.logical_and(index, key_idx)
return index 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() @cache.memoize()
def annotation(self, filter, axis, fields=None): def annotation(self, filter, axis, fields=None):
""" """
@@ -219,11 +268,11 @@ class ScanpyEngine(CXGDriver):
df_axis = getattr(df, axis) df_axis = getattr(df, axis)
if not fields: if not fields:
fields = df_axis.columns.tolist() fields = df_axis.columns.tolist()
annotations = DataFrame(df_axis[fields], index=df_axis.index) result = {
return {
"names": fields, "names": fields,
"data": annotations.reset_index().values.tolist() "data": DataFrame(df_axis[fields]).to_records(index=True).tolist()
} }
return result
@cache.memoize() @cache.memoize()
def data_frame(self, filter, axis): def data_frame(self, filter, axis):
@@ -237,28 +286,20 @@ class ScanpyEngine(CXGDriver):
} }
""" """
try: try:
df = self.filter_dataframe(filter) slice = self.filter_dataframe(filter)
except KeyError as e: except KeyError as e:
raise FilterError(f"Error parsing filter: {e}") from e raise FilterError(f"Error parsing filter: {e}") from e
var_idx = df.var.index.tolist() # convert sparse slice to dense
obs_idx = df.obs.index.tolist() X = slice._X.toarray() if sparse.issparse(slice._X) else slice._X
values = df.X
df_shape = df.shape
if df_shape[0] == 1:
values = values[None, :]
elif df_shape[1] == 1:
values = values[:, None]
if axis == Axis.OBS: if axis == Axis.OBS:
expression = DataFrame(values, index=obs_idx)
result = { result = {
"var": var_idx, "var": slice.var.index.tolist(),
"obs": expression.reset_index().values.tolist() "obs": DataFrame(X, index=slice.obs.index).to_records(index=True).tolist()
} }
else: else:
expression = DataFrame(values.T, index=var_idx)
result = { result = {
"obs": obs_idx, "obs": slice.obs.index.tolist(),
"var": expression.reset_index().values.tolist(), "var": DataFrame(X.T, index=slice.var.index).to_records(index=True).tolist()
} }
return result return result
@@ -298,14 +339,18 @@ class ScanpyEngine(CXGDriver):
top_n = DEFAULT_TOP_N top_n = DEFAULT_TOP_N
genes_idx = df1.var.index 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) pval = self._nan_to_one(diffexp_result.pvalue)
bonferroni_pval = 1 - (1 - pval) ** self.gene_count bonferroni_pval = 1 - (1 - pval) ** self.gene_count
ave_exp_set1 = np.mean(df1.X, axis=0) ave_exp_set1 = np.mean(X1, axis=0)
ave_exp_set2 = np.mean(df2.X, axis=0) ave_exp_set2 = np.mean(X2, axis=0)
ave_diff = ave_exp_set1 - ave_exp_set2 ave_diff = ave_exp_set1 - ave_exp_set2
if mode == DiffExpMode.TOP_N: 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 # If top_n > length it will just return length
genes = self._top_sort(genes_idx, sort_order, top_n) genes = self._top_sort(genes_idx, sort_order, top_n)
pval = self._top_sort(pval, sort_order, top_n) pval = self._top_sort(pval, sort_order, top_n)
@@ -350,6 +395,5 @@ class ScanpyEngine(CXGDriver):
index=df.obs.index) index=df.obs.index)
return { return {
"ndims": normalized_layout.shape[1], "ndims": normalized_layout.shape[1],
# reset_index gets obs' id into output "coordinates": normalized_layout.to_records(index=True).tolist()
"coordinates": normalized_layout.reset_index().values.tolist()
} }
+1 -1
View File
@@ -1,4 +1,4 @@
anndata==0.6.11 anndata>=0.6.12
click==6.7 click==6.7
Flask==0.12.4 Flask==0.12.4
Flask-Caching==1.4.0 Flask-Caching==1.4.0
+8 -8
View File
@@ -6,6 +6,10 @@
}, },
"annotations": { "annotations": {
"obs": [ "obs": [
{
"name": "name",
"type": "string"
},
{ {
"name": "n_genes", "name": "n_genes",
"type": "int32" "type": "int32"
@@ -31,20 +35,16 @@
"Dendritic cells", "Dendritic cells",
"Megakaryocytes" "Megakaryocytes"
] ]
},
{
"name": "name",
"type": "string"
} }
], ],
"var": [ "var": [
{
"name": "n_cells",
"type": "int32"
},
{ {
"name": "name", "name": "name",
"type": "string" "type": "string"
},
{
"name": "n_cells",
"type": "int32"
} }
] ]
} }
+4 -4
View File
@@ -102,7 +102,7 @@ class EndPoints(unittest.TestCase):
result = self.session.get(url) result = self.session.get(url)
self.assertEqual(result.status_code, HTTPStatus.OK) self.assertEqual(result.status_code, HTTPStatus.OK)
result_data = result.json() 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"]), 2638)
self.assertEqual(len(result_data["data"][0]), 6) self.assertEqual(len(result_data["data"][0]), 6)
@@ -140,7 +140,7 @@ class EndPoints(unittest.TestCase):
result = self.session.put(url, json=obs_filter) result = self.session.put(url, json=obs_filter)
self.assertEqual(result.status_code, HTTPStatus.OK) self.assertEqual(result.status_code, HTTPStatus.OK)
result_data = result.json() 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) self.assertEqual(len(result_data["data"]), 15)
def test_filter_put_annotations_obs(self): def test_filter_put_annotations_obs(self):
@@ -224,7 +224,7 @@ class EndPoints(unittest.TestCase):
result = self.session.get(url) result = self.session.get(url)
self.assertEqual(result.status_code, HTTPStatus.OK) self.assertEqual(result.status_code, HTTPStatus.OK)
result_data = result.json() 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"]), 1838)
self.assertEqual(len(result_data["data"][0]), 3) self.assertEqual(len(result_data["data"][0]), 3)
@@ -260,7 +260,7 @@ class EndPoints(unittest.TestCase):
result = self.session.put(url, json=var_filter) result = self.session.put(url, json=var_filter)
self.assertEqual(result.status_code, HTTPStatus.OK) self.assertEqual(result.status_code, HTTPStatus.OK)
result_data = result.json() 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) self.assertEqual(len(result_data["data"]), 2)
def test_filter_put_annotations_var(self): def test_filter_put_annotations_var(self):
+18 -10
View File
@@ -3,6 +3,7 @@ from os import path
import pytest import pytest
import time import time
import unittest import unittest
import argparse
import numpy as np import numpy as np
from pandas import Series from pandas import Series
@@ -12,7 +13,14 @@ from server.app.scanpy_engine.scanpy_engine import ScanpyEngine
class UtilTest(unittest.TestCase): class UtilTest(unittest.TestCase):
def setUp(self): 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() self.data._create_schema()
def test_init(self): def test_init(self):
@@ -30,7 +38,7 @@ class UtilTest(unittest.TestCase):
@pytest.mark.filterwarnings("ignore:Scanpy data matrix") @pytest.mark.filterwarnings("ignore:Scanpy data matrix")
def test_data_type(self): def test_data_type(self):
self.data.data.X = self.data.data.X.astype("float64") 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): def test_filter_idx(self):
filter_ = { filter_ = {
@@ -130,10 +138,10 @@ class UtilTest(unittest.TestCase):
def test_annotations(self): def test_annotations(self):
annotations = self.data.annotation(None, "obs") 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) self.assertEqual(len(annotations["data"]), 2638)
annotations = self.data.annotation(None, "var") 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) self.assertEqual(len(annotations["data"]), 1838)
def test_annotation_fields(self): def test_annotation_fields(self):
@@ -160,10 +168,10 @@ class UtilTest(unittest.TestCase):
} }
} }
annotations = self.data.annotation(filter_["filter"], "obs") 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) self.assertEqual(len(annotations["data"]), 497)
annotations = self.data.annotation(filter_["filter"], "var") 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) self.assertEqual(len(annotations["data"]), 2)
def test_filtered_layout(self): def test_filtered_layout(self):
@@ -222,12 +230,12 @@ class UtilTest(unittest.TestCase):
data_frame_obs = self.data.data_frame(filter_["filter"], "obs") data_frame_obs = self.data.data_frame(filter_["filter"], "obs")
self.assertEqual(len(data_frame_obs["var"]), 1838) self.assertEqual(len(data_frame_obs["var"]), 1838)
self.assertEqual(len(data_frame_obs["obs"]), 497) 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) self.assertEqual(type(data_frame_obs["var"][0]), int)
data_frame_var = self.data.data_frame(filter_["filter"], "var") data_frame_var = self.data.data_frame(filter_["filter"], "var")
self.assertEqual(len(data_frame_var["var"]), 1838) self.assertEqual(len(data_frame_var["var"]), 1838)
self.assertEqual(len(data_frame_var["obs"]), 497) 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) self.assertEqual(type(data_frame_var["obs"][0]), int)
def test_data_single_gene(self): def test_data_single_gene(self):
@@ -244,10 +252,10 @@ class UtilTest(unittest.TestCase):
data_frame_var = self.data.data_frame(filter_["filter"], axis) data_frame_var = self.data.data_frame(filter_["filter"], axis)
if axis == "obs": if axis == "obs":
self.assertEqual(type(data_frame_var["var"][0]), int) 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": elif axis == "var":
self.assertEqual(type(data_frame_var["obs"][0]), int) 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__': if __name__ == '__main__':
unittest.main() unittest.main()