如何使用PyTorch将CSV数据集切分为train_loader和test_loader
修复方案
你现有代码存在4处需要调整的点:
- 特征列选取错误:你当前取前5列作为特征,不符合前6列是特征的要求
- 未使用归一化后的数据:你当前返回的特征是原始未归一化的数值
- 未按要求返回float32类型的输入和标签
- 缺少数据集随机切分、封装DataLoader的逻辑
以下是完整修改后的可运行代码:
import torch import torch.utils.data as data import torch.utils.data.dataset as dataset from torch.utils.data import random_split, DataLoader import numpy as np import pickle from sklearn.preprocessing import MinMaxScaler, StandardScaler class Nav_Dataset(dataset.Dataset): def __init__(self): self.data = np.genfromtxt('saved/training_data.csv', delimiter=',') # 修正:取前6列作为特征,最后一列作为标签 self.x = self.data[:, 0:6] self.y = self.data[:, [6]] # 归一化数据和保存scaler逻辑保留 self.scaler = MinMaxScaler() self.normalized_data = self.scaler.fit_transform(self.data) pickle.dump(self.scaler, open("saved/scaler.pkl", "wb")) def __len__(self): return len(self.y) def __getitem__(self, idx): if not isinstance(idx, int): idx = idx.item() # 修正:从归一化后的数据中取特征和标签,同时转换为float32 x = self.normalized_data[idx, 0:6].astype(np.float32) y = self.normalized_data[idx, [6]].astype(np.float32) sample = {'input': x, 'label': y} return sample class Data_Loaders(): def __init__(self, batch_size): self.nav_dataset = Nav_Dataset() # 修正:随机切分数据集,测试集占比可以按需调整,这里取20%作为测试集 dataset_size = len(self.nav_dataset) train_size = int(0.8 * dataset_size) test_size = dataset_size - train_size train_dataset, test_dataset = random_split(self.nav_dataset, [train_size, test_size]) # 封装为DataLoader self.train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True) self.test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False) def main(): batch_size = 16 data_loaders = Data_Loaders(batch_size) for idx, sample in enumerate(data_loaders.train_loader): _, _ = sample['input'], sample['label'] for idx, sample in enumerate(data_loaders.test_loader): _, _ = sample['input'], sample['label'] if __name__ == '__main__': main()
关键逻辑说明
- 切分比例可以按需调整,上述代码使用80%数据作为训练集、20%作为测试集
random_split会自动完成随机切分,无需手动打乱索引DataLoader的shuffle参数设置为训练集打开、测试集关闭,符合常规训练逻辑- 所有返回的输入和标签都已经转换为float32类型,满足自动判分要求
内容的提问来源于stack exchange,提问作者user11665646
相关产品推荐
相关产品推荐

