core library (#711)

* move app creation to function

* create engine without load

* flake 8 fixes

* cleanup original scanpy test

* add default config

* handle missing data

* test data changes

* unify update

* load data isn't static anymore

* make app a class
This commit is contained in:
Charlotte Weaver
2019-04-22 12:24:07 -07:00
committed by GitHub
parent 1f735abe2b
commit 9c6273eb94
10 changed files with 171 additions and 60 deletions
+21 -19
View File
@@ -5,27 +5,29 @@ from flask_caching import Cache
from flask_compress import Compress from flask_compress import Compress
from flask_cors import CORS from flask_cors import CORS
from .rest_api.rest import get_api_resources from server.app.rest_api.rest import get_api_resources
from .util.utils import Float32JSONEncoder from server.app.util.utils import Float32JSONEncoder
from .web import webapp from server.app.web import webapp
REACTIVE_LIMIT = 1_000_000
app = Flask(__name__, static_folder="web/static") class Server:
app.json_encoder = Float32JSONEncoder def __init__(self):
cache = Cache(app, config={"CACHE_TYPE": "simple", "CACHE_DEFAULT_TIMEOUT": 860_000}) self.data = None
Compress(app) self.cache = Cache(config={"CACHE_TYPE": "simple", "CACHE_DEFAULT_TIMEOUT": 860_000})
CORS(app)
# Config def create_app(self):
SECRET_KEY = os.environ.get("CXG_SECRET_KEY", default="SparkleAndShine") app = Flask(__name__, static_folder="web/static")
app.json_encoder = Float32JSONEncoder
self.cache.init_app(app)
Compress(app)
CORS(app)
app.config.update(SECRET_KEY=SECRET_KEY) # Config
SECRET_KEY = os.environ.get("CXG_SECRET_KEY", default="SparkleAndShine")
app.config.update(SECRET_KEY=SECRET_KEY)
# Application Data resources = get_api_resources()
data = None app.register_blueprint(webapp.bp)
app.register_blueprint(resources.blueprint)
resources = get_api_resources() app.add_url_rule("/", endpoint="index")
app.register_blueprint(webapp.bp) return app
app.register_blueprint(resources.blueprint)
app.add_url_rule("/", endpoint="index")
+24 -13
View File
@@ -11,13 +11,27 @@ Sort order for methods
class CXGDriver(metaclass=ABCMeta): class CXGDriver(metaclass=ABCMeta):
def __init__(self, data, args): def __init__(self, data=None, args={}):
self.data = self._load_data(data) self.config = self._get_default_config()
self.layout_method = args["layout"] self.config.update(args)
self.diffexp_method = args["diffexp"] if data:
self.max_category_items = args["max_category_items"] self._load_data(data)
self.diffexp_lfc_cutoff = args["diffexp_lfc_cutoff"] else:
self.cluster = None self.data = None
def update(self, data=None, args={}):
self.config.update(args)
if data:
self._load_data(data)
@staticmethod
def _get_default_config():
return {
"layout": None,
"diffexp": None,
"max_category_items": None,
"diffexp_lfc_cutoff": None
}
@property @property
def features(self): def features(self):
@@ -27,18 +41,15 @@ class CXGDriver(metaclass=ABCMeta):
"diffexp": {"available": False}, "diffexp": {"available": False},
} }
# TODO - Interactive limit should be generated from the actual available methods see GH issue #94 # TODO - Interactive limit should be generated from the actual available methods see GH issue #94
if self.layout_method: if self.config["layout"]:
# TODO handle "var" when gene layout becomes available # TODO handle "var" when gene layout becomes available
features["layout"]["obs"] = {"available": True, "interactiveLimit": 50000} features["layout"]["obs"] = {"available": True, "interactiveLimit": 50000}
if self.diffexp_method: if self.config["diffexp"]:
features["diffexp"] = {"available": True, "interactiveLimit": 50000} features["diffexp"] = {"available": True, "interactiveLimit": 50000}
if self.cluster:
features["cluster"] = {"available": True, "interactiveLimit": 50000}
return features return features
@staticmethod
@abstractmethod @abstractmethod
def _load_data(data): def _load_data(self, data):
pass pass
@abstractmethod @abstractmethod
+1 -1
View File
@@ -63,7 +63,7 @@ class ConfigAPI(Resource):
"dataset": current_app.config["DATASET_TITLE"], "dataset": current_app.config["DATASET_TITLE"],
}, },
"parameters": { "parameters": {
"max_category_items": current_app.data.max_category_items "max_category_items": current_app.data.config["max_category_items"]
}, },
"library_versions": { "library_versions": {
"scanpy": pkg_resources.get_distribution("scanpy").version, "scanpy": pkg_resources.get_distribution("scanpy").version,
+48 -23
View File
@@ -12,7 +12,7 @@ from server.app.util.errors import (
PrepareError, PrepareError,
ScanpyFileError, ScanpyFileError,
) )
from server.app.util.utils import jsonify_scanpy from server.app.util.utils import jsonify_scanpy, requires_data
from server.app.scanpy_engine.diffexp import diffexp_ttest from server.app.scanpy_engine.diffexp import diffexp_ttest
from server.app.util.fbs.matrix import encode_matrix_fbs from server.app.util.fbs.matrix import encode_matrix_fbs
@@ -27,17 +27,26 @@ Sort order for methods
class ScanpyEngine(CXGDriver): class ScanpyEngine(CXGDriver):
def __init__(self, data, args): def __init__(self, data=None, args={}):
super().__init__(data, args) super().__init__(data, args)
self._alias_annotation_names(Axis.OBS, args["obs_names"]) if self.data:
self._alias_annotation_names(Axis.VAR, args["var_names"]) self._validate_and_initialize()
self._validate_data_types()
self._validate_data_calculations() def update(self, data=None, args={}):
self.cell_count = self.data.shape[0] super().__init__(data, args)
self.gene_count = self.data.shape[1] if self.data:
self.layout_options = ["umap", "tsne"] self._validate_and_initialize()
self.diffexp_options = ["ttest"]
self._create_schema() @staticmethod
def _get_default_config():
return {
"layout": "umap",
"diffexp": "ttest",
"max_category_items": 100,
"obs_names": None,
"var_names": None,
"diffexp_lfc_cutoff": 0.01,
}
def _alias_annotation_names(self, axis, name): def _alias_annotation_names(self, axis, name):
""" """
@@ -95,6 +104,7 @@ class ScanpyEngine(CXGDriver):
return True return True
return False return False
@requires_data
def _create_schema(self): def _create_schema(self):
self.schema = { self.schema = {
"dataframe": { "dataframe": {
@@ -128,13 +138,12 @@ class ScanpyEngine(CXGDriver):
) )
self.schema["annotations"][ax].append(ann_schema) self.schema["annotations"][ax].append(ann_schema)
@staticmethod def _load_data(self, data):
def _load_data(data):
# Based on benchmarking, cache=True has no impact on perf. # Based on benchmarking, cache=True has no impact on perf.
# Note: as of current scanpy/anndata release, setting backed='r' will # Note: as of current scanpy/anndata release, setting backed='r' will
# result in an error. https://github.com/theislab/anndata/issues/79 # result in an error. https://github.com/theislab/anndata/issues/79
try: try:
result = sc.read(data, cache=True) self.data = sc.read(data, cache=True)
except ValueError: except ValueError:
raise ScanpyFileError( raise ScanpyFileError(
"File must be in the .h5ad format. Please read " "File must be in the .h5ad format. Please read "
@@ -151,8 +160,18 @@ class ScanpyEngine(CXGDriver):
f"Error while loading file: {e}, File must be in the .h5ad format, please check " f"Error while loading file: {e}, File must be in the .h5ad format, please check "
f"that your input and try again." f"that your input and try again."
) )
return result
@requires_data
def _validate_and_initialize(self):
self._alias_annotation_names(Axis.OBS, self.config["obs_names"])
self._alias_annotation_names(Axis.VAR, self.config["var_names"])
self._validate_data_types()
self._validate_data_calculations()
self.cell_count = self.data.shape[0]
self.gene_count = self.data.shape[1]
self._create_schema()
@requires_data
def _validate_data_types(self): def _validate_data_types(self):
if self.data.X.dtype != "float32": if self.data.X.dtype != "float32":
warnings.warn( warnings.warn(
@@ -176,7 +195,7 @@ class ScanpyEngine(CXGDriver):
) )
if isinstance(datatype, CategoricalDtype): if isinstance(datatype, CategoricalDtype):
category_num = len(curr_axis[ann].dtype.categories) category_num = len(curr_axis[ann].dtype.categories)
if category_num > 500 and category_num > self.max_category_items: if category_num > 500 and category_num > self.config['max_category_items']:
warnings.warn( warnings.warn(
f"{str(ax).title()} annotation '{ann}' has {category_num} categories, this may be " f"{str(ax).title()} annotation '{ann}' has {category_num} categories, this may be "
f"cumbersome or slow to display. We recommend setting the " f"cumbersome or slow to display. We recommend setting the "
@@ -184,16 +203,17 @@ class ScanpyEngine(CXGDriver):
f"annotations with more than 500 categories in the UI" f"annotations with more than 500 categories in the UI"
) )
@requires_data
def _validate_data_calculations(self): def _validate_data_calculations(self):
layout_key = f"X_{self.layout_method}" layout_key = f"X_{self.config['layout']}"
try: try:
assert layout_key in self.data.obsm_keys() assert layout_key in self.data.obsm_keys()
except AssertionError: except AssertionError:
raise PrepareError( raise PrepareError(
f"Cannot find a field with coordinates for the {self.layout_method} layout requested. A different" f"Cannot find a field with coordinates for the {self.config['layout']} layout requested. A different"
f" layout may have been computed. The requested layout must be pre-calculated and saved " f" layout may have been computed. The requested layout must be pre-calculated and saved "
f"back in the h5ad file. You can run " f"back in the h5ad file. You can run "
f"`cellxgene prepare --layout {self.layout_method} <datafile>` " f"`cellxgene prepare --layout {self.config['layout']} <datafile>` "
f"to solve this problem. " f"to solve this problem. "
) )
@@ -220,7 +240,7 @@ class ScanpyEngine(CXGDriver):
mask = np.zeros((count,), dtype=bool) mask = np.zeros((count,), dtype=bool)
for i in filter: for i in filter:
if type(i) == list: if type(i) == list:
mask[i[0] : i[1]] = True mask[i[0]: i[1]] = True
else: else:
mask[i] = True mask[i] = True
return mask return mask
@@ -241,6 +261,7 @@ class ScanpyEngine(CXGDriver):
) )
return mask return mask
@requires_data
def _filter_to_mask(self, filter, use_slices=True): def _filter_to_mask(self, filter, use_slices=True):
if use_slices: if use_slices:
obs_selector = slice(0, self.data.n_obs) obs_selector = slice(0, self.data.n_obs)
@@ -260,6 +281,7 @@ class ScanpyEngine(CXGDriver):
) )
return obs_selector, var_selector return obs_selector, var_selector
@requires_data
def annotation_to_fbs_matrix(self, axis, fields=None): def annotation_to_fbs_matrix(self, axis, fields=None):
if axis == Axis.OBS: if axis == Axis.OBS:
df = self.data.obs df = self.data.obs
@@ -269,6 +291,7 @@ class ScanpyEngine(CXGDriver):
df = df[fields] df = df[fields]
return encode_matrix_fbs(df, col_idx=df.columns) return encode_matrix_fbs(df, col_idx=df.columns)
@requires_data
def data_frame_to_fbs_matrix(self, filter, axis): def data_frame_to_fbs_matrix(self, filter, axis):
""" """
Retrieves data 'X' and returns in a flatbuffer Matrix. Retrieves data 'X' and returns in a flatbuffer Matrix.
@@ -295,6 +318,7 @@ class ScanpyEngine(CXGDriver):
X = X[:, var_selector] X = X[:, var_selector]
return encode_matrix_fbs(X, col_idx=np.nonzero(var_selector)[0], row_idx=None) return encode_matrix_fbs(X, col_idx=np.nonzero(var_selector)[0], row_idx=None)
@requires_data
def diffexp_topN(self, obsFilterA, obsFilterB, top_n=None, interactive_limit=None): def diffexp_topN(self, obsFilterA, obsFilterB, top_n=None, interactive_limit=None):
if Axis.VAR in obsFilterA or Axis.VAR in obsFilterB: if Axis.VAR in obsFilterA or Axis.VAR in obsFilterB:
raise FilterError("Observation filters may not contain vaiable conditions") raise FilterError("Observation filters may not contain vaiable conditions")
@@ -310,7 +334,7 @@ class ScanpyEngine(CXGDriver):
if top_n is None: if top_n is None:
top_n = DEFAULT_TOP_N top_n = DEFAULT_TOP_N
result = diffexp_ttest( result = diffexp_ttest(
self.data, obs_mask_A, obs_mask_B, top_n, self.diffexp_lfc_cutoff self.data, obs_mask_A, obs_mask_B, top_n, self.config['diffexp_lfc_cutoff']
) )
try: try:
return jsonify_scanpy(result) return jsonify_scanpy(result)
@@ -319,6 +343,7 @@ class ScanpyEngine(CXGDriver):
"Error encoding differential expression to JSON" "Error encoding differential expression to JSON"
) )
@requires_data
def layout_to_fbs_matrix(self): def layout_to_fbs_matrix(self):
""" """
Return the default 2-D layout for cells as a FBS Matrix. Return the default 2-D layout for cells as a FBS Matrix.
@@ -328,14 +353,14 @@ class ScanpyEngine(CXGDriver):
* only returns Matrix in columnar layout * only returns Matrix in columnar layout
""" """
try: try:
full_embedding = self.data.obsm[f"X_{self.layout_method}"] full_embedding = self.data.obsm[f"X_{self.config['layout']}"]
if full_embedding.shape[1] > 2: if full_embedding.shape[1] > 2:
warnings.warn(f"Warning: found {full_embedding.shape[1]} \ warnings.warn(f"Warning: found {full_embedding.shape[1]} \
components of embedding. Using the first two for layout display.") components of embedding. Using the first two for layout display.")
df_layout = full_embedding[:, :2] df_layout = full_embedding[:, :2]
except ValueError as e: except ValueError as e:
raise PrepareError( raise PrepareError(
f"Layout has not been calculated using {self.layout_method}, " f"Layout has not been calculated using {self.config['layout']}, "
f"please prepare your datafile and relaunch cellxgene") from e f"please prepare your datafile and relaunch cellxgene") from e
normalized_layout = (df_layout - df_layout.min()) / (df_layout.max() - df_layout.min()) normalized_layout = (df_layout - df_layout.min()) / (df_layout.max() - df_layout.min())
+1
View File
@@ -30,3 +30,4 @@ class DiffExpMode(AugmentedEnum):
JSON_NaN_to_num_warning_msg = ( JSON_NaN_to_num_warning_msg = (
"JSON encoding failure - please verify all data are finite values (no NaN or Infinities)" "JSON encoding failure - please verify all data are finite values (no NaN or Infinities)"
) )
REACTIVE_LIMIT = 1_000_000
+9
View File
@@ -50,3 +50,12 @@ class ScanpyFileError(Exception):
def __init__(self, message): def __init__(self, message):
self.message = message self.message = message
class DriverError(Exception):
"""
Raised when file loaded into scanpy is misformatted
"""
def __init__(self, message):
self.message = message
+13
View File
@@ -1,6 +1,10 @@
from functools import wraps
from flask import json from flask import json
from numpy import float32, integer from numpy import float32, integer
from server.app.util.errors import DriverError
class Float32JSONEncoder(json.JSONEncoder): class Float32JSONEncoder(json.JSONEncoder):
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
@@ -28,3 +32,12 @@ def custom_format_warning(msg, *args, **kwargs):
def jsonify_scanpy(data): def jsonify_scanpy(data):
return json.dumps(data, cls=Float32JSONEncoder, allow_nan=False) return json.dumps(data, cls=Float32JSONEncoder, allow_nan=False)
def requires_data(func):
@wraps(func)
def wrapped_function(self, *args, **kwargs):
if self.data is None:
raise DriverError(f"error data must be loaded before you call {func.__name__}")
return func(self, *args, **kwargs)
return wrapped_function
+3 -1
View File
@@ -7,6 +7,7 @@ import webbrowser
import click import click
from server.app.app import Server
from server.app.util.errors import ScanpyFileError from server.app.util.errors import ScanpyFileError
from server.app.util.utils import custom_format_warning from server.app.util.utils import custom_format_warning
@@ -116,8 +117,9 @@ def launch(
cellxgene_url = f"http://{host}:{port}" cellxgene_url = f"http://{host}:{port}"
# Import Flask app # Import Flask app
from server.app.app import app server = Server()
app = server.create_app()
app.config.update(DATASET_TITLE=title) app.config.update(DATASET_TITLE=title)
if not verbose: if not verbose:
+1 -3
View File
@@ -12,7 +12,7 @@ from server.app.scanpy_engine.scanpy_engine import ScanpyEngine
from server.app.util.errors import FilterError from server.app.util.errors import FilterError
class UtilTest(unittest.TestCase): class EngineTest(unittest.TestCase):
def setUp(self): def setUp(self):
args = { args = {
"layout": "umap", "layout": "umap",
@@ -22,9 +22,7 @@ class UtilTest(unittest.TestCase):
"var_names": None, "var_names": None,
"diffexp_lfc_cutoff": 0.01, "diffexp_lfc_cutoff": 0.01,
} }
self.data = ScanpyEngine("example-dataset/pbmc3k.h5ad", args) self.data = ScanpyEngine("example-dataset/pbmc3k.h5ad", args)
self.data._create_schema()
def test_init(self): def test_init(self):
self.assertEqual(self.data.cell_count, 2638) self.assertEqual(self.data.cell_count, 2638)
@@ -0,0 +1,50 @@
import unittest
import json
from server.app.scanpy_engine.scanpy_engine import ScanpyEngine
from server.app.util.errors import DriverError
class DataLoadEngineTest(unittest.TestCase):
def setUp(self):
self.data_file = "example-dataset/pbmc3k.h5ad"
self.data = ScanpyEngine()
def test_init(self):
self.assertIsNone(self.data.data)
def test_delayed_load_args(self):
args = {
"layout": "tsne",
"diffexp": "ttest",
"max_category_items": 1000,
"obs_names": "foo",
"var_names": "bar",
"diffexp_lfc_cutoff": 0.1,
}
self.data.update(args=args)
self.assertEqual(args, self.data.config)
def test_requires_data(self):
with self.assertRaises(DriverError):
self.data._create_schema()
def test_delayed_load_data(self):
self.data.update(data=self.data_file)
self.data._create_schema()
self.assertEqual(self.data.cell_count, 2638)
self.assertEqual(self.data.gene_count, 1838)
epsilon = 0.000_005
self.assertTrue(self.data.data.X[0, 0] - -0.171_469_51 < epsilon)
def test_diffexp_topN(self):
self.data.update(data=self.data_file)
f1 = {"filter": {"obs": {"index": [[0, 500]]}}}
f2 = {"filter": {"obs": {"index": [[500, 1000]]}}}
result = json.loads(self.data.diffexp_topN(f1["filter"], f2["filter"]))
self.assertEqual(len(result), 10)
result = json.loads(self.data.diffexp_topN(f1["filter"], f2["filter"], 20))
self.assertEqual(len(result), 20)
if __name__ == "__main__":
unittest.main()