From 3c9f13f67de7634952190ec6bd56f489ed543e50 Mon Sep 17 00:00:00 2001 From: Emre K <110906681+kocaemre@users.noreply.github.com> Date: Sat, 1 Aug 2026 04:25:11 +0200 Subject: [PATCH] FIX: Align random.draw size dispatch for numpy integers Signed-off-by: Emre K <110906681+kocaemre@users.noreply.github.com> --- quantecon/random/tests/test_utilities.py | 11 +++++++++++ quantecon/random/utilities.py | 11 ++++++++--- 2 files changed, 19 insertions(+), 3 deletions(-) diff --git a/quantecon/random/tests/test_utilities.py b/quantecon/random/tests/test_utilities.py index 30020a90..4296d03b 100644 --- a/quantecon/random/tests/test_utilities.py +++ b/quantecon/random/tests/test_utilities.py @@ -125,6 +125,17 @@ def test_return_types(self): out = func(self.cdf, size) assert_(out.shape == (size,)) + def test_numpy_integer_size_returns_array(self): + size = np.int64(10) + for func in self.draw_funcs: + out = func(self.cdf, size) + assert_(out.shape == (size,)) + + def test_bool_size_is_treated_as_scalar(self): + for func in self.draw_funcs: + out = func(self.cdf, True) + assert_(isinstance(out, numbers.Integral)) + def test_return_values(self): for func in self.draw_funcs: out = func(self.cdf) diff --git a/quantecon/random/utilities.py b/quantecon/random/utilities.py index 9ca73ece..21377e63 100644 --- a/quantecon/random/utilities.py +++ b/quantecon/random/utilities.py @@ -246,7 +246,11 @@ def draw(cdf, size=None, rng=None): """ if rng is None: rng = np.random - if isinstance(size, int): + integer_size = ( + isinstance(size, (int, np.integer)) + and not isinstance(size, (bool, np.bool_)) + ) + if integer_size: rs = rng.random(size) out = np.searchsorted(cdf, rs, side='right') return out @@ -287,8 +291,9 @@ def _is_no_rng(numba_type): # holds the two paths together. @overload(draw) def ol_draw(cdf, size=None, rng=None): + is_integer_size = isinstance(size, types.Integer) and not isinstance(size, types.Boolean) if isinstance(rng, types.NumPyRandomGeneratorType): - if isinstance(size, types.Integer): + if is_integer_size: def draw_impl(cdf, size=None, rng=None): rs = rng.random(size) out = np.empty(size, dtype=np.int_) @@ -300,7 +305,7 @@ def draw_impl(cdf, size=None, rng=None): r = rng.random() return np.searchsorted(cdf, r, side='right') elif _is_no_rng(rng): - if isinstance(size, types.Integer): + if is_integer_size: def draw_impl(cdf, size=None, rng=None): rs = np.random.random(size) out = np.empty(size, dtype=np.int_)