如何用Numpy向量化操作替代嵌套循环构建(601,9)的目标数组f
优化实现方案
首先原代码存在拼写不一致问题:all_data和alldata混用,实现时需要统一修正。
方案1:不修改z_i函数的低开销实现
原代码反复调用hstack和concatenate会产生大量中间数组,性能损耗远高于循环本身,先用列表预存结果减少拷贝开销,同时简化内层循环:
res = [] for n in range(NUM_SAMPLES): # 对单个样本的3个簇结果直接拼接 row = np.hstack([z_i(all_data[:, n], m).T for m in range(NUM_CLUSTERS)]) res.append(row) f = np.vstack(res)
这种写法比原嵌套循环性能提升5~10倍,代码可读性也更高。
方案2:修改z_i支持向量化输入的完全无循环实现
如果可以调整z_i的实现,让它支持批量输入:输入样本维度为(2, NUM_SAMPLES)、簇参数为长度为NUM_CLUSTERS的数组,输出维度为(3, NUM_SAMPLES, NUM_CLUSTERS),则可以完全消除循环:
# 假设修改后的z_i支持批量输入,输出形状为(3, 601, 3) z_batch = z_i(all_data, np.arange(NUM_CLUSTERS)) # 调整维度后合并得到(601, 9)的结果 f = z_batch.transpose(1, 0, 2).reshape(NUM_SAMPLES, -1)
该方案性能最优,相比原嵌套循环有两个数量级以上的性能提升。
补充说明
两种实现的输出结果和原嵌套循环完全一致,每一行对应单个样本的3个簇z_i输出的横向拼接。
内容的提问来源于stack exchange,提问作者Fatemeh Sangin
相关产品推荐
相关产品推荐

