在njit函数中调用Numba cfunc遇data_as方法缺失问题求助
解决Numba njit函数中调用cfunc时的指针转换问题
问题核心:Numba的nopython模式不支持numpy数组的ctypes.data_as方法,因此无法直接在njit编译的函数中用该方式获取C指针。
解决方案
使用Numba提供的nb.address_as_void_pointer获取数组数据的内存地址,再通过nb.cast将其转换为目标类型的C指针,替代ctypes.data_as的功能。
修改后的代码
import ctypes import numpy as np import numba as nb @nb.cfunc(nb.types.void( nb.types.CPointer(nb.types.double), nb.types.CPointer(nb.types.double), nb.types.int64, nb.types.int64, nb.types.int64, nb.types.int64, )) def get_param2(xn_, x_, idx, n, m1, m2): in_array = nb.carray(x_, (n, m1, m2)) out_array = nb.carray(xn_, (m1, m2)) if idx >= n: idx = n - 1 out_array[:, :] = in_array[idx] def test_get_param(): # 正常运行 A = np.zeros((100, 2, 3)) Ai = np.empty((2, 3)) get_param2( Ai.ctypes.data_as(ctypes.POINTER(ctypes.c_double)), A.ctypes.data_as(ctypes.POINTER(ctypes.c_double)), 40, *A.shape, ) assert np.array_equal(A[40], Ai) @nb.jit(nopython=True) def get_param_njit(A, i): Ai = np.empty((2, 3)) # 替换成Numba支持的指针转换方式 ai_ptr = nb.cast(nb.address_as_void_pointer(Ai.ravel()), nb.types.CPointer(nb.types.double)) a_ptr = nb.cast(nb.address_as_void_pointer(A.ravel()), nb.types.CPointer(nb.types.double)) get_param2( ai_ptr, a_ptr, i, *A.shape ) return Ai def test_get_param_njit(): A = np.zeros((100, 2, 3)) Ai = get_param_njit(A, 40) assert np.array_equal(A[40], Ai)
说明
nb.address_as_void_pointer:获取数组数据的内存地址(返回void*类型),Numba的nopython模式完全支持该函数。nb.cast:将void*地址转换为目标类型的CPointer(double),匹配cfunc的参数类型要求。- 使用
ravel()确保获取的是数组连续内存的起始地址(即使原数组是连续的,这一步也不会有性能损耗,只是视图转换)。
内容的提问来源于stack exchange,提问作者kesh
相关产品推荐
相关产品推荐

