如何修复Numba CUDA中argtypes参数弃用的DeprecationError?
修复Numba CUDA
argtypes/restype 弃用错误 错误原因是新版Numba CUDA已弃用restype和argtypes关键字参数,要求将函数签名作为第一个位置参数传入cuda.jit。
修复后的代码
from numba import cuda, uint32, f8 def mandel(x, y, max_iters): """ Given the real and imaginary parts of a complex number, determine if it is a candidate for membership in the Mandelbrot set given a fixed number of iterations. """ c = complex(x, y) z = 0.0j for i in range(max_iters): z = z*z + c if (z.real*z.real + z.imag*z.imag) >= 4: return i return max_iters # 使用字符串形式的完整签名作为第一个参数 mandel_gpu = cuda.jit('uint32(f8, f8, uint32)', device=True)(mandel)
关键修改点
- 删除原代码中的
restype=uint32和argtypes=[f8, f8, uint32]参数 - 将函数签名(返回类型
uint32+ 三个参数类型f8, f8, uint32)以字符串形式'uint32(f8, f8, uint32)'作为第一个位置参数传递给cuda.jit - 保留
device=True以声明这是一个CUDA设备函数
你也可以使用Numba类型元组的形式传递签名(效果一致):
mandel_gpu = cuda.jit((f8, f8, uint32), restype=uint32, device=True)(mandel)
不过更推荐字符串签名的写法,可读性更强且符合新版Numba的最佳实践。
内容的提问来源于stack exchange,提问作者Đỗ Như Vỹ
相关产品推荐
相关产品推荐

