mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-10-01 04:58:12 +08:00
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:
committed by
Colin Megill
parent
eaae6df5e3
commit
154d099fef
@@ -1,3 +1,4 @@
|
||||
from typing import Tuple
|
||||
import numba
|
||||
import concurrent.futures
|
||||
import numpy as np
|
||||
@@ -6,7 +7,7 @@ from backend.common.constants import XApproximateDistribution
|
||||
|
||||
|
||||
@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."""
|
||||
|
||||
# initialize to first finite value in array. Normally,
|
||||
@@ -47,6 +48,24 @@ def min_max(arr: np.ndarray):
|
||||
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:
|
||||
"""
|
||||
Estimate the distribution (normal, count) of the X matrix.
|
||||
@@ -72,6 +91,8 @@ def estimate_approximate_distribution(X) -> XApproximateDistribution:
|
||||
else:
|
||||
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
|
||||
if Xdata.size > CHUNKSIZE:
|
||||
min_val = max_val = Xdata[0]
|
||||
|
||||
@@ -65,7 +65,13 @@ def path_join(base, *urls):
|
||||
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):
|
||||
"""
|
||||
NaN/Infinities are illegal in standard JSON. Python extends JSON with
|
||||
@@ -78,9 +84,11 @@ class Float32JSONEncoder(json.JSONEncoder):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
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)
|
||||
elif isinstance(obj, np.integer):
|
||||
if isinstance(obj, np.integer):
|
||||
return int(obj)
|
||||
return json.JSONEncoder.default(self, obj)
|
||||
|
||||
@@ -89,8 +97,8 @@ def custom_format_warning(msg, *args, **kwargs):
|
||||
return f"[cellxgene] Warning: {msg} \n"
|
||||
|
||||
|
||||
def jsonify_numpy(data):
|
||||
return json.dumps(data, cls=Float32JSONEncoder, allow_nan=False)
|
||||
def jsonify_strict(data):
|
||||
return json.dumps(data, cls=StrictJSONEncoder, allow_nan=False)
|
||||
|
||||
|
||||
def import_plugins(plugin_module):
|
||||
|
||||
Reference in New Issue
Block a user