如何像zip()一样用np.nditer()遍历不同维度数组的前N维
用
np.nditer替代zip处理不同维度数组,仅遍历前若干维度 我尝试将np.nditer()当作zip()使用,处理不同维度的数组,且仅遍历数组的前若干维度,但遇到了维度不匹配的报错。
最简示例
import numpy as np a_all = np.arange(6).reshape(2,3) idx_all = np.arange(12).reshape(2,3,2) for a, idx in np.nditer([a_all, idx_all]): print((a, idx))
运行后抛出错误:
ValueError: operands could not be broadcast together with shapes (2,3) (2,3,2)
实际应用场景
我有两个数据数组需要运算,同时还有一个用于其他数组的索引列表,尝试如下代码:
import numpy as np a_all = np.arange(6).reshape(2,3) b_all = np.arange(6).reshape(2,3) idx_all = ( ((0,0), (0,1), (0,2)), ((1,0), (1,1), (1,2)) ) result = np.zeros((2,3)) for a, b, idx in np.nditer([a_all, b_all, idx_all]): result[idx] += a*b
同样出现维度不匹配的错误。我推测问题在于np.nditer()会遍历idx_all的所有维度,但找不到仅限制遍历前两个维度的方法。
我不想使用zip(),因为需要嵌套两层循环:
for a_, b_, idx_ in zip(a_all, b_all, idx_all): for a, b, idx in zip(a_, b_, idx_): result[idx] += a*b
更贴合实际的示例
import numpy as np a_all = np.random.randn(2,3) b_all = np.random.randn(2) idx_all = ( ((1,1), (2,2)) ) result = np.zeros(2) for a, b, idx, res in np.nditer([a_all, b_all, idx_all, result], op_flags=['readwrite']): res += a[idx] + b
解决方法
要让np.nditer只遍历前N个维度,核心是让所有输入数组的遍历维度逻辑一致,可以通过以下两种方式实现:
方法1:扩展低维数组的维度,匹配遍历结构
对维度较少的数组,用[:, :, None]或np.expand_dims()扩展维度,让它和高维数组的前N个维度对齐,np.nditer会自动按对齐的维度遍历,多余维度做广播处理。
以最简示例为例:
import numpy as np a_all = np.arange(6).reshape(2,3)[:, :, None] # 扩展为(2,3,1) idx_all = np.arange(12).reshape(2,3,2) for a, idx in np.nditer([a_all, idx_all]): print((a, idx))
此时a_all和idx_all的前两个维度都是(2,3),nditer会按这个维度遍历,每次迭代对应一组a和该位置下的所有idx元素。
方法2:用op_axes指定遍历轴
通过nditer的op_axes参数,为每个数组指定哪些维度参与遍历,哪些维度保留为内部维度,配合external_loop flag可以返回对应维度的块。
示例:
import numpy as np a_all = np.arange(6).reshape(2,3) idx_all = np.arange(12).reshape(2,3,2) # op_axes=[None, [0,1,-1]]:对idx_all遍历前两个轴,第三个轴不展开 for a, idx in np.nditer([a_all, idx_all], flags=['external_loop'], op_axes=[None, [0,1,-1]]): print("a块:", a) print("idx块:", idx)
针对实际场景的优化代码
import numpy as np a_all = np.arange(6).reshape(2,3) b_all = np.arange(6).reshape(2,3) idx_all = np.array([ [(0,0), (0,1), (0,2)], [(1,0), (1,1), (1,2)] ]) result = np.zeros((2,3)) # 扩展a_all和b_all的维度,匹配idx_all的前两个维度 for a, b, idx in np.nditer([a_all[:, :, None], b_all[:, :, None], idx_all]): result[tuple(idx)] += a * b print(result)
更贴合实际示例的修改
import numpy as np a_all = np.random.randn(2,3) b_all = np.random.randn(2)[:, None] # 扩展为(2,1) idx_all = np.array([[(1,1), (2,2)]]) result = np.zeros(2)[:, None] # 扩展为(2,1) for a, b, idx, res in np.nditer([a_all, b_all, idx_all, result], op_flags=['readwrite']): res += a[tuple(idx)] + b print(result.flatten()) # 转回一维数组
内容的提问来源于stack exchange,提问作者domist07
相关产品推荐
相关产品推荐

