高维数组下numpy.where返回异常?是理解偏差、Bug还是环境问题?
这绝对是NumPy旧版本中的已知Bug,既不是你对numpy.where的理解偏差,也不是系统环境的问题(Windows 10本身不存在这类限制),核心原因是你使用的Python 3.5.3对应的NumPy版本大概率比较老旧(比如1.10.x或更早),这类版本在处理高维数组(通常≥8维)时,numpy.where的内部索引计算逻辑存在错误,导致返回的各维度坐标数组被错误地填充为相同内容。
为什么不是理解偏差?
正常情况下,numpy.where(A)应该返回一个元组,元组长度等于数组的维度数,每个元素是一个一维数组,对应该维度下所有非零元素的索引。比如你给出的5维数组示例,每个维度的坐标数组内容都正确匹配你设置的非零位置,这说明你对numpy.where的用法理解完全正确。
解决方案
升级NumPy版本
Python 3.5.x最高支持到NumPy 1.18.x版本(后续NumPy版本不再兼容Python 3.5),你可以通过pip升级:pip install numpy==1.18.5 --upgrade升级后重新运行测试代码,
numpy.where就能正确返回各维度的坐标数组了。临时替代方案(不升级的情况下)
如果暂时无法升级NumPy,可以改用np.argwhere来获取正确坐标,它会返回一个二维数组,每行代表一个非零元素的完整坐标,转置后转换成元组即可,输出格式和numpy.where完全一致:import numpy as np for i in range(5,11): print("dims:", i) A = np.zeros([5]*i) print("shape:", A.shape) for j in range(10): c = np.random.randint(low=0, high=5, size=i) A[tuple(c)] = j print(tuple(c), ":", j) # 用argwhere替代where获取正确坐标 coords = tuple(np.argwhere(A).T) print(coords)这个方法在高维数组下也能稳定返回正确的坐标结果。
补充说明
这类高维数组的索引Bug在NumPy 1.11及以后的版本中已经被逐步修复,所以升级是最彻底的解决办法。如果你的项目允许,也可以考虑升级Python版本到3.6+,这样能支持更新的NumPy版本,避免更多旧版本的兼容性问题。
内容的提问来源于stack exchange,提问作者Raketenolli

