You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何反向还原Sklearn中train_test_split的拆分结果?

如何获取train_test_split拆分后样本的原数据索引?

嘿,这个问题我刚好碰到过!用train_test_split的时候默认不会返回原数据的索引,但其实有两个实用的办法能搞定,让你轻松找到拆分后样本对应的原数据位置~

方法一:拆分时同步保存原索引(推荐!)

这是最稳妥的方式,拆分前先生成原数据的索引数组,把它和X、Y一起传入train_test_split,这样就能直接拿到拆分后的索引列表:

  1. 生成原数据的索引:
    import numpy as np
    from sklearn.model_selection import train_test_split
    
    # 假设X是你的数据集,形状(10000,50,50)
    indices = np.arange(X.shape[0])  # 生成0到9999的索引数组
    
  2. 同步拆分数据、标签和索引:
    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
    )
    
  3. 现在要找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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.22 09:39:46