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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 02:24:03