mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-10-06 01:58:11 +08:00
CXG performance improvements (#1371)
* initial kv cache * add per-key locks * comments * memoize schema * add missing initialization * fix sever timing * fetch only what is requested * fix tests to not require strict ordering of columns * clean up annotation request * remove debugging print
This commit is contained in:
@@ -0,0 +1,71 @@
|
|||||||
|
import threading
|
||||||
|
from collections.abc import MutableMapping
|
||||||
|
|
||||||
|
|
||||||
|
class ImmutableKVCache(MutableMapping):
|
||||||
|
"""
|
||||||
|
Guarantees that the factory will be called for each key once, and
|
||||||
|
only once.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, factory):
|
||||||
|
self.factory = factory # user-provided factory function
|
||||||
|
self.lock = threading.Lock() # guards factory_calls
|
||||||
|
self.factory_calls = {} # per-key factory condition variables
|
||||||
|
self.cache = {} # result cache, indexed by key
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
def __getitem__(self, key):
|
||||||
|
if key in self.cache:
|
||||||
|
return self.cache[key]
|
||||||
|
|
||||||
|
# we need to call factory. First grab the main lock and the per-key CV.
|
||||||
|
factory_calls = None
|
||||||
|
creation_thr = False
|
||||||
|
with self.lock:
|
||||||
|
if key in self.cache:
|
||||||
|
return self.cache[key]
|
||||||
|
if key not in self.factory_calls:
|
||||||
|
creation_thr = True
|
||||||
|
self.factory_calls[key] = {'cv': threading.Condition(), 'is_done': False, 'error': None}
|
||||||
|
factory_calls = self.factory_calls[key]
|
||||||
|
|
||||||
|
# with the CV, create the value (or wait for it to be created)
|
||||||
|
cv = factory_calls['cv']
|
||||||
|
with cv:
|
||||||
|
if creation_thr:
|
||||||
|
try:
|
||||||
|
self.cache[key] = self.factory(key)
|
||||||
|
except Exception as e:
|
||||||
|
factory_calls['error'] = e
|
||||||
|
|
||||||
|
factory_calls['is_done'] = True
|
||||||
|
cv.notify_all()
|
||||||
|
else:
|
||||||
|
""" wait for the value to be available """
|
||||||
|
while not factory_calls['is_done']:
|
||||||
|
cv.wait()
|
||||||
|
|
||||||
|
with self.lock:
|
||||||
|
if key in self.factory_calls:
|
||||||
|
del self.factory_calls[key]
|
||||||
|
|
||||||
|
return self.cache[key]
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
""" weak iter, don't call factory """
|
||||||
|
return self.cache.__iter__()
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self.cache.__len__()
|
||||||
|
|
||||||
|
def __contains__(self, key):
|
||||||
|
""" weak contain - don't call factory """
|
||||||
|
return self.cache.__contains__(key)
|
||||||
|
|
||||||
|
def __delitem__(self, key):
|
||||||
|
del self.cache[key]
|
||||||
|
|
||||||
|
def __setitem__(self, key, value):
|
||||||
|
""" unsupported """
|
||||||
|
raise NotImplementedError
|
||||||
@@ -7,6 +7,7 @@ from server.common.utils import path_join
|
|||||||
from server.common.constants import Axis
|
from server.common.constants import Axis
|
||||||
from server.data_common.data_adaptor import DataAdaptor
|
from server.data_common.data_adaptor import DataAdaptor
|
||||||
from server.data_common.fbs.matrix import encode_matrix_fbs
|
from server.data_common.fbs.matrix import encode_matrix_fbs
|
||||||
|
from server.common.immutable_kvcache import ImmutableKVCache
|
||||||
import tiledb
|
import tiledb
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
@@ -23,7 +24,6 @@ class CxgAdaptor(DataAdaptor):
|
|||||||
|
|
||||||
def __init__(self, data_locator, config=None):
|
def __init__(self, data_locator, config=None):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.arrays = {}
|
|
||||||
self.lock = threading.Lock()
|
self.lock = threading.Lock()
|
||||||
|
|
||||||
self.data_locator = data_locator
|
self.data_locator = data_locator
|
||||||
@@ -31,6 +31,11 @@ class CxgAdaptor(DataAdaptor):
|
|||||||
if self.url[-1] != "/":
|
if self.url[-1] != "/":
|
||||||
self.url += "/"
|
self.url += "/"
|
||||||
|
|
||||||
|
# caching immutable state
|
||||||
|
self.lsuri_results = ImmutableKVCache(lambda key: self._lsuri(uri=key, tiledb_ctx=self.tiledb_ctx))
|
||||||
|
self.arrays = ImmutableKVCache(lambda key: self._open_array(uri=key, tiledb_ctx=self.tiledb_ctx))
|
||||||
|
self.schema = None
|
||||||
|
|
||||||
self._validate_and_initialize()
|
self._validate_and_initialize()
|
||||||
|
|
||||||
def cleanup(self):
|
def cleanup(self):
|
||||||
@@ -83,6 +88,18 @@ class CxgAdaptor(DataAdaptor):
|
|||||||
def get_path(self, *urls):
|
def get_path(self, *urls):
|
||||||
return path_join(self.url, *urls)
|
return path_join(self.url, *urls)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _lsuri(uri, tiledb_ctx):
|
||||||
|
def _cleanpath(p):
|
||||||
|
if p[-1] == "/":
|
||||||
|
return p[:-1]
|
||||||
|
else:
|
||||||
|
return p
|
||||||
|
|
||||||
|
result = []
|
||||||
|
tiledb.ls(uri, lambda path, type: result.append((_cleanpath(path), type)), ctx=tiledb_ctx)
|
||||||
|
return result
|
||||||
|
|
||||||
def lsuri(self, uri):
|
def lsuri(self, uri):
|
||||||
"""
|
"""
|
||||||
given a URI, do a tiledb.ls but normalizing for all path weirdness:
|
given a URI, do a tiledb.ls but normalizing for all path weirdness:
|
||||||
@@ -92,19 +109,9 @@ class CxgAdaptor(DataAdaptor):
|
|||||||
returns list of (absolute paths, type) *without* trailing slash
|
returns list of (absolute paths, type) *without* trailing slash
|
||||||
in the path.
|
in the path.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def _cleanpath(p):
|
|
||||||
if p[-1] == "/":
|
|
||||||
return p[:-1]
|
|
||||||
else:
|
|
||||||
return p
|
|
||||||
|
|
||||||
if uri[-1] != "/":
|
if uri[-1] != "/":
|
||||||
uri += "/"
|
uri += "/"
|
||||||
|
return self.lsuri_results[uri]
|
||||||
result = []
|
|
||||||
tiledb.ls(uri, lambda path, type: result.append((_cleanpath(path), type)), ctx=self.tiledb_ctx)
|
|
||||||
return result
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def isvalid(url):
|
def isvalid(url):
|
||||||
@@ -160,19 +167,14 @@ class CxgAdaptor(DataAdaptor):
|
|||||||
self.about = about
|
self.about = about
|
||||||
self.cxg_version = cxg_version
|
self.cxg_version = cxg_version
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _open_array(uri, tiledb_ctx):
|
||||||
|
return tiledb.DenseArray(uri, mode="r", ctx=tiledb_ctx)
|
||||||
|
|
||||||
def open_array(self, name):
|
def open_array(self, name):
|
||||||
try:
|
try:
|
||||||
with self.lock:
|
p = self.get_path(name)
|
||||||
array = self.arrays.get(name)
|
return self.arrays[p]
|
||||||
if array:
|
|
||||||
return array
|
|
||||||
p = self.get_path(name)
|
|
||||||
try:
|
|
||||||
array = tiledb.DenseArray(p, mode="r", ctx=self.tiledb_ctx)
|
|
||||||
except tiledb.libtiledb.TileDBError:
|
|
||||||
raise DatasetAccessError(name)
|
|
||||||
self.arrays[name] = array
|
|
||||||
return array
|
|
||||||
except tiledb.libtiledb.TileDBError:
|
except tiledb.libtiledb.TileDBError:
|
||||||
raise DatasetAccessError(name)
|
raise DatasetAccessError(name)
|
||||||
|
|
||||||
@@ -280,7 +282,10 @@ class CxgAdaptor(DataAdaptor):
|
|||||||
schema["categories"] = schema_hints["categories"]
|
schema["categories"] = schema_hints["categories"]
|
||||||
return schema
|
return schema
|
||||||
|
|
||||||
def get_schema(self):
|
def _get_schema(self):
|
||||||
|
if self.schema:
|
||||||
|
return self.schema
|
||||||
|
|
||||||
shape = self.get_shape()
|
shape = self.get_shape()
|
||||||
dtype = self.get_X_array_dtype()
|
dtype = self.get_X_array_dtype()
|
||||||
|
|
||||||
@@ -320,20 +325,77 @@ class CxgAdaptor(DataAdaptor):
|
|||||||
schema = {"dataframe": dataframe, "annotations": annotations, "layout": {"obs": obs_layout}}
|
schema = {"dataframe": dataframe, "annotations": annotations, "layout": {"obs": obs_layout}}
|
||||||
return schema
|
return schema
|
||||||
|
|
||||||
|
def get_schema(self):
|
||||||
|
if self.schema is None:
|
||||||
|
with self.lock:
|
||||||
|
self.schema = self._get_schema()
|
||||||
|
return self.schema
|
||||||
|
|
||||||
|
def _annotations_field_split(self, axis, fields, A, labels):
|
||||||
|
"""
|
||||||
|
fields: requested fields, may be None (all)
|
||||||
|
labels: writable user annotations dataframe, if any
|
||||||
|
|
||||||
|
Remove redundant fields, raise KeyError on non-existant fields,
|
||||||
|
and split into three lists:
|
||||||
|
fields_to_fetch_from_cxg
|
||||||
|
fields_to_fetch_from_labels
|
||||||
|
fields_to_return
|
||||||
|
|
||||||
|
if we have to return from labels, the fetch fields will contain the index
|
||||||
|
to join on, which may not be in fields_to_return
|
||||||
|
"""
|
||||||
|
need_labels = axis == Axis.OBS and labels is not None and not labels.empty
|
||||||
|
index_key = self.get_obs_names() if need_labels else None
|
||||||
|
|
||||||
|
if not fields:
|
||||||
|
return (None, None, None, index_key)
|
||||||
|
|
||||||
|
cxg_keys = frozenset([a.name for a in A.schema])
|
||||||
|
user_anno_keys = frozenset(labels.columns.tolist()) if need_labels else frozenset()
|
||||||
|
return_keys = frozenset(fields)
|
||||||
|
|
||||||
|
label_join_index = (
|
||||||
|
frozenset([index_key]) if need_labels and (return_keys & user_anno_keys) else frozenset()
|
||||||
|
)
|
||||||
|
|
||||||
|
unknown_fields = return_keys - (cxg_keys | user_anno_keys)
|
||||||
|
if unknown_fields:
|
||||||
|
raise KeyError("_".join(unknown_fields))
|
||||||
|
|
||||||
|
return (
|
||||||
|
list((return_keys & cxg_keys) | label_join_index),
|
||||||
|
list(return_keys & user_anno_keys),
|
||||||
|
list(return_keys),
|
||||||
|
index_key
|
||||||
|
)
|
||||||
|
|
||||||
def annotation_to_fbs_matrix(self, axis, fields=None, labels=None):
|
def annotation_to_fbs_matrix(self, axis, fields=None, labels=None):
|
||||||
with ServerTiming.time(f"annotations.{axis}.query"):
|
with ServerTiming.time(f"annotations.{axis}.query"):
|
||||||
A = self.open_array(str(axis))
|
A = self.open_array(str(axis))
|
||||||
if axis == Axis.OBS:
|
|
||||||
if labels is not None and not labels.empty:
|
|
||||||
df = pd.DataFrame.from_dict(A[:])
|
|
||||||
df = df.join(labels, self.get_obs_names())
|
|
||||||
else:
|
|
||||||
df = pd.DataFrame.from_dict(A[:])
|
|
||||||
else:
|
|
||||||
df = pd.DataFrame.from_dict(A[:])
|
|
||||||
|
|
||||||
if fields is not None and len(fields) > 0:
|
# may raise if fields contains unknown key
|
||||||
df = df[fields]
|
cxg_fields, anno_fields, return_fields, index_field = self._annotations_field_split(axis, fields, A, labels)
|
||||||
|
|
||||||
|
if cxg_fields is None:
|
||||||
|
data = A[:]
|
||||||
|
elif cxg_fields:
|
||||||
|
data = A.query(attrs=cxg_fields)[:]
|
||||||
|
else:
|
||||||
|
data = {}
|
||||||
|
|
||||||
|
df = pd.DataFrame.from_dict(data)
|
||||||
|
|
||||||
|
if axis == Axis.OBS and labels is not None and not labels.empty:
|
||||||
|
if anno_fields is None:
|
||||||
|
assert index_field
|
||||||
|
df = df.join(labels, index_field)
|
||||||
|
elif anno_fields:
|
||||||
|
assert index_field
|
||||||
|
df = df.join(labels[anno_fields], index_field)
|
||||||
|
|
||||||
|
if return_fields:
|
||||||
|
df = df[return_fields]
|
||||||
|
|
||||||
with ServerTiming.time(f"annotations.{axis}.encode"):
|
with ServerTiming.time(f"annotations.{axis}.encode"):
|
||||||
fbs = encode_matrix_fbs(df, col_idx=df.columns)
|
fbs = encode_matrix_fbs(df, col_idx=df.columns)
|
||||||
|
|||||||
@@ -84,7 +84,7 @@ class EndPoints(object):
|
|||||||
self.assertIsNone(df["row_idx"])
|
self.assertIsNone(df["row_idx"])
|
||||||
self.assertEqual(len(df["columns"]), df["n_cols"])
|
self.assertEqual(len(df["columns"]), df["n_cols"])
|
||||||
obs_index_col_name = self.schema["schema"]["annotations"]["obs"]["index"]
|
obs_index_col_name = self.schema["schema"]["annotations"]["obs"]["index"]
|
||||||
self.assertListEqual(
|
self.assertCountEqual(
|
||||||
df["col_idx"],
|
df["col_idx"],
|
||||||
[obs_index_col_name, "n_genes", "percent_mito", "n_counts", "louvain"]
|
[obs_index_col_name, "n_genes", "percent_mito", "n_counts", "louvain"]
|
||||||
+ (["cluster-test"] if self.ANNOTATIONS_ENABLED else []),
|
+ (["cluster-test"] if self.ANNOTATIONS_ENABLED else []),
|
||||||
@@ -104,7 +104,7 @@ class EndPoints(object):
|
|||||||
self.assertIsNotNone(df["columns"])
|
self.assertIsNotNone(df["columns"])
|
||||||
self.assertIsNone(df["row_idx"])
|
self.assertIsNone(df["row_idx"])
|
||||||
self.assertEqual(len(df["columns"]), df["n_cols"])
|
self.assertEqual(len(df["columns"]), df["n_cols"])
|
||||||
self.assertListEqual(df["col_idx"], ["n_genes", "percent_mito"])
|
self.assertCountEqual(df["col_idx"], ["n_genes", "percent_mito"])
|
||||||
|
|
||||||
def test_get_annotations_obs_error(self):
|
def test_get_annotations_obs_error(self):
|
||||||
endpoint = "annotations/obs"
|
endpoint = "annotations/obs"
|
||||||
@@ -157,7 +157,7 @@ class EndPoints(object):
|
|||||||
self.assertIsNone(df["row_idx"])
|
self.assertIsNone(df["row_idx"])
|
||||||
self.assertEqual(len(df["columns"]), df["n_cols"])
|
self.assertEqual(len(df["columns"]), df["n_cols"])
|
||||||
var_index_col_name = self.schema["schema"]["annotations"]["var"]["index"]
|
var_index_col_name = self.schema["schema"]["annotations"]["var"]["index"]
|
||||||
self.assertListEqual(df["col_idx"], [var_index_col_name, "n_cells"])
|
self.assertCountEqual(df["col_idx"], [var_index_col_name, "n_cells"])
|
||||||
|
|
||||||
def test_get_annotations_var_keys_fbs(self):
|
def test_get_annotations_var_keys_fbs(self):
|
||||||
endpoint = "annotations/var"
|
endpoint = "annotations/var"
|
||||||
@@ -173,7 +173,7 @@ class EndPoints(object):
|
|||||||
self.assertIsNotNone(df["columns"])
|
self.assertIsNotNone(df["columns"])
|
||||||
self.assertIsNone(df["row_idx"])
|
self.assertIsNone(df["row_idx"])
|
||||||
self.assertEqual(len(df["columns"]), df["n_cols"])
|
self.assertEqual(len(df["columns"]), df["n_cols"])
|
||||||
self.assertListEqual(df["col_idx"], ["n_cells"])
|
self.assertCountEqual(df["col_idx"], ["n_cells"])
|
||||||
|
|
||||||
def test_get_annotations_var_error(self):
|
def test_get_annotations_var_error(self):
|
||||||
endpoint = "annotations/var"
|
endpoint = "annotations/var"
|
||||||
|
|||||||
Reference in New Issue
Block a user