如何沿指定轴将多维NumPy数组对半分割(舍去中间元素)
问题描述
我有一个如下所示的多维NumPy数组:
[ [ [1,2,3,4,5], [6,7,8,9,10], [11,12,13,14,15] ], [ [16,17,18,19,20], [21,22,23,24,25], [26,27,28,29,30] ] ]
希望编写一个函数,沿指定轴将其对半分割,当轴长度为奇数时舍去中间元素。例如:
- 调用
my_function(my_ndarray, 0)时得到:
[ [ [1,2,3,4,5], [6,7,8,9,10], [11,12,13,14,15] ] ]
- 调用
my_function(my_ndarray, 1)时得到:
[ [ [1,2,3,4,5] ], [ [16,17,18,19,20] ] ]
- 调用
my_function(my_ndarray, 2)时得到:
[ [ [1,2], [6,7], [11,12] ], [ [16,17], [21,22], [26,27] ] ]
最初尝试使用np.split()方法,但当轴长度为奇数时无法舍去中间元素。理论上可以用条件语句处理,但想了解更高效的解决方法。
高效解决方案
不用复杂的条件判断,直接利用NumPy的切片特性就能高效实现需求,核心思路是计算目标轴的半长(向下取整),然后对数组进行切片保留前半部分。
实现代码
import numpy as np def my_function(arr, axis): half_length = arr.shape[axis] // 2 # 构建切片元组:仅指定轴取前half_length个元素,其余轴保留全部内容 slices = tuple(slice(None) if i != axis else slice(half_length) for i in range(arr.ndim)) return arr[slices]
代码说明
- 计算半长:用整数除法
//直接得到向下取整的半长,不管轴长度是奇数还是偶数都能统一处理——偶数时正好取一半,奇数时自动舍去中间元素。 - 构建切片元组:通过生成器表达式创建对应每个维度的切片规则,只有目标轴执行截断,其他维度保持完整。
- 切片操作:NumPy的数组切片是高效的视图操作(不会复制原始数据),性能远优于条件判断后拆分的方式。
测试验证
用示例数组测试:
# 构造测试数组 my_ndarray = np.array([ [[1,2,3,4,5], [6,7,8,9,10], [11,12,13,14,15]], [[16,17,18,19,20], [21,22,23,24,25], [26,27,28,29,30]] ]) # 测试轴0 print(my_function(my_ndarray, 0)) # 输出符合预期,仅保留第一个外层元素组 # 测试轴1 print(my_function(my_ndarray, 1)) # 输出符合预期,每个外层组仅保留第一个子数组 # 测试轴2 print(my_function(my_ndarray, 2)) # 输出符合预期,每个最内层数组仅保留前2个元素
内容的提问来源于stack exchange,提问作者Yes
相关产品推荐
相关产品推荐

