random.shuffle洗牌机器学习数据集时数据重复丢失,求解决方案
问题分析:numpy数组用random.shuffle洗牌异常的原因及正确方案
异常原因
你遇到的问题核心是Python标准库的random.shuffle()不适合直接处理numpy数组:
random.shuffle是为Python原生可变序列(比如列表)设计的,它通过交换序列中的元素实现洗牌。- 当传入numpy数组时,它会将数组的每一行视为一个"元素"进行交换,但numpy数组的行是原数组的**视图(view)**而非独立副本,洗牌过程中会出现元素引用错乱,最终导致数据重复、类别样本丢失这类异常。
正确的数据集洗牌方式
针对numpy数组,推荐以下几种可靠的洗牌方法:
1. 用numpy原生的np.random.shuffle()(原地洗牌)
这是最直接的方案,专门针对numpy数组设计,会原地打乱数组的第一维度(即样本维度):
np.random.shuffle(data)
2. 用np.random.permutation()(返回新数组)
如果不想修改原数组,这个方法会返回一个打乱后的新数组,原数组保持不变:
shuffled_data = np.random.permutation(data)
3. 用PyTorch的索引洗牌(适合深度学习场景)
如果你已经用到了PyTorch,可以通过生成随机索引来实现洗牌:
# 生成随机排列的索引 idx = torch.randperm(data.shape[0]) # 转换为tensor后索引,再转回numpy(如果需要) shuffled_data = torch.tensor(data)[idx].numpy()
4. 转成列表洗牌(不推荐,仅作参考)
如果一定要用random.shuffle,可以先把numpy数组转成Python列表,洗牌后再转回numpy数组(大样本下效率低):
data_list = data.tolist() random.shuffle(data_list) shuffled_data = np.array(data_list)
额外小提示
你生成样本时,Y0和Y1用了同一个X数组,这会导致两类样本的分布完全对应(只是均值偏移),如果需要更真实的随机样本,可以分别生成X0和X1:
X0 = np.random.randn(10,2) X1 = np.random.randn(10,2) Y0 = mu0 + X0@L.T Y1 = mu1 + X1@L.T
内容的提问来源于stack exchange,提问作者Scott
相关产品推荐
相关产品推荐

