如何用Numpy或其他方法加速这段Python代码的运行?
优化方案:用Numpy向量化操作替代三重循环
原代码的三重嵌套循环效率极低——Python循环的单步开销很高,面对(144,192,256)规模的数组,循环会浪费大量时间。用Numpy的向量化操作可以彻底解决这个问题,同时保证结果完全一致:
步骤1:一次性计算所有位置的类别索引
直接对prediction数组的最后一维(axis=3)执行argmax,一次性得到所有位置的类别标记,替代逐个循环计算的过程:
d = np.argmax(prediction, axis=3)
此时d的形状为(144, 192, 256),和原循环中每个i,j,k对应的d完全一致。
步骤2:创建类别到目标值的映射
用数组直接建立类别索引到目标值的对应关系:
mapping = np.array([0, 30, 150, 250], dtype=int)
这里索引0对应值0,索引1对应30,以此类推,完美匹配原代码的if逻辑。
步骤3:生成最终结果数组
通过映射数组直接索引,再添加最后一维保证形状和原代码输出一致:
tt = mapping[d][..., np.newaxis]
[..., np.newaxis]会把d的形状从(144,192,256)扩展为(144,192,256,1),和原代码中tt的形状完全相同。
完整代码
# prediction.shape = (144, 192, 256, 4) d = np.argmax(prediction, axis=3) mapping = np.array([0, 30, 150, 250], dtype=int) tt = mapping[d][..., np.newaxis]
这种方式完全依赖Numpy的底层优化,运行速度会比原循环快几十甚至上百倍,且代码更简洁易读。
内容的提问来源于stack exchange,提问作者C1SKO
相关产品推荐
相关产品推荐

