如何按y值高效拆分numpy数组?求更简洁实现方案
根据y的取值拆分NumPy数组的简洁实现
问题描述
现有以下NumPy数组:
import numpy as np x = np.array([[2,3,5,6], [1,2,4,3], [1,5,6,4], [2,8,9,5]]) y = np.array([1,0,1,2])
需要依据y的取值将x拆分为多个子数组,最终期望输出:
x_0 = np.array([1,2,4,3]) x_1 = np.array([[2,3,5,6], [1,5,6,4]]) x_2 = np.array([2,8,9,5])
已尝试用字典、列表实现,寻求更简洁的方案。
简洁实现方案
方案1:字典推导式(推荐,适配任意数量的y取值)
通过np.unique获取y的唯一取值,再用字典推导直接生成分组:
import numpy as np x = np.array([[2,3,5,6], [1,2,4,3], [1,5,6,4], [2,8,9,5]]) y = np.array([1,0,1,2]) # 生成分组字典,键为y的取值,值为对应x子数组 grouped_x = {val: x[y == val] for val in np.unique(y)}
此时grouped_x[0]对应目标中的x_0,grouped_x[1]对应x_1,以此类推。若需要单独赋值给变量,可直接:
x_0, x_1, x_2 = grouped_x.values()
字典存储的方式更灵活,能适配x、y维度或元素数量变化的场景。
方案2:结合排序与np.split
先对y和x按y值排序,再计算拆分点完成拆分:
import numpy as np x = np.array([[2,3,5,6], [1,2,4,3], [1,5,6,4], [2,8,9,5]]) y = np.array([1,0,1,2]) # 按y值排序对应的索引 sorted_indices = np.argsort(y) sorted_x = x[sorted_indices] sorted_y = y[sorted_indices] # 找到y值发生变化的位置,作为拆分点 split_points = np.where(np.diff(sorted_y))[0] + 1 # 拆分数组 x_0, x_1, x_2 = np.split(sorted_x, split_points)
该方法会按y的升序输出拆分后的数组,结果与期望一致。
内容的提问来源于stack exchange,提问作者bobr
相关产品推荐
相关产品推荐

