如何更简洁地编写numpy数组子集索引的遍历循环代码
问题:如下
for循环是否有更简洁的编写方式? 原始实现代码:
import itertools import numpy as np def f(a, b, c): # 复杂函数占位 print(a+b+c) a = np.arange(12).reshape(3, 4) for y, x in itertools.product(range(a.shape[0]-1), range(a.shape[1]-1)): f(a[y, x], a[y, x+1], a[y+1, x])
提问者表示自己尝试过其他实现,但写法反而更晦涩繁琐,例如:
it = np.nditer(a[:-1, :-1], flags=['multi_index']) for e in it: y, x = it.multi_index f(a[y, x], a[y, x+1], a[y+1, x])
回答
可以根据f的支持能力选更简洁的实现,两种方案都比原始写法和尝试的nditer写法可读性更好:
- 若
f仅支持单值输入、必须保留逐元素遍历逻辑:直接用np.ndenumerate处理切片后的数组,不需要手动生成索引序列、也不需要额外导入itertools:
for (y, x), val in np.ndenumerate(a[:-1, :-1]): f(val, a[y, x+1], a[y+1, x])
切片a[:-1, :-1]天然排除了最后一行、最后一列不需要遍历的元素,ndenumerate直接返回每个元素的坐标和值,逻辑和原始代码完全一致。
- 若
f支持numpy数组批量运算(numpy原生运算、逐元素处理的函数都支持):可以完全删除循环,用数组切片一次性取出所有待计算的值批量传入,性能会有数量级提升:
f(a[:-1, :-1], a[:-1, 1:], a[1:, :-1])
示例里的a+b+c原生支持数组运算,用这个写法可以直接得到所有位置的计算结果,不需要任何循环。
内容的提问来源于stack exchange,提问作者Paul Jurczak
相关产品推荐
相关产品推荐

