python numba error when passing a np.random.Generator object to a njit-decorated function

Viewed 139

I have a module that looks as following:

import numba as nb
import numpy as np

@nb.njit
def random_numba_function(rng):
    print(rng.choice([0,1,2]))

seed = 1
rng = np.random.default_rng(seed)
random_numba_function(rng)

This returns the error message for the decorated function random_numba_function:

argument 0: Cannot determine Numba type of <class 'numpy.random._generator.Generator'>

The documentation states that the numpy.random.seed function is supported, but the np.random.default_rng function is not listed as a supported function which probably means that you cannot pass a random number generator this way.

In my particular situation, I also have a conditional inside of the decorated function which decides whether or not to "draw" a number from the random number generator. This means that I can't just calculate the index of the random choice ahead of time and pass it to the function as an integer, since I only know if I am going to be using it once I am inside of the function. Is there an efficient way to substitute the way random seed numbers are generated with numba?

0 Answers
Related