如何通过numpy数组索引加速for循环与数组广播操作
问题解答
你完全可以使用numpy的索引特性替换这段for循环,数据规模越大,性能提升越明显。
你的代码逻辑是给数组x的第i行、第idxs[i]列的位置赋值为values[i],正好适配numpy的整数数组索引规则,直接用行索引数组+列索引数组批量定位赋值即可,不需要逐次循环。
替换后的完整代码如下:
import numpy as np np.random.seed(2021) x = np.zeros((10, 3)) values = np.random.randn(10) idxs = np.random.randint(0, 3, (10,)) # 替代原for循环的向量化写法 x[np.arange(len(idxs)), idxs] = values print(x)
运行后输出的结果和你原for循环的输出完全一致:
[[ 0. 0. 1.48860905] [ 0. 0. 0.67601087] [-0.41845137 0. 0. ] [ 0. -0.80652081 0. ] [ 0.55587583 0. 0. ] [ 0. 0. -0.70550429] [ 0. 0. 1.13085826] [ 0.64500184 0. 0. ] [ 0. 0.10641374 0. ] [ 0.42215483 0. 0. ]]
这种写法的所有运算都在numpy底层C层面执行,避免了Python原生循环的对象遍历开销,当数据量达到十万、百万级别时,性能会比纯Python循环高几个数量级。
内容的提问来源于stack exchange,提问作者discort
相关产品推荐
相关产品推荐

