为何在循环中使用numpy shuffle无法打乱数组指定列?
问题分析与解决
你的代码核心问题是错误操作了原数组X_test的元素,而非拷贝后的X_test_s:
- 你先通过
X_test_s = X_test.copy()得到原数组的拷贝,但循环里tempmat = X_test[i]取的是原数组的视图,对tempmat的修改会直接改变原数组X_test,之后再把这个修改后的视图赋值给X_test_s[i],相当于把原数组的修改结果同步到拷贝数组,完全没起到单独打乱X_test_s的作用。 - 另外,
np.random.shuffle是原地打乱数组,temp2 = tempmat[:,1]得到的是视图,执行shuffle后已经直接修改了tempmat的对应列,不需要额外执行tempmat[:,1] = temp2。
修正后的简化代码
X_test_s = X_test.copy() for i in range(len(X_test_s)): # 直接操作拷贝数组的第i个样本的第2列(索引为1) np.random.shuffle(X_test_s[i][:, 1])
执行这段代码后,X_test_s[3][:,1]会输出打乱后的结果,且原数组X_test的对应列不会被修改。
更高效的向量化实现(避免循环)
如果数据量较大,可通过numpy向量化操作提升效率:
X_test_s = X_test.copy() # 为每个样本的第2列生成随机排序索引 shuffled_indices = np.random.rand(X_test_s.shape[0], X_test_s.shape[1]).argsort(axis=1) # 根据索引重新排列第2列元素 X_test_s[:, :, 1] = X_test_s[np.arange(X_test_s.shape[0])[:, None], shuffled_indices, 1]
内容的提问来源于stack exchange,提问作者eng
相关产品推荐
相关产品推荐

