How can I create a numba callable that has adjustable parameters?

Viewed 507

I would like to create a numba-compiled python callable (a function that I can use in another Numba-compiled function) that has an internal array that I can adjust to influence the result of the function call. In pure python, this would correspond to a class with a __call__ method:

class Test:
    def __init__(self, arr):
        self.arr = arr
    def __call__(self, idx):
        res = 0
        for i in idx:
            res += self.arr[i]
        return res

t = Test([0, 1, 2])
print(t([1, 2]))
t.arr = [1, 2, 3]
print(t([1, 2]))

which prints 3 and 5, respectively, so the result was different after I modified the internal array arr.

A literal translation to Numba using jitclass and numpy arrays looks like this

import numpy as np
import numba as nb

@nb.jitclass([('arr', nb.double[:])])
class Test:
    def __init__(self, arr):
        self.arr = arr.astype(np.double)
    def __call__(self, idx):
        res = 0
        for i in idx:
            res += self.arr[i]
        return res

t = Test(np.arange(3))
print(t(np.array([1, 2])))
t.arr = np.arange(3) + 1
print(t(np.array([1, 2])))

Unfortunately, this fails with TypeError: 'Test' object is not callable, since Numba does not seem to support __call__, yet.

I then tried to solve the problem using closures

import numpy as np
import numba as nb

arr = np.arange(5)

@nb.jit
def call(idx):
    res = 0
    for i in idx:
        res += arr[i]
    return res


print(call(np.array([1, 2])))
arr += 1
print(call(np.array([1, 2])))

but this prints 3 twice, since closures copy the data in arr into an internal representation, which I then cannot (easily?) change from the outside. I even tried to trick Numba, by using ctypes pointers on Numpy arrays I combination with numba.carray, but Numba still seems to copy the data, so I cannot manipulate it.

I understand that Numba wants to control the memory and avoid access to memory regions that might not be used anymore. However, I have a specific use case where I would like to avoid passing around the extra array arr and rather adjust the internal copy somehow. Is there any way to achieve this?

EDIT: I tried the suggestion by Daniel in the comments to use a method different than __call__, but this also does not work. Here is what I thought might work:

@nb.jitclass([('arr', nb.double[:])])
class Test:
    def __init__(self, arr):
        self.arr = arr

    def call(self, idx):
        return self.arr[idx]

a = Test(np.arange(5).astype(np.double))
print(a.call(3))
a.arr += 1
print(a.call(3))


@nb.njit
def rhs(idx):
    return a.call(idx)

rhs(3)

This prints 3 and 4, so the array arr can indeed be manipulated. However, using the instance a in a compiled method fails with a NotImplementedError, so I suspect this use case is not (yet) supported by Numba.

3 Answers

Divide the problem in two parts, a numba function and a pure python class:

import numpy as np
import numba

@numba.jit
def calc(arr, idx):
    res = 0
    for i in idx:
        res += arr[i]
    return res

class Test:
    def __init__(self, arr):
        self.arr = arr.astype(np.double)

    def __call__(self, idx):
        return calc(self.arr, idx)

t = Test(np.arange(3))
print(t(np.array([1, 2])))
t.arr = np.arange(3) + 1
print(t(np.array([1, 2])))

I believe you need @property before the methods of the class but this may not be the only issue

@nb.jitclass([('arr', nb.double[:])])
class Test:
    def __init__(self, arr):
        self.arr = arr

    @property
    def call(self, idx):
        return self.arr[idx]

a = Test(np.arange(5).astype(np.double))
print(a.call(3))
a.arr += 1
print(a.call(3))


@nb.njit
def rhs(idx):
    return a.call(idx)

rhs(3)

This effect is the result of nopython compilation. If your goal is to create such callable at any costs, even possibly without taking benefits from jit-compilation - object compilation mode is a simple solution for your problem. This may be acheived in your closure example code simply by providing forceobj=True parameter to @nb.jit decorator.

This code prints 3 and 5 respectively:

import numpy as np
import numba as nb

arr = np.arange(5)

@nb.jit(forceobj=True)
def call(idx):
    res = 0
    for i in idx:
        res += arr[i]
    return res


print(call(np.array([1, 2])))
arr += 1
print(call(np.array([1, 2])))
Related