Typing annotate a `njit()`-decorated function accepting another function as parameter

Viewed 32

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]"
0 Answers
Related