mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-10-07 13:58:11 +08:00
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:
+21
-19
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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())
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user