mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-10-01 23:28:12 +08:00
This PR contains a refactoring to make adding new features easier. The new features include supporting the tiledb format, and the multi dataset application. The refactoring includes Simplifying the directory structure and files. a class structure to handle annotations (currently one type: AnnotationsLocalFile). a class to handle application configuration a class structure to handle matrix data (currently AnndataAdaptor and CxgAdaptor). CxgAdaptor uses tiledb. Algorithms that were previously dependent on the scanpy anndata object are now generalized to work with an abstract interface. The multi dataset option is not fully supported yet, and so the option to use it is hidden. Use "cli launch --dataroot ..." To access this feature. All combinations of app single dataset/ app multi dataset and AnndataAdaptor/CxgAdaptor work with all the features, such as annotations, ontologies, diffexp.
322 lines
10 KiB
Python
322 lines
10 KiB
Python
import os
|
|
import json
|
|
from server.common.utils import dtype_to_schema
|
|
from server.common.errors import DatasetAccessError
|
|
from server.common.utils import path_join
|
|
from server.common.constants import Axis
|
|
from server.data_common.data_adaptor import DataAdaptor
|
|
from server.data_common.fbs.matrix import encode_matrix_fbs
|
|
import tiledb
|
|
import numpy as np
|
|
import pandas as pd
|
|
from server_timing import Timing as ServerTiming
|
|
import threading
|
|
|
|
|
|
class CxgAdaptor(DataAdaptor):
|
|
|
|
# TODO: The tiledb context parameters should be a configuration option
|
|
tiledb_ctx = tiledb.Ctx({
|
|
'sm.tile_cache_size': 8 * 1024 * 1024 * 1024,
|
|
'sm.num_reader_threads': 32,
|
|
})
|
|
|
|
def __init__(self, location, config=None):
|
|
super().__init__(config)
|
|
self.url = location
|
|
self.arrays = {}
|
|
self.lock = threading.Lock()
|
|
|
|
self.url = location
|
|
if self.url[-1] != '/':
|
|
self.url += '/'
|
|
|
|
self._validate_and_initialize()
|
|
|
|
def cleanup(self):
|
|
"""close all the open tiledb arrays"""
|
|
for array in self.arrays.values():
|
|
array.close()
|
|
self.arrays.clear()
|
|
|
|
@staticmethod
|
|
def pre_load_validation(location):
|
|
if not CxgAdaptor.isvalid(location):
|
|
raise DatasetAccessError(f"cxg matrix is not valid: {location}")
|
|
|
|
@staticmethod
|
|
def file_size(location):
|
|
return 0
|
|
|
|
@staticmethod
|
|
def open(location, args):
|
|
return CxgAdaptor(location, args)
|
|
|
|
def get_location(self):
|
|
return self.url
|
|
|
|
def get_name(self):
|
|
return "cellxgene cxcxgg adaptor version"
|
|
|
|
def get_library_versions(self):
|
|
return dict(tiledb=tiledb.__version__)
|
|
|
|
def get_path(self, *urls):
|
|
return path_join(self.url, *urls)
|
|
|
|
def lsuri(self, uri):
|
|
"""
|
|
given a URI, do a tiledb.ls but normalizing for all path weirdness:
|
|
* S3 URIs require trailing slash. file: doesn't care.
|
|
* results on S3 *have* a trailing slash, Posix does not.
|
|
|
|
returns list of (absolute paths, type) *without* trailing slash
|
|
in the path.
|
|
"""
|
|
def _cleanpath(p):
|
|
if p[-1] == '/':
|
|
return p[:-1]
|
|
else:
|
|
return p
|
|
|
|
if uri[-1] != '/':
|
|
uri += '/'
|
|
|
|
result = []
|
|
tiledb.ls(uri,
|
|
lambda path, type: result.append((_cleanpath(path), type)),
|
|
ctx=self.tiledb_ctx)
|
|
return result
|
|
|
|
@staticmethod
|
|
def isvalid(url):
|
|
"""
|
|
Return True if this looks like a valid CXG, False if not. Just a quick/cheap
|
|
test, not to be fully trusted.
|
|
"""
|
|
if not tiledb.object_type(url) == "group":
|
|
return False
|
|
if not tiledb.object_type(path_join(url, "obs")) == "array":
|
|
return False
|
|
if not tiledb.object_type(path_join(url, "var")) == "array":
|
|
return False
|
|
if not tiledb.object_type(path_join(url, "X")) == "array":
|
|
return False
|
|
if not tiledb.object_type(path_join(url, "emb")) == "group":
|
|
return False
|
|
return True
|
|
|
|
def _validate_and_initialize(self):
|
|
if not self.isvalid(self.url):
|
|
raise DatasetAccessError(f"invalid cxg dataset {self.url}")
|
|
|
|
def open_array(self, name):
|
|
try:
|
|
with self.lock:
|
|
array = self.arrays.get(name)
|
|
if array:
|
|
return array
|
|
p = self.get_path(name)
|
|
try:
|
|
array = tiledb.DenseArray(p, mode="r", ctx=self.tiledb_ctx)
|
|
except tiledb.libtiledb.TileDBError as e:
|
|
raise AttributeError(str(e))
|
|
self.arrays[name] = array
|
|
return array
|
|
except tiledb.libtiledb.TileDBError as e:
|
|
raise AttributeError(str(e))
|
|
|
|
def get_embedding_array(self, ename, dims=2):
|
|
array = self.open_array(f"emb/{ename}")
|
|
return array[:, 0:dims]
|
|
|
|
def get_X_array(self, obs_mask=None, var_mask=None):
|
|
obs_items = self._convert_mask(obs_mask)
|
|
var_items = self._convert_mask(var_mask)
|
|
X = self.open_array("X")
|
|
if obs_items == slice(None) and var_items == slice(None):
|
|
data = X[:, :]
|
|
else:
|
|
data = X.multi_index[obs_items, var_items]['']
|
|
return data
|
|
|
|
def get_shape(self):
|
|
X = self.open_array("X")
|
|
return X.shape
|
|
|
|
def get_X_array_dtype(self):
|
|
X = self.open_array("X")
|
|
return X.dtype
|
|
|
|
def query_var_array(self, term_name):
|
|
var = self.open_array("var")
|
|
data = var.query(attrs=[term_name])[:][term_name]
|
|
return data
|
|
|
|
def query_obs_array(self, term_name):
|
|
var = self.open_array("obs")
|
|
try:
|
|
data = var.query(attrs=[term_name])[:][term_name]
|
|
except tiledb.libtiledb.TileDBError as e:
|
|
raise AttributeError(str(e))
|
|
return data
|
|
|
|
def get_obs_names(self):
|
|
# get the index from the meta data
|
|
obs = self.open_array("obs")
|
|
meta = json.loads(obs.meta["cxg_schema"])
|
|
index_name = meta["index"]
|
|
return index_name
|
|
|
|
def get_obs_index(self):
|
|
obs = self.open_array("obs")
|
|
meta = json.loads(obs.meta["cxg_schema"])
|
|
index_name = meta["index"]
|
|
data = obs.query(attrs=[index_name])[:][index_name]
|
|
return data
|
|
|
|
def get_obs_columns(self):
|
|
obs = self.open_array("obs")
|
|
schema = obs.schema
|
|
col_names = [attr.name for attr in schema]
|
|
return pd.Index(col_names)
|
|
|
|
# function to get the embedding
|
|
# this function to iterate through embeddings.
|
|
def get_embedding_names(self):
|
|
with ServerTiming.time(f'layout.lsuri'):
|
|
pemb = self.get_path("emb")
|
|
embeddings = [
|
|
os.path.basename(p) for (p, t) in self.lsuri(pemb) if t == 'array'
|
|
]
|
|
return embeddings
|
|
|
|
@staticmethod
|
|
def _get_col_type(attr, schema_hints={}):
|
|
type_hint = schema_hints.get(attr.name, {})
|
|
dtype = attr.dtype
|
|
schema = {}
|
|
# type hints take precedence
|
|
if 'type' in type_hint:
|
|
schema['type'] = type_hint['type']
|
|
elif dtype == np.float32:
|
|
schema['type'] = 'float32'
|
|
elif dtype == np.int32:
|
|
schema['type'] = 'int32'
|
|
elif dtype == np.bool_:
|
|
schema['type'] = 'boolean'
|
|
elif dtype == np.str:
|
|
schema['type'] = 'string'
|
|
elif dtype == "category":
|
|
schema["type"] = "categorical"
|
|
schema["categories"] = dtype.categories.tolist()
|
|
else:
|
|
raise TypeError(
|
|
f"Annotations of type {dtype} are unsupported."
|
|
)
|
|
|
|
if schema['type'] == 'categorical' and 'categories' in schema_hints:
|
|
schema['categories'] = schema_hints['categories']
|
|
return schema
|
|
|
|
def get_schema(self):
|
|
shape = self.get_shape()
|
|
dtype = self.get_X_array_dtype()
|
|
|
|
dataframe = {
|
|
'nObs': shape[0],
|
|
'nVar': shape[1],
|
|
'type': dtype.name
|
|
}
|
|
|
|
annotations = {}
|
|
for ax in ('obs', 'var'):
|
|
A = self.open_array(ax)
|
|
schema_hints = json.loads(A.meta['cxg_schema']) if 'cxg_schema' in A.meta else {}
|
|
if type(schema_hints) is not dict:
|
|
raise TypeError(f'Array schema was malformed.')
|
|
|
|
cols = []
|
|
for attr in A.schema:
|
|
schema = dict(name=attr.name, writable=False)
|
|
type_hint = schema_hints.get(attr.name, {})
|
|
# type hints take precedence
|
|
if 'type' in type_hint:
|
|
schema['type'] = type_hint['type']
|
|
if schema['type'] == 'categorical' and 'categories' in type_hint:
|
|
schema['categories'] = type_hint['categories']
|
|
else:
|
|
schema.update(dtype_to_schema(attr.dtype))
|
|
cols.append(schema)
|
|
|
|
annotations[ax] = dict(columns=cols)
|
|
|
|
if 'index' in schema_hints:
|
|
annotations[ax].update({'index': schema_hints['index']})
|
|
|
|
obs_layout = []
|
|
embeddings = self.get_embedding_names()
|
|
for ename in embeddings:
|
|
A = self.open_array(f"emb/{ename}")
|
|
obs_layout.append({
|
|
'name': ename,
|
|
'type': A.dtype.name,
|
|
'dims': [f'{ename}_{d}' for d in range(0, A.ndim)]
|
|
})
|
|
|
|
schema = {
|
|
'dataframe': dataframe,
|
|
'annotations': annotations,
|
|
'layout': {'obs': obs_layout}
|
|
}
|
|
return schema
|
|
|
|
def annotation_to_fbs_matrix(self, axis, fields=None, labels=None):
|
|
with ServerTiming.time(f'annotations.{axis}.query'):
|
|
A = self.open_array(str(axis))
|
|
if fields is not None and len(fields) > 0:
|
|
try:
|
|
df = pd.DataFrame(A.query(attrs=fields)[:])
|
|
except tiledb.libtiledb.TileDBError:
|
|
raise KeyError("bad field {fields}")
|
|
|
|
else:
|
|
df = pd.DataFrame.from_dict(A[:])
|
|
|
|
if axis == Axis.OBS:
|
|
if labels is not None and not labels.empty:
|
|
obs_names = self.get_obs_names()
|
|
df = df.join(labels, obs_names)
|
|
|
|
with ServerTiming.time(f'annotations.{axis}.encode'):
|
|
fbs = encode_matrix_fbs(df, col_idx=df.columns)
|
|
|
|
return fbs
|
|
|
|
@staticmethod
|
|
def _convert_mask(boolarray):
|
|
"""Convert an index mask to a list of ranges or indices that can be used in a multi_index."""
|
|
if boolarray is None:
|
|
return slice(None)
|
|
assert type(boolarray) == np.ndarray
|
|
assert(boolarray.dtype) == bool
|
|
|
|
selector = np.nonzero(boolarray)[0]
|
|
|
|
if len(selector) == 0:
|
|
return slice(None)
|
|
|
|
result = []
|
|
current = slice(selector[0], selector[0])
|
|
for sel in selector[1:]:
|
|
if sel == current.stop + 1:
|
|
current = slice(current.start, sel)
|
|
else:
|
|
result.append(current if current.start != current.stop else current.start)
|
|
current = slice(sel, sel)
|
|
|
|
if len(result) == 0 or result[-1] != current:
|
|
result.append(current if current.start != current.stop else current.start)
|
|
|
|
return result
|