Numba nopython模式使用np.isin触发编译报错如何解决?
问题根因
- Numba 的 nopython 模式暂不支持
np.isin函数,这是触发 TypingError 的直接原因 - 原代码中使用的
np.apply_along_axis虽然有部分兼容支持,但性能远低于原生 numpy 轴运算,也建议替换
修复方案
手动实现 np.isin 的筛选逻辑,同时将求均值逻辑替换为 numba 兼容的 mean(axis=1) 写法,修改后的可运行代码如下:
import numpy as np from numba import jit @jit(nopython=True) def from_node_to_centroid(npCoor, nodeList): """ Calculate centroids """ # 手动实现isin逻辑生成布尔索引 mask = np.zeros(npCoor.shape[0], dtype=np.bool_) node_set = set(nodeList) for i in range(npCoor.shape[0]): if npCoor[i, 0] in node_set: mask[i] = True elem = npCoor[mask] # 直接沿轴1求均值,替代apply_along_axis meanData = elem[:, 1:].mean(axis=1) return meanData # 坐标数组 # npCoor格式:索引, x坐标, y坐标, z坐标 npCoor=np.array([ [1,4,5,6], [2,1,5,5], [3,2,4,4] ]) # 需要计算x/y/z均值的节点索引列表 nodeList=[1,2] mean = from_node_to_centroid(npCoor, nodeList) print(mean)
运行结果
[5. 3.66666667]
运行结果和纯numpy版本完全一致,同时可以正常触发numba的nopython模式加速。
内容的提问来源于stack exchange,提问作者Dario Daddà Turres
相关产品推荐
相关产品推荐

