如何将自定义CSV数据集划分为训练集与测试集?
嘿,我来帮你解决这个问题~其实两种方式都可行:既可以在现有类内部实现训练/测试集的划分逻辑,也可以在外部划分好数据后再用同一个类分别加载。先给你纠正下原代码里的一个关键问题,再分别讲两种方案:
先修正原代码的核心错误
你的__getitem__方法写得不对哦!Dataset类的__getitem__应该根据传入的index返回单个样本,而不是遍历所有样本返回整个数据集。另外pd.get_dummies().as_matrix()已经被Pandas弃用了,改用.values或者.to_numpy()更稳妥。修正后的基础版本如下:
import pandas as pd import numpy as np import cv2 from torch.utils.data.dataset import Dataset class CustomDatasetFromCSV(Dataset): def __init__(self, csv_path, transform=None): self.data = pd.read_csv(csv_path) # 替换弃用的as_matrix() self.labels = pd.get_dummies(self.data['emotion']).values self.height = 48 self.width = 48 self.transform = transform def __getitem__(self, index): # 获取单个样本的像素序列 pixel_sequence = self.data.iloc[index]['pixels'] # 解析像素为数组 face = np.array([int(p) for p in pixel_sequence.split(' ')], dtype=np.uint8) # 重塑为48x48的灰度图 face = face.reshape(self.height, self.width) # 转为float32并增加通道维度(适配PyTorch的输入格式) face = face.astype(np.float32) face = np.expand_dims(face, axis=-1) # 获取对应标签 label = self.labels[index] # 应用transform if self.transform is not None: face = self.transform(face) return face, label def __len__(self): return len(self.data)
方案1:在类内部实现划分逻辑
只需要给__init__方法添加几个参数,控制是否加载训练集/测试集、划分比例等,这样不用额外写新类,实例化两次就能得到训练和测试数据集:
import pandas as pd import numpy as np import cv2 from torch.utils.data.dataset import Dataset from sklearn.model_selection import train_test_split # 需要导入这个 class CustomDatasetFromCSV(Dataset): def __init__(self, csv_path, transform=None, is_train=True, train_ratio=0.8, random_state=42): self.data = pd.read_csv(csv_path) # 按比例划分训练/测试索引,stratify保证类别分布一致 train_indices, test_indices = train_test_split( self.data.index, train_size=train_ratio, random_state=random_state, stratify=self.data['emotion'] ) # 根据is_train选择对应的子集 if is_train: self.data = self.data.loc[train_indices] else: self.data = self.data.loc[test_indices] # 处理标签 self.labels = pd.get_dummies(self.data['emotion']).values self.height = 48 self.width = 48 self.transform = transform def __getitem__(self, index): pixel_sequence = self.data.iloc[index]['pixels'] face = np.array([int(p) for p in pixel_sequence.split(' ')], dtype=np.uint8) face = face.reshape(self.height, self.width) face = face.astype(np.float32) face = np.expand_dims(face, axis=-1) label = self.labels[index] if self.transform is not None: face = self.transform(face) return face, label def __len__(self): return len(self.data)
使用方式:
# 实例化训练集 train_dataset = CustomDatasetFromCSV( 'your_data.csv', transform=train_transform, is_train=True ) # 实例化测试集 test_dataset = CustomDatasetFromCSV( 'your_data.csv', transform=test_transform, is_train=False )
方案2:外部划分后再用类加载
如果觉得把划分逻辑放在类外面更清晰,也可以先在外部把数据分成训练和测试两部分,再用同一个类分别加载:
方式A:保存为两个CSV文件
# 外部划分代码 import pandas as pd from sklearn.model_selection import train_test_split data = pd.read_csv('your_data.csv') # 划分数据,保证类别分布 train_data, test_data = train_test_split( data, train_size=0.8, random_state=42, stratify=data['emotion'] ) # 保存为CSV train_data.to_csv('train_data.csv', index=False) test_data.to_csv('test_data.csv', index=False) # 然后用你的类加载 train_dataset = CustomDatasetFromCSV('train_data.csv', transform=train_transform) test_dataset = CustomDatasetFromCSV('test_data.csv', transform=test_transform)
方式B:直接传入划分好的DataFrame(更高效,不用存文件)
可以稍微修改类的__init__,让它支持传入文件路径或者已有的DataFrame:
class CustomDatasetFromCSV(Dataset): def __init__(self, data_source, transform=None): # 支持传入CSV路径或者DataFrame if isinstance(data_source, str): self.data = pd.read_csv(data_source) elif isinstance(data_source, pd.DataFrame): self.data = data_source else: raise ValueError("data_source必须是CSV文件路径或者pandas DataFrame") self.labels = pd.get_dummies(self.data['emotion']).values self.height = 48 self.width = 48 self.transform = transform # __getitem__和__len__和之前一致,省略
使用方式:
# 外部划分得到DataFrame data = pd.read_csv('your_data.csv') train_data, test_data = train_test_split( data, train_size=0.8, random_state=42, stratify=data['emotion'] ) # 直接传入DataFrame train_dataset = CustomDatasetFromCSV(train_data, transform=train_transform) test_dataset = CustomDatasetFromCSV(test_data, transform=test_transform)
两种方案的选择
- 如果希望类的复用性更强,不用每次都写外部划分代码,选方案1;
- 如果希望划分逻辑和数据集类解耦,代码结构更清晰,选方案2。
内容的提问来源于stack exchange,提问作者nirvair
相关产品推荐
相关产品推荐

