在Numba中调用NumPy histogram时,weights关键字参数为何失效?
使用Python 3.13.5、numpy 2.2.6和numba 0.61.2编写了如下脚本:
import numpy as np, numba as nb @nb.njit(fastmath=True) def f(a, b): return np.histogram(a, 10, weights=b) a = np.random.randint(0, 256, (100,)).astype(np.uint8) b = np.random.randint(0, 256, (100,)).astype(np.uint8) print(np.histogram(a, 10, weights=b)) # 无问题 print(f(a, b)) # 此处报错
运行后出现以下错误:
numba.core.errors.TypingError: Failed in nopython mode pipeline (step: nopython frontend)
No implementation of function Function(<function histogram at 0x7244459565c0>) found for signature:histogram(array(uint8, 1d, C), Literalint, weights=array(uint8, 1d, C))
There are 2 candidate implementations:
- Of which 2 did not match due to:
Overload in function 'np_histogram': File: numba/np/old_arraymath.py: Line 3942.
With argument(s): '(array(uint8, 1d, C), int64, weights=array(uint8, 1d, C))':
Rejected as the implementation raised a specific error:
TypingError: got an unexpected keyword argument 'weights'
raised from /home/paul/st-python/test-3.13/.venv/lib/python3.13/site-packages/numba/core/typing/templates.py:791During: resolving callee type: Function(<function histogram at 0x7244459565c0>)
During: typing of call at /home/paul/st-python/test-3.13/so-numba-hist.py (5)File "so-numba-hist.py", line 5:
def f(a, b):
return np.histogram(a, 10, weights=b)
^During: Pass nopython_type_inference
问题:为何Numba不接受合法的weights=参数?
原因
Numba对NumPy函数的支持并非完全覆盖,仅实现了NumPy部分函数的子集,且部分实现不支持原函数的全部参数。你使用的Numba 0.61.2版本中,np.histogram的内置实现不支持weights关键字参数,这就是报错提示“got an unexpected keyword argument 'weights'”的核心原因。
解决方法
有两种可行的处理方式:
- 手动实现带权重的直方图计算:在Numba装饰的函数内自行编写计算逻辑,绕开Numba不支持的参数。示例如下:
import numpy as np, numba as nb @nb.njit(fastmath=True) def weighted_histogram(a, bins, weights): min_val = a.min() max_val = a.max() bin_width = (max_val - min_val) / bins hist = np.zeros(bins, dtype=np.uint64) for val, w in zip(a, weights): bin_idx = int((val - min_val) / bin_width) # 处理边界值,避免索引越界 if bin_idx >= bins: bin_idx = bins - 1 hist[bin_idx] += w # 生成区间数组,和np.histogram输出格式对齐 edges = np.linspace(min_val, max_val, bins + 1) return hist, edges a = np.random.randint(0, 256, (100,)).astype(np.uint8) b = np.random.randint(0, 256, (100,)).astype(np.uint8) print(np.histogram(a, 10, weights=b)) print(weighted_histogram(a, 10, b))
- 使用
nb.objmode临时退出Numba编译模式:在需要调用np.histogram(weights=...)的地方,用nb.objmode包裹,让这部分代码回到普通Python解释器执行。示例:
import numpy as np, numba as nb @nb.njit(fastmath=True) def f(a, b): with nb.objmode(hist='uint64[:]', edges='float64[:]'): hist, edges = np.histogram(a, 10, weights=b) return hist, edges a = np.random.randint(0, 256, (100,)).astype(np.uint8) b = np.random.randint(0, 256, (100,)).astype(np.uint8) print(np.histogram(a, 10, weights=b)) print(f(a, b))
注意:objmode会带来一定性能损耗,适合对这部分代码性能要求不高的场景。
内容的提问来源于stack exchange,提问作者Paul Jurczak

