numpy内存引用难题:矢量化替代慢循环并获一致结果
问题描述
给定以下代码:
import numpy as np a = np.array([0,0,0,0,0,0,0,0]) b = np.array([4,2,3]) c = np.array([5,5,2]) for i, e in enumerate(c): a[e] += b[i] print(a) # 输出:[0 0 3 0 0 6 0 0] a = np.array([0,0,0,0,0,0,0,0]) b = np.array([4,2,3]) c = np.array([5,5,2]) a[c] += b[np.arange(len(b))] print(a) # 输出:[0 0 3 0 0 2 0 0]
实际场景中数组c的长度极长,Python for循环执行速度过慢,因此希望使用numpy矢量化操作替代循环,但直接采用上述索引赋值方式得到的结果与循环结果不一致,推测是内存引用机制导致的问题,请问如何修改才能让两种方式得到相同的结果?
解决方案
使用np.add.at()函数替代直接索引赋值,它会对重复索引对应的位置执行原地累加,完美复现循环的计算逻辑:
import numpy as np a = np.array([0,0,0,0,0,0,0,0]) b = np.array([4,2,3]) c = np.array([5,5,2]) np.add.at(a, c, b) print(a) # 输出:[0 0 3 0 0 6 0 0]
原因说明
直接使用a[c] += b时,numpy的执行流程是:
- 提取
a[c]生成临时副本数组 - 将临时数组与
b对应元素相加 - 把计算结果赋值回原数组
a的对应位置
当c中存在重复索引时,临时数组里的重复位置会被多次计算,但最终赋值回原数组时,仅保留最后一次的计算结果(比如示例中索引5先后被加4和2,但临时数组里该位置最终值为2,因此a[5]最终只得到2)。
而np.add.at()是原地操作,会遍历所有索引,对每个索引对应的a元素直接进行累加,完全模拟for循环的逐个累加逻辑,既能得到和循环一致的结果,又能保持矢量化操作的高效性,非常适合处理超长数组的场景。
内容的提问来源于stack exchange,提问作者string
相关产品推荐
相关产品推荐

