Splitting the format validation and mandatory annotations

This commit is contained in:
Charlotte Weaver
2018-08-10 16:38:07 -07:00
parent 58000c1815
commit 55fa9b892a
2 changed files with 11 additions and 58 deletions
-4
View File
@@ -10,10 +10,6 @@ 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
+11 -54
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, schema=None, graph_method="umap", diffexp_method="ttest"): def __init__(self, data, graph_method="umap", diffexp_method="ttest"):
self.data = self._load_data(data) self.data = self._load_data(data)
self._format_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,59 +39,17 @@ 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"))
# TODO delete after v0.1 v2.0 transition def _add_mandatory_annotations(self):
def _load_or_infer_schema(self, data, schema):
if not os.path.isfile(os.path.join(data, schema)):
# Initialize with cell name which is built off the index
data_schema = {
"CellName": {
"type": "string",
"variabletype": "categorical",
"displayname": "Name",
"include": True
}
}
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 _format_data(self):
# ensure gene # ensure gene
self.data.var["name"] = list(self.data.var.index) self.data.var["name"] = list(self.data.var.index)
self.data.var.index = list(range(self.data.var.shape[0])) self.data.var.index = Series(list(range(self.data.var.shape[0])), dtype='int32')
# ensure cell name # ensure cell name
self.data.obs["name"] = list(self.data.obs.index) self.data.obs["name"] = list(self.data.obs.index)
self.data.obs.index = list(range(self.data.obs.shape[0])) self.data.obs.index = Series(list(range(self.data.obs.shape[0])), dtype='int32')
# ensure formats correct/ reform old format
for annotation in self.data.obs: def _validatate_data_types(self):
data_type = self.data.obs[annotation].dtype
if data_type.kind == 'f' and data_type != "float32":
self.data.obs[annotation] = self.data.obs[annotation].astype("float32")
elif data_type.kind in ['i', 'u'] and data_type != "int32":
self.data.obs[annotation] = self.data.obs[annotation].astype("int32")
if self.data.X.dtype != 'float32': if self.data.X.dtype != 'float32':
self.data.X = self.data.X.astype("float32") warnings.warn(f"Scanpy data matrix is in {self.data.X.dtype} format not float32. Precision may be truncated.")
def cells(self): def cells(self):
return list(self.data.obs.index) return list(self.data.obs.index)