Python中按指定长度拆分多维数组的高效实现方法
问题分析与解决方案
先明确你原代码里的两个核心问题:
- 拆分轴选反了:要按行数拆分,
np.split的axis参数得设为0(你写了axis=1,这是按列拆分); - 拆分参数格式错了:
np.split的第二个参数是拆分点的索引位置,不是每个子数组的长度。需要把i转换成累积和的前N-1个值,才能作为正确的拆分位置。
修正后的高效实现
import numpy as np N = 100 i = np.random.poisson(10, N) v = np.random.uniform(0, 200, sum(i)) r = np.vstack([v] * 91).T # 此时r的形状为(sum(i), 91) # 生成拆分点:取累积和的前N-1个值(避免最后拆分出空数组) split_points = np.cumsum(i)[:-1] # 按行拆分目标数组 splitted_r = np.split(r, split_points, axis=0)
关键说明
- 正确性验证:每个子数组的行数完全匹配
i的对应元素,比如splitted_r[0].shape[0] == i[0]; - 性能最优:
np.split是numpy原生C扩展实现,比手动循环拆分快几个数量级; - 灵活性适配:只要保证
sum(i)等于目标数组对应轴的长度(按行拆分行数、按列拆分列数),代码可以适配任意形状的多维数组,只需调整axis参数即可。
比如要按列拆分某个数组arr,只需修改轴参数,同时确保sum(i)等于arr.shape[1]:
splitted_arr = np.split(arr, np.cumsum(i)[:-1], axis=1)
内容的提问来源于stack exchange,提问作者Dreamer93
相关产品推荐
相关产品推荐

