You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.22 00:35:25