X float16 support (#2406)

* float16 support

* fix type checks

* PR review comments

* add tests for custom json encoder; rename and comment for posterity

* lint

* typos
This commit is contained in:
Bruce Martin
2021-08-25 11:12:49 -07:00
committed by Colin Megill
parent eaae6df5e3
commit 154d099fef
8 changed files with 102 additions and 15 deletions
@@ -1,3 +1,4 @@
from typing import Tuple
import numba import numba
import concurrent.futures import concurrent.futures
import numpy as np import numpy as np
@@ -6,7 +7,7 @@ from backend.common.constants import XApproximateDistribution
@numba.njit(error_model="numpy", nogil=True) @numba.njit(error_model="numpy", nogil=True)
def min_max(arr: np.ndarray): def min_max_fast(arr: np.ndarray) -> Tuple[float, float]:
"""Return (min, max) values for the ndarray.""" """Return (min, max) values for the ndarray."""
# initialize to first finite value in array. Normally, # initialize to first finite value in array. Normally,
@@ -47,6 +48,24 @@ def min_max(arr: np.ndarray):
return min_val, max_val return min_val, max_val
def min_max_numpy(arr: np.ndarray) -> Tuple[float, float]:
return arr.min(), arr.max()
def numba_has_support_for_scalar_type(arr: np.ndarray) -> bool:
"""Numba does not support half-floats, 128 bit floats, ints > 64 bit or non-scalars."""
if arr.dtype == np.float32 or arr.dtype == np.float64:
return True
if np.issubdtype(arr.dtype, np.integer) and arr.dtype <= np.int64:
return True
if arr.dtype == np.bool_:
return True
return False
def estimate_approximate_distribution(X) -> XApproximateDistribution: def estimate_approximate_distribution(X) -> XApproximateDistribution:
""" """
Estimate the distribution (normal, count) of the X matrix. Estimate the distribution (normal, count) of the X matrix.
@@ -72,6 +91,8 @@ def estimate_approximate_distribution(X) -> XApproximateDistribution:
else: else:
raise TypeError(f"Unsupported matrix format: {str(type(X))}") raise TypeError(f"Unsupported matrix format: {str(type(X))}")
min_max = min_max_fast if numba_has_support_for_scalar_type(Xdata) else min_max_numpy
CHUNKSIZE = 1 << 24 CHUNKSIZE = 1 << 24
if Xdata.size > CHUNKSIZE: if Xdata.size > CHUNKSIZE:
min_val = max_val = Xdata[0] min_val = max_val = Xdata[0]
+13 -5
View File
@@ -65,7 +65,13 @@ def path_join(base, *urls):
return btpl._replace(path=path).geturl() return btpl._replace(path=path).geturl()
class Float32JSONEncoder(json.JSONEncoder): class StrictJSONEncoder(json.JSONEncoder):
"""
Custom JSON encoder set-up performing two tasks:
1. Strict JSON conformance with non-finite floats (NaN, +/-Inf) via allow_nan=False
2. Convert various Numpy types into python types so the encoder will correctly encode.
"""
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
""" """
NaN/Infinities are illegal in standard JSON. Python extends JSON with NaN/Infinities are illegal in standard JSON. Python extends JSON with
@@ -78,9 +84,11 @@ class Float32JSONEncoder(json.JSONEncoder):
super().__init__(*args, **kwargs) super().__init__(*args, **kwargs)
def default(self, obj): def default(self, obj):
if isinstance(obj, np.float32): """This helps us convert types not supported by the native JSON encoder into
standard python types, eg, np.int64."""
if isinstance(obj, np.floating):
return float(obj) return float(obj)
elif isinstance(obj, np.integer): if isinstance(obj, np.integer):
return int(obj) return int(obj)
return json.JSONEncoder.default(self, obj) return json.JSONEncoder.default(self, obj)
@@ -89,8 +97,8 @@ def custom_format_warning(msg, *args, **kwargs):
return f"[cellxgene] Warning: {msg} \n" return f"[cellxgene] Warning: {msg} \n"
def jsonify_numpy(data): def jsonify_strict(data):
return json.dumps(data, cls=Float32JSONEncoder, allow_nan=False) return json.dumps(data, cls=StrictJSONEncoder, allow_nan=False)
def import_plugins(plugin_module): def import_plugins(plugin_module):
+2 -2
View File
@@ -24,7 +24,7 @@ import backend.czi_hosted.common.rest as common_rest
from backend.common.utils.data_locator import DataLocator from backend.common.utils.data_locator import DataLocator
from backend.common.errors import DatasetAccessError, RequestException from backend.common.errors import DatasetAccessError, RequestException
from backend.czi_hosted.common.health import health_check from backend.czi_hosted.common.health import health_check
from backend.common.utils.utils import path_join, Float32JSONEncoder from backend.common.utils.utils import path_join, StrictJSONEncoder
from backend.czi_hosted.data_common.matrix_loader import MatrixDataLoader from backend.czi_hosted.data_common.matrix_loader import MatrixDataLoader
webbp = Blueprint("webapp", "backend.czi_hosted.common.web", template_folder="templates") webbp = Blueprint("webapp", "backend.czi_hosted.common.web", template_folder="templates")
@@ -396,7 +396,7 @@ class Server:
self.app = Flask(__name__, static_folder=None) self.app = Flask(__name__, static_folder=None)
handle_api_base_url(self.app, app_config) handle_api_base_url(self.app, app_config)
self._before_adding_routes(self.app, app_config) self._before_adding_routes(self.app, app_config)
self.app.json_encoder = Float32JSONEncoder self.app.json_encoder = StrictJSONEncoder
server_config = app_config.server_config server_config = app_config.server_config
if server_config.app__server_timing_headers: if server_config.app__server_timing_headers:
ServerTiming(self.app, force_debug=True) ServerTiming(self.app, force_debug=True)
@@ -15,7 +15,7 @@ from backend.common.errors import (
UnsupportedSummaryMethod, UnsupportedSummaryMethod,
DatasetAccessError, DatasetAccessError,
) )
from backend.common.utils.utils import jsonify_numpy from backend.common.utils.utils import jsonify_strict
from backend.common.fbs.matrix import encode_matrix_fbs from backend.common.fbs.matrix import encode_matrix_fbs
@@ -336,7 +336,7 @@ class DataAdaptor(metaclass=ABCMeta):
) )
try: try:
return jsonify_numpy(result) return jsonify_strict(result)
except ValueError: except ValueError:
raise JSONEncodingValueError("Error encoding differential expression to JSON") raise JSONEncodingValueError("Error encoding differential expression to JSON")
+2 -2
View File
@@ -17,7 +17,7 @@ from flask_restful import Api, Resource
import backend.server.common.rest as common_rest import backend.server.common.rest as common_rest
from backend.common.errors import DatasetAccessError, RequestException from backend.common.errors import DatasetAccessError, RequestException
from backend.server.common.health import health_check from backend.server.common.health import health_check
from backend.common.utils.utils import Float32JSONEncoder from backend.common.utils.utils import StrictJSONEncoder
webbp = Blueprint("webapp", "backend.server.common.web", template_folder="templates") webbp = Blueprint("webapp", "backend.server.common.web", template_folder="templates")
@@ -257,7 +257,7 @@ class Server:
def __init__(self, app_config): def __init__(self, app_config):
self.app = Flask(__name__, static_folder=None) self.app = Flask(__name__, static_folder=None)
self._before_adding_routes(self.app, app_config) self._before_adding_routes(self.app, app_config)
self.app.json_encoder = Float32JSONEncoder self.app.json_encoder = StrictJSONEncoder
server_config = app_config.server_config server_config = app_config.server_config
# enable session data # enable session data
@@ -232,7 +232,7 @@ class AnndataAdaptor(DataAdaptor):
"Anndata data matrix is sparse, but not a CSC (columnar) matrix. " "Anndata data matrix is sparse, but not a CSC (columnar) matrix. "
"Performance may be improved by using CSC." "Performance may be improved by using CSC."
) )
if self.data.X.dtype != "float32": if self.data.X.dtype > np.dtype(np.float32):
warnings.warn( warnings.warn(
f"Anndata data matrix is in {self.data.X.dtype} format not float32. " f"Precision may be truncated." f"Anndata data matrix is in {self.data.X.dtype} format not float32. " f"Precision may be truncated."
) )
+2 -2
View File
@@ -8,7 +8,7 @@ from server_timing import Timing as ServerTiming
from backend.server.common.config.app_config import AppConfig from backend.server.common.config.app_config import AppConfig
from backend.common.constants import Axis, XApproximateDistribution from backend.common.constants import Axis, XApproximateDistribution
from backend.common.errors import FilterError, JSONEncodingValueError, ExceedsLimitError, UnsupportedSummaryMethod from backend.common.errors import FilterError, JSONEncodingValueError, ExceedsLimitError, UnsupportedSummaryMethod
from backend.common.utils.utils import jsonify_numpy from backend.common.utils.utils import jsonify_strict
from backend.common.fbs.matrix import encode_matrix_fbs from backend.common.fbs.matrix import encode_matrix_fbs
from backend.common.genesets import validate_gene_sets from backend.common.genesets import validate_gene_sets
@@ -331,7 +331,7 @@ class DataAdaptor(metaclass=ABCMeta):
) )
try: try:
return jsonify_numpy(result) return jsonify_strict(result)
except ValueError: except ValueError:
raise JSONEncodingValueError("Error encoding differential expression to JSON") raise JSONEncodingValueError("Error encoding differential expression to JSON")
@@ -0,0 +1,58 @@
import unittest
import numpy as np
from backend.common.utils.utils import (
jsonify_strict,
)
class TestJsonifyStrict(unittest.TestCase):
def test_jsonify_numpy_general_cases(self):
self.assertEqual(jsonify_strict({}), "{}")
self.assertEqual(jsonify_strict({"a": [], "b": "hello", "c": True}), '{"a": [], "b": "hello", "c": true}')
def test_jsonify_numpy_float_edges(self):
with self.assertRaises(ValueError):
jsonify_strict({"nan": [np.nan]})
with self.assertRaises(ValueError):
jsonify_strict({"pinf": [np.PINF]})
with self.assertRaises(ValueError):
jsonify_strict({"ninf": [np.NINF]})
def test_jsonify_numpy_ndarray(self):
values = {
"integer": [
np.int8(0),
np.int16(1),
np.int32(2),
np.int64(3),
np.uint8(4),
np.uint16(5),
np.uint32(6),
np.uint64(7),
],
"floating": [
np.float16(100.0),
np.float32(101.0),
np.float64(102.0),
],
}
# these just confirm our test assumptions
self.assertTrue(isinstance(values["floating"][0], np.float16))
self.assertTrue(isinstance(values["floating"][1], np.float32))
self.assertTrue(isinstance(values["floating"][2], np.float64))
self.assertTrue(isinstance(values["integer"][0], np.int8))
self.assertTrue(isinstance(values["integer"][1], np.int16))
self.assertTrue(isinstance(values["integer"][2], np.int32))
self.assertTrue(isinstance(values["integer"][3], np.int64))
self.assertTrue(isinstance(values["integer"][4], np.uint8))
self.assertTrue(isinstance(values["integer"][5], np.uint16))
self.assertTrue(isinstance(values["integer"][6], np.uint32))
self.assertTrue(isinstance(values["integer"][7], np.uint64))
# the actual test!
self.assertEqual(
jsonify_strict(values),
'{"floating": [100.0, 101.0, 102.0], "integer": [0, 1, 2, 3, 4, 5, 6, 7]}',
)