numpy使用frompyfunc传数组参数报'float'不可下标错误怎么解决?
错误原因
np.frompyfunc 会自动对所有输入参数做元素级的遍历拆分,你传入shape为(2,8)的pval数组时,它会逐个取出数组内的单个float值传给getnpx的mt参数,此时执行mt[0]就相当于对float值做下标索引,自然抛出对应报错。
方案1:直接用原生numpy广播(优先推荐)
你的需求本身完全符合numpy的广播规则,不需要额外做矢量化包装,直接调用函数即可:
import numpy as np pval=np.array([[0., 0.,0., 0., 0.,0., 0., 0.], [0., 0., 0., 0., 0.,0., 0., 0.]]) def getnpx(mt, age): return mt[0] + age # 直接调用即可,不需要额外包装 print(getnpx(pval, 1)) # 输出:[1. 1. 1. 1. 1. 1. 1. 1.]
如果需要对多个age值批量计算,numpy广播也原生支持,比如传入多个age的数组:
# 批量计算age为1、2、3的结果 print(getnpx(pval, np.array([1,2,3])[:, np.newaxis]))
方案2:确需矢量化包装时用np.vectorize指定排除参数
如果你的实际业务逻辑比示例复杂,必须要做矢量化包装,可以用np.vectorize的excluded参数指定mt参数不参与元素级拆分:
def getnpx(mt, age): return mt[0] + age # 排除第一个参数(参数下标为0的mt)不做矢量化拆分 vgetnpx = np.vectorize(getnpx, excluded=[0]) print(vgetnpx(pval, 1)) # 输出:[1. 1. 1. 1. 1. 1. 1. 1.]
内容的提问来源于stack exchange,提问作者William
相关产品推荐
相关产品推荐

