如何在Numba中正确指定CUDA本地数组的参数类型?
问题
我在CUDA编译的Numba内核中声明了一个本地数组:
from numba import cuda, types @cuda.jit((types.void(), device=True)) def A(): arr = cuda.local.array(30, dtype=types.int64)
我需要将该数组传递给另一个CUDA编译的函数:
@cuda.jit((types.void(), device=True)) def A(): arr = cuda.local.array(30, dtype=types.int64) B(arr) @cuda.jit((types.void( # 此处应填写什么参数类型? ), device=True)) def B(arr): # 对arr进行操作 pass
我尝试了以下几种类型标注方式,但均无法正常运行,出现多种错误:
# 尝试1 @cuda.jit((types.void( types.Array(types.int64, 1, 'A') )), device=True) # 尝试2 @cuda.jit((types.void( types.Array(types.Literal, 1, 'A') )), device=True) # 尝试3 @cuda.jit((types.void( types.Array(types.Literal[int64](-1), 1, 'A') )), device=True) # 尝试4 @cuda.jit((types.void( types.Array(types.Literal[int](-1), 1, 'A') )), device=True) # 尝试5 @cuda.jit((types.void( types.Literal[int](-1) )), device=True)
错误示例包括:
numba.core.errors.TypingError: Failed in cuda mode pipeline (step: nopython frontend) Internal error at <numba.core.typeinfer.ArgConstraint object at 0x7fb691c9d950>. 'property' object has no attribute 'is_precise' During: typing of argument at /home/file.py (94) Enable logging at debug level for details.
或者:
Traceback (most recent call last): File "/home/file.py", line 43, in <module> types.Array(types.Literal[int], 1, 'A'), ~~~~~~~~~~~~~^^^^^ TypeError: type 'Literal' is not subscriptable
或者:
numba.core.errors.TypingError: Failed in cuda mode pipeline (step: nopython frontend) No implementation of function Function(<built-in function eq>) found for signature: >>> eq(array(int64, 1d, C), Literal[int](-1))
请问正确的参数类型应该是什么?
解决方案
正确的做法是使用types.Array类型,并指定匹配的内存布局标记。cuda.local.array创建的数组默认是'C'(C风格连续内存)布局,而非你尝试的'A'(任意布局)。
正确的函数B类型标注
from numba import cuda, types @cuda.jit((types.void(), device=True)) def A(): arr = cuda.local.array(30, dtype=types.int64) B(arr) @cuda.jit((types.void(types.Array(types.int64, 1, 'C')), device=True)) def B(arr): # 示例操作:给数组第一个元素赋值 arr[0] = 100
错误原因分析
- 尝试1:使用
'A'布局标记与本地数组实际的'C'布局不匹配,导致类型推断失败。 - 尝试2-4:
Literal类型用于表示固定字面量值,并非数组类型,完全不适用于此场景。 - 尝试5:错误地将数组类型标注为单个字面量值,类型完全不匹配。
另外注意:原代码中@cuda.jit装饰器存在括号遗漏问题(如@cuda.jit((types.void(), device=True)缺少闭合)),实际运行时需修正该语法错误。
内容的提问来源于stack exchange,提问作者Edy Bourne
相关产品推荐
相关产品推荐

