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

如何让Numba装饰的函数返回Numpy标量而非Python类型?

取消Numba njit函数默认转换Numpy标量为Python类型的方法

默认情况下,被njit装饰的函数返回的Numpy标量会被自动转换为Python的float/int类型,这与预期的返回类型约定不符。以下是具体示例、预期行为及解决方法:

代码示例

import numpy as np
import numba as nb
from numba import uint64

@nb.njit
def foo():
    return np.float32(1)

@nb.njit(uint64())
def bar():
    return np.float32(1)

float64_not_np_float32 = foo()
int64_not_np_uint64 = bar()

print(type(float64_not_np_float32))
print(foo.nopython_signatures)
print(type(int64_not_np_uint64))
print(bar.nopython_signatures)

实际输出

<class 'float'>
[() -> float32]
<class 'int'>
[() -> uint64]

预期行为

def foo():
    return np.float32(1)

np_float32 = foo()
print(type(np_float32))

预期输出

<class 'numpy.float32'>

解决方法

全局禁用自动转换

通过修改Numba的全局配置参数SCALAR_PYTHON_COMPAT为False,可以取消Numpy标量到Python标量的自动转换:

import numba as nb
import numpy as np

# 全局配置生效,后续njit函数返回原始Numpy标量类型
nb.config.SCALAR_PYTHON_COMPAT = False

@nb.njit
def foo():
    return np.float32(1)

result = foo()
print(type(result))  # 输出 <class 'numpy.float32'>

注意事项

  • 该配置是全局生效的,会影响所有后续定义的njit函数。
  • 对于显式指定了非Numpy标量类型(如示例中的uint64())的函数,即使修改配置,仍会按照签名强制返回对应的Python标量。
  • 禁用转换后,返回的Numpy标量行为与普通Numpy标量一致,部分依赖Python标量的场景可能需要额外适配。

内容的提问来源于stack exchange,提问作者Marcel de Haan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 15:20:25