如何将PyTorch random_split得到的Dataset转为Pandas DataFrame?
解决方法
首先你原代码里有个小错误:len_train = len(df) - len(test) 这里test还没定义,应该改成len_train = len(df) - len_test。
接下来,random_split返回的是Subset对象,不是直接的数据集内容,所以不能直接用pd.DataFrame()转换。正确的做法是通过Subset的indices属性获取拆分后的索引,再从原DataFrame中提取对应行:
import pandas as pd import torch from torch.utils.data import Dataset, random_split # 读取数据 df = pd.read_csv('some_txt_file.txt', sep=' ') # 修正计算长度的错误 len_test = len(df) // 10 len_train = len(df) - len_test lengths = [len_train, len_test] # 拆分数据集 train_subset, test_subset = random_split(df, lengths, torch.Generator().manual_seed(42)) # 转换回DataFrame df_train = df.iloc[train_subset.indices].copy() df_test = df.iloc[test_subset.indices].copy()
原理说明
random_split生成的Subset对象内部保存了两个关键信息:原数据集(这里就是你的df)和拆分后对应的索引列表(存在indices属性里)。通过df.iloc[索引列表]就能精准提取出训练集和测试集的行数据,最后用.copy()避免后续操作触发链式索引警告。
内容的提问来源于stack exchange,提问作者CoolMathematician
相关产品推荐
相关产品推荐

