使用Numpy高效为数组新增多维度并生成对应份数拷贝的最优方案
Numpy高效为数组新增多维度并生成对应份数拷贝的最优方案
你好呀!你的问题我太有共鸣了——用嵌套np.repeat处理高维度扩展时,中间数组的反复创建确实会拖慢速度,特别是当目标维度很大的时候(比如你这里的361x721x11)。下面给你几个比嵌套repeat高效得多的方案,亲测速度能提升不少:
方案一:广播赋值(最推荐,代码简洁且速度快)
原理是先创建一个符合目标形状的空数组,然后利用numpy的广播机制,把原数组一次性填充到所有新增维度的位置。这种方法只做一次内存分配,避免了中间数组的冗余拷贝:import numpy as np ll = np.ones((13,25)) # 先创建目标形状的空数组 ll_ijk = np.empty((13,25,361,721,11), dtype=ll.dtype) # 广播赋值,一步到位 ll_ijk[:] = ll[..., None, None, None]方案二:broadcast_to + copy(和方案一效率接近)
np.broadcast_to会创建一个原数组的广播视图(不占额外内存),然后我们再拷贝成实际的数组,效果和方案一差不多:ll_expanded = ll[..., None, None, None] ll_ijk = np.broadcast_to(ll_expanded, (13,25,361,721,11)).copy()如果不需要实际数组:直接用广播视图参与计算
如果你只是需要在后续计算中使用这个扩展后的数组,完全不用生成实际的拷贝——直接用ll[..., None, None, None]就行!numpy会自动在计算时进行广播,既省内存又快,比如:# 假设后续要和另一个形状为(361,721,11)的数组arr计算 result = ll[..., None, None, None] * arr
为什么原方法会慢?
每次调用np.repeat都会生成一个全新的中间数组(比如第一次repeat得到(13,25,361),第二次得到(13,25,361,721)),每一步都要分配新内存并拷贝数据,三次下来就多了两次不必要的内存操作。而上面的方案只做一次内存分配和一次数据填充,自然快很多。
备注:内容来源于stack exchange,提问作者Researcher R
相关产品推荐
相关产品推荐

