Numba中使用b[:,None]实现numpy数组逐行相乘报错如何解决
报错原因
Numba的nopython模式暂不支持将None作为数组索引的维度占位符,你在Numpy中常用的[:, None]新增维度的写法,在@njit装饰器下没有对应的类型实现,因此会抛出getitem方法匹配失败的报错。
可行实现方案
以下三种写法都可以在@njit环境下正常运行,输出结果和原生Numpy写法完全一致:
- 方案1:使用
reshape调整数组维度(兼容性最高,推荐)
import numpy as np from numba import njit @njit() def fct(a,b): # 将一维的b调整为(行数, 1)的二维数组,符合广播规则 c = a * b.reshape((-1, 1)) return c a = np.array([[3,5,6,9,10],[4,7,8,11,12]]) b = np.array([-1,2]) A = fct(a, b) print(A)
- 方案2:使用
np.expand_dims显式新增维度
只需要修改函数内的计算逻辑即可:
@njit() def fct(a,b): c = a * np.expand_dims(b, axis=1) return c
- 方案3:显式循环实现(适合需要自定义计算逻辑的场景,Numba对循环有专门优化,效率足够)
@njit() def fct(a,b): rows, cols = a.shape c = np.empty_like(a) for row_idx in range(rows): c[row_idx, :] = a[row_idx, :] * b[row_idx] return c
内容的提问来源于stack exchange,提问作者user9393931
相关产品推荐
相关产品推荐

