如何按第一列相同值分组后批量打乱数组行?
解决numpy数组打乱时保持同第一列值行连续的问题
我之前做数据预处理的时候刚好碰到过一模一样的需求,直接用numpy.random.shuffle()确实不行——它会把所有行彻底打乱,完全不管分组。其实咱们可以换个思路,先按第一列的值分组,再打乱分组的顺序(组内的行可以选打乱或不打乱),最后把这些分组拼起来就行,具体实现看下面:
核心思路
- 先把数组按第一列的取值分成若干个独立分组
- 随机打乱这些分组的整体排列顺序
- (可选)每个组内部的行也可以随机打乱(不影响同组连续的要求)
- 把打乱后的各组重新拼接成完整数组
代码实现
import numpy as np # 生成测试数组方便验证效果 np.random.seed(42) # 固定随机种子,结果可复现 m, n = 12, 3 # 第一列用0、1、2的重复值,模拟题目中存在相同值的场景 test_arr = np.hstack([np.random.randint(0, 3, size=(m,1)), np.random.rand(m, n-1)]) print("原始数组:") print(test_arr) # 步骤1:提取第一列,获取所有唯一分组值 first_column = test_arr[:, 0] unique_groups = np.unique(first_column) # 步骤2:随机打乱分组的顺序 shuffled_group_order = np.random.permutation(unique_groups) # 步骤3:按新的分组顺序拼接,同时可选打乱组内行 final_result = [] for group_val in shuffled_group_order: # 取出当前分组的所有行 group_data = test_arr[first_column == group_val] # 可选操作:打乱组内的行顺序(不需要的话直接删掉这行) np.random.shuffle(group_data) final_result.append(group_data) # 把所有分组垂直拼接成最终数组 shuffled_arr = np.vstack(final_result) print("\n处理后的数组(同第一列值的行保持连续,整体顺序打乱):") print(shuffled_arr)
代码说明
np.unique()用来快速提取第一列的所有唯一值,这就是我们的分组依据np.random.permutation()用来打乱分组的顺序,它会返回新的打乱后的数组,不会修改原分组列表,比直接用shuffle()更稳妥- 循环里的
np.random.shuffle(group_data)是可选操作:如果希望同组内的行也随机打乱就保留,要是想让同组内的行保持原始相对顺序就删掉这一行 - 最后用
np.vstack()把所有分组拼接起来,得到符合要求的最终数组
这样处理后,既能保证第一列值相同的行始终连续,又能让整个数组的行顺序(分组层面)是随机打乱的,完全满足你的需求~
内容的提问来源于stack exchange,提问作者Mathews24
相关产品推荐
相关产品推荐

