为何numpy.split的结果使用numpy.take比普通索引慢?
numpy.split生成的数组用numpy.take访问比普通索引慢? 这个问题的核心在于**numpy.split返回的是原数组的非连续视图**,以及take和普通索引完全不同的底层实现逻辑。让我一步步拆解:
1. 先搞清楚split返回的是什么:视图而非连续副本
当你用np.split(MyTwoArrays, 2, axis=2)[0]时,你得到的不是一个独立的、内存连续的新数组,而是原数组MyTwoArrays的视图。原数组的内存布局是C顺序(numpy默认),也就是元素按(i,j,0) → (i,j,1) → (i+1,j,0) → (i+1,j,1)这样的顺序存储。split出来的第一个数组,其实是从原数组中每隔一个元素取一个,这些元素在内存中是分散、不连续的。
你可以通过打印数组的strides属性验证这一点:
print("原数组strides:", MyTwoArrays.strides) print("split后数组strides:", MyArray.strides)
比如如果原数组是float64类型,原数组的strides会是(160000, 16, 8),而split后的数组strides会是(160000, 16, 16)——最后一维的步长变成了16,意味着每次访问下一个元素要跳过16字节(也就是原数组里的另一个元素),这直接说明内存不连续。
2. 普通索引MyArray[0]为什么快?
普通索引(比如[0])是直接利用数组的strides信息来计算目标元素的内存地址的。对于split后的视图数组,MyArray[0]本质上是通过strides直接定位到(0,0,0)对应的内存位置,不需要额外的遍历或临时数组,一步到位,所以速度极快。
3. numpy.take为什么慢?
take的设计目标是处理任意的索引集合(比如多个分散的索引值),它的底层逻辑不会针对单元素访问做特殊优化,反而会做这些额外操作:
- 默认情况下,
take会先把数组展平成一维,这对于非连续的视图来说,展平过程需要遍历所有元素的位置来构建一维索引映射,开销很大。 - 即使指定了
axis参数,take也不会利用视图的strides来直接定位元素,而是会按照索引值逐个从原数组中提取元素,对于非连续内存的视图,这意味着每次访问都要计算一次偏移量,远不如普通索引的直接地址计算高效。
补充验证:连续数组下的对比
你测试里的第一个MyArray = np.empty((10000,10000,1))是内存连续的数组,这时候take(0)和普通索引的速度差距很小——因为连续内存下,take可以直接批量读取,不需要额外的偏移计算。这也反过来证明了,非连续视图才是导致性能差距的关键。
内容的提问来源于stack exchange,提问作者bers

