diff --git a/jaxtyping/_array_types.py b/jaxtyping/_array_types.py index fadeaeb..5ac236f 100644 --- a/jaxtyping/_array_types.py +++ b/jaxtyping/_array_types.py @@ -498,6 +498,16 @@ def _make_array(array_type, dim_str, dtypes, name): return array_type else: return _not_made + elif array_type is np.bool_: + if _check_scalar("bool", dtypes, dims): + return array_type + else: + return _not_made + elif array_type is np.generic or array_type is np.number: + if _check_scalar("", dtypes, dims): + return array_type + else: + return _not_made if issubclass(array_type, AbstractArray): if dtypes is _any_dtype: dtypes = array_type.dtypes diff --git a/test/test_array.py b/test/test_array.py index 3d7c620..3be95f1 100644 --- a/test/test_array.py +++ b/test/test_array.py @@ -496,22 +496,18 @@ def test_arraylike(typecheck, getkey): assert set(get_args(floatlike1)) == { Float32[Array, ""], Float32[np.ndarray, ""], - Float32[np.bool_, ""], Float32[np.number, ""], float, } assert set(get_args(floatlike2)) == { Float[Array, ""], Float[np.ndarray, ""], - Float[np.bool_, ""], Float[np.number, ""], float, } assert set(get_args(floatlike3)) == { Float32[Array, "4"], Float32[np.ndarray, "4"], - Float32[np.bool_, "4"], - Float32[np.number, "4"], } shaped1 = Shaped[ArrayLike, ""] @@ -531,8 +527,6 @@ def test_arraylike(typecheck, getkey): assert set(get_args(shaped2)) == { Shaped[Array, "4"], Shaped[np.ndarray, "4"], - Shaped[np.bool_, "4"], - Shaped[np.number, "4"], }