numpy报错:多元素数组真值判断歧义,如何解决该问题?
解决numpy.linspace传入多维数组触发的ValueError问题
你遇到的这个问题我之前也碰到过,其实是numpy版本更新后,linspace内部的逻辑变化导致的。之前旧版本的numpy可能允许直接传入多维数组作为start和stop参数,但新版本在判断step是否为0的时候,因为step会是一个和输入同形状的数组,直接用step == 0做判断就触发了数组布尔值的歧义错误。
先看一下你遇到的完整报错信息:
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) <ipython-input-19-187bbe847597> in <module> ----> 1 t = np.linspace(np.zeros((2, 2)), np.ones((2, 2)), 20) ~\Anaconda3\lib\site-packages\numpy\core\function_base.py in linspace(start, stop, num, endpoint, retstep, dtype) 122 if num > 1: 123 step = delta / div --> 124 if step == 0: 125 # Special handling for denormal numbers, gh-5437 126 y /= div ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()
解决方法
其实我们可以利用numpy的广播机制来实现原来的需求——生成20个从全0到全1的2x2数组,每个数组的元素都是对应位置的线性插值,具体有两种简洁的写法:
方法一:一维序列扩展维度后广播
先生成一维的0到1线性序列,通过reshape扩展维度,再和目标形状的全1数组相乘,利用广播完成每个元素的插值:
import numpy as np # 生成一维序列并扩展为(20,1,1),与(2,2)的全1数组广播相乘 t = np.linspace(0, 1, 20).reshape(-1, 1, 1) * np.ones((2, 2))
方法二:直接添加维度实现广播
另一种更直观的写法,通过np.newaxis给一维序列添加维度,直接匹配目标数组的形状:
t = np.linspace(0, 1, 20)[:, np.newaxis, np.newaxis] * np.ones((2, 2))
这两种方法生成的结果和你原来期望的完全一致,而且都是numpy的向量化操作,效率很高。
为什么之前没问题?
旧版本的numpy(大概1.20版本之前)的linspace函数内部没有这个step == 0的判断逻辑,或者对多维数组的处理更宽松,所以当时可以直接传入多维的start和stop。新版本修复了一些边界情况的处理,但也导致了这种场景下的报错。
内容的提问来源于stack exchange,提问作者JH Y
相关产品推荐
相关产品推荐

