Merge pull request #190 from chanzuckerberg/revert-179-csweaver/api-v2-init

Revert "Format loaded dataset"
This commit is contained in:
Charlotte Weaver
2018-08-14 10:43:15 -07:00
committed by GitHub
4 changed files with 96 additions and 28 deletions
+3
View File
@@ -14,3 +14,6 @@ script:
- set -eo pipefail - set -eo pipefail
- flake8 server/app/ - flake8 server/app/
- pytest -s server/test/test_filter.py server/test/test_scanpy_engine.py - pytest -s server/test/test_filter.py server/test/test_scanpy_engine.py
- cellxgene scanpy example-dataset/ &
- for i in {1..90}; do if http :5005/api/v0.1/initialize > /dev/null; then break; else echo "Waiting for server..."; sleep 1; fi; done
- pytest server/test/test_api.py
+4
View File
@@ -10,6 +10,10 @@ class CXGDriver(metaclass=ABCMeta):
def _load_data(data): def _load_data(data):
pass pass
@abstractmethod
def _load_or_infer_schema(data):
pass
@abstractmethod @abstractmethod
def cells(self): def cells(self):
pass pass
+39 -17
View File
@@ -1,21 +1,20 @@
import os import os
import warnings
import numpy as np import numpy as np
from pandas import Series
import scanpy.api as sc import scanpy.api as sc
from scipy import stats from scipy import stats
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
from server.app.util.schema_parse import parse_schema
class ScanpyEngine(CXGDriver): class ScanpyEngine(CXGDriver):
def __init__(self, data, graph_method="umap", diffexp_method="ttest"): def __init__(self, data, schema=None, graph_method="umap", diffexp_method="ttest"):
self.data = self._load_data(data) self.data = self._load_data(data)
self._validatate_data_types() self.schema = self._load_or_infer_schema(data, schema)
self._add_mandatory_annotations() self._set_cell_names()
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.graph_method = graph_method self.graph_method = graph_method
@@ -40,18 +39,41 @@ class ScanpyEngine(CXGDriver):
def _load_data(data): def _load_data(data):
return sc.read(os.path.join(data, "data.h5ad")) return sc.read(os.path.join(data, "data.h5ad"))
def _add_mandatory_annotations(self): def _load_or_infer_schema(self, data, schema):
# ensure gene if not os.path.isfile(os.path.join(data, schema)):
self.data.var["name"] = list(self.data.var.index) # Initialize with cell name which is built off the index
self.data.var.index = Series(list(range(self.data.var.shape[0])), dtype="int32") data_schema = {
# ensure cell name "CellName": {
self.data.obs["name"] = list(self.data.obs.index) "type": "string",
self.data.obs.index = Series(list(range(self.data.obs.shape[0])), dtype="int32") "variabletype": "categorical",
"displayname": "Name",
def _validatate_data_types(self): "include": True
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.") metadata_fields = list(self.data.obs)
for m in metadata_fields:
# Since there are many type of float/int in numpy datatypes the kind attribute of a datatype object
# offers a decent insight into whether it can be lumped in with floats or ints, which is what we
# care about here.
data_kind = self.data.obs[m].dtype.kind
variable_type = "categorical"
data_type = "string"
if data_kind == 'f':
variable_type = "continuous"
data_type = "float"
elif data_kind in ['i', 'u']:
data_type = "int"
if self.data.obs[m].nunique() > 50:
variable_type = "continuous"
data_schema[m] = {
"type": data_type,
"variabletype": variable_type,
"displayname": m,
"include": True
}
else:
data_schema = parse_schema(os.path.join(data, schema))
return data_schema
def cells(self): def cells(self):
return list(self.data.obs.index) return list(self.data.obs.index)
+50 -11
View File
@@ -1,12 +1,11 @@
import unittest import unittest
import pytest
from server.app.scanpy_engine.scanpy_engine import ScanpyEngine 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/") self.data = ScanpyEngine("example-dataset/", schema="data_schema.json")
def test_init(self): def test_init(self):
self.assertEqual(self.data.cell_count, 2638) self.assertEqual(self.data.cell_count, 2638)
@@ -14,16 +13,56 @@ class UtilTest(unittest.TestCase):
epsilon = 0.000005 epsilon = 0.000005
self.assertTrue(self.data.data.X[0,0] - -0.17146951 < epsilon) self.assertTrue(self.data.data.X[0,0] - -0.17146951 < epsilon)
def test_mandatory_annotations(self): def test_schema(self):
self.assertIn("name", self.data.data.obs) self.assertEqual(self.data.schema, {'CellName': {'type': 'string', 'variabletype': 'categorical', 'displayname': 'Name', 'include': True}, 'n_genes': {'type': 'int', 'variabletype': 'continuous', 'displayname': 'Num Genes', 'include': True}, 'percent_mito': {'type': 'float', 'variabletype': 'continuous', 'displayname': 'Mitochondrial Percentage', 'include': True}, 'n_counts': {'type': 'float', 'variabletype': 'continuous', 'displayname': 'Num Counts', 'include': True}, 'louvain': {'type': 'string', 'variabletype': 'categorical', 'displayname': 'Louvain Cluster', 'include': True}})
self.assertEqual(list(self.data.data.obs.index), list(range(2638)))
self.assertIn("name", self.data.data.var)
self.assertEqual(list(self.data.data.var.index), list(range(1838)))
@pytest.mark.filterwarnings("ignore:Scanpy data matrix") def test_cells(self):
def test_data_type(self): cells = self.data.cells()
self.data.data.X = self.data.data.X.astype("float64") self.assertIn("AAACATACAACCAC-1", cells)
self.assertWarns(UserWarning, self.data._validatate_data_types()) self.assertEqual(len(cells), 2638)
def test_genes(self):
genes = self.data.genes()
self.assertIn("SEPT4", genes)
self.assertEqual(len(genes), 1838)
def test_filter_categorical(self):
filter = {"louvain": {"variable_type": "categorical", "value_type": "string", "query": ["B cells"]}}
filtered_data = self.data.filter_cells(filter)
self.assertEqual(filtered_data.shape, (342, 1838))
louvain_vals = filtered_data.obs['louvain'].tolist()
self.assertIn("B cells", louvain_vals)
self.assertNotIn("NK cells", louvain_vals)
def test_filter_continuous(self):
# print(self.data.data.obs["n_genes"].tolist())
filter = {"n_genes": {"variable_type": "continuous", "value_type": "int", "query": {"min": 300, "max": 400}}}
filtered_data = self.data.filter_cells(filter)
self.assertEqual(filtered_data.shape, (71, 1838))
n_genes_vals = filtered_data.obs['n_genes'].tolist()
for val in n_genes_vals:
self.assertTrue(300 <= val <= 400)
def test_metadata(self):
metadata = self.data.metadata(df=self.data.data)
self.assertEqual(len(metadata), 2638)
self.assertIn('louvain', metadata[0])
@unittest.skip("Umap not producing the same graph on different systems, even with the same seed. Skipping for now")
def test_create_graph(self):
graph = self.data.create_graph(df=self.data.data)
self.assertEqual(graph[0][1], 0.5545382653143183)
self.assertEqual(graph[0][2], 0.6021833809031731)
def test_diffexp(self):
diffexp = self.data.diffexp(["AAACATACAACCAC-1", "AACCGATGGTCATG-1"], ["CCGATAGACCTAAG-1", "GGTGGAGAAGTAGA-1"], 0.5, 7)
self.assertEqual(diffexp["celllist1"]["topgenes"], ['EBNA1BP2', 'DIAPH1', 'SLC25A11', 'SNRNP27', 'COMMD8', 'COTL1', 'GTF3A'])
def test_expression(self):
expression = self.data.expression(cells=["AAACATACAACCAC-1"])
data_exp = self.data.data[["AAACATACAACCAC-1"], :].X
for idx in range(len(expression["cells"][0]["e"])):
self.assertEqual(expression["cells"][0]["e"][idx], data_exp[idx])
if __name__ == '__main__': if __name__ == '__main__':