Assume I have the following toy example:
import typing as _T
def foo(fn: _T.Callable[[int, int], int]) -> int:
return fn(0, 1)
Now, for some reason I want to accelerate foo() with Numba in Non-Python mode:
import numba as nb
@nb.njit
def foo_nb(fn: _T.Callable[[int, int], int]) -> int:
return fn(0, 1)
However, this no longer captures the fact that fn now needs to be compatible with Numba.
What should I replace _T.Callable[[int, int], int] with?
I could use a string like "nb.njit(callable)[[int, int], int]" but I suspect there is something better.
Clearly, _T.Callable[[int, int], int] is wrong, since passing a normal Python callable will result in an error:
TypingError: Failed in nopython mode pipeline (step: nopython frontend)
non-precise type pyobject
During: typing of argument at <ipython-input-21-6068134a3c84> (16)
....
This error may have been caused by the following argument(s):
- argument 1: cannot determine Numba type of <class 'function'>
The result of type(foo_nb) is numba.core.registry.CPUDispatcher but I am unsure it serves well for documenting the type I should pass to fn_nb in foo_nb().
I know that mypy will fail to work for Numba-accelerated and @-decorated functions, as per this issue.
A minimal example follows:
import typing as _T
import numba as nb
def foo_nb(fn: _T.Callable[[int, int], int]) -> int:
return fn(0, 1)
def f_nb(a: int, b: int) -> int:
return a + b
foo_nb = nb.njit(foo_nb)
f_nb = nb.njit(f_nb)
foo_nb(f_nb)
foo_nb(lambda x, y: x + y)
# `mypy` will not detect an issue, but the running this code will fail
foo_nb(1)
# error: Argument 1 to "foo_nb" has incompatible type "int"; expected "Callable[[int, int], int]"