如何反向还原Sklearn中train_test_split的拆分结果?
如何获取train_test_split拆分后样本的原数据索引?
嘿,这个问题我刚好碰到过!用train_test_split的时候默认不会返回原数据的索引,但其实有两个实用的办法能搞定,让你轻松找到拆分后样本对应的原数据位置~
方法一:拆分时同步保存原索引(推荐!)
这是最稳妥的方式,拆分前先生成原数据的索引数组,把它和X、Y一起传入train_test_split,这样就能直接拿到拆分后的索引列表:
- 生成原数据的索引:
import numpy as np from sklearn.model_selection import train_test_split # 假设X是你的数据集,形状(10000,50,50) indices = np.arange(X.shape[0]) # 生成0到9999的索引数组 - 同步拆分数据、标签和索引:
train_data, test_data, train_label, test_label, train_indices, test_indices = train_test_split( X, Y, indices, test_size=0.2, random_state=11, shuffle=True ) - 现在要找
train_data[123]对应的原索引,直接取train_indices[123]就可以了!这个方法完全不受数据重复的影响,绝对可靠。
方法二:拆分后反向匹配(仅适用于无重复样本)
如果已经拆分完数据,才想起要找原索引,那可以通过逐样本匹配的方式来查找,但注意:只有当你的数据中没有完全相同的样本时,这个方法才准确:
# 找到train_data[123]对应的原索引 match_mask = (X == train_data[123]).all(axis=(1, 2)) # 检查每个原样本是否和目标样本完全一致 original_idx = np.where(match_mask)[0][0] # 取出第一个匹配的索引
这里的all(axis=(1,2))是因为你的X是三维数组,要确保每个50x50的样本所有元素都完全匹配,再用np.where定位到原索引。
内容的提问来源于stack exchange,提问作者null
相关产品推荐
相关产品推荐

