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

如何为本地CSV数据创建torch.utils.data.Dataset并导入DataLoader

自定义PyTorch Dataset读取本地CSV文件实现方案

你只需要继承torch.utils.data.Dataset抽象类,实现3个核心内置方法即可完成自定义数据集的构建,后续可直接对接DataLoader使用。

前置依赖安装

需要用到pandas读取CSV文件,执行以下命令安装依赖:
pip install pandas torch numpy

自定义CSV Dataset实现代码

import torch
import pandas as pd
from torch.utils.data import Dataset, DataLoader

class CustomCSVDataset(Dataset):
    def __init__(self, csv_path, feature_columns=None, label_column=None, transform=None):
        """
        初始化方法
        :param csv_path: 本地CSV文件的绝对/相对路径
        :param feature_columns: 要作为特征使用的列名列表,不传则默认取除标签列外的所有列
        :param label_column: 要作为标签使用的列名,无监督场景可不传
        :param transform: 自定义特征预处理的可调用对象
        """
        self.data_df = pd.read_csv(csv_path)
        self.transform = transform
        self.label_column = label_column
        
        # 自动匹配特征列
        if feature_columns is None:
            if label_column is not None:
                self.feature_columns = self.data_df.columns.drop(label_column)
            else:
                self.feature_columns = self.data_df.columns
        else:
            self.feature_columns = feature_columns

    def __len__(self):
        """返回数据集总样本数量"""
        return len(self.data_df)

    def __getitem__(self, idx):
        """根据索引返回单条样本数据"""
        # 读取指定行的特征并转为张量
        feature = self.data_df.loc[idx, self.feature_columns].values.astype(float)
        if self.transform:
            feature = self.transform(feature)
        feature_tensor = torch.tensor(feature, dtype=torch.float32)

        # 存在标签列时返回特征+标签,否则只返回特征
        if self.label_column is not None:
            label = self.data_df.loc[idx, self.label_column]
            # 分类任务标签用long类型,回归任务用float类型
            label_tensor = torch.tensor(label, dtype=torch.long if isinstance(label, int) else torch.float32)
            return feature_tensor, label_tensor
        return feature_tensor

使用示例

实例化数据集

# 替换为你的本地CSV路径,假设CSV的标签列名为"target"
csv_dataset = CustomCSVDataset(
    csv_path="./your_local_data.csv",
    label_column="target"
)

封装为DataLoader

data_loader = DataLoader(
    csv_dataset,
    batch_size=32, # 可自定义批次大小
    shuffle=True, # 训练集建议打乱,验证/测试集设为False
    num_workers=2 # Windows环境运行报错可修改为0
)

遍历验证

for batch_x, batch_y in data_loader:
    print("批次特征维度:", batch_x.shape)
    print("批次标签维度:", batch_y.shape)
    break

特殊场景适配建议

  • 如果CSV包含字符串类型的分类特征,可在__init__方法中加入编码逻辑,比如用LabelEncoder将分类值转为数值类型再使用
  • 如果CSV文件过大无法一次性加载到内存,可在__init__中仅记录CSV的行索引,__getitem__时按需读取对应行数据,避免内存溢出
  • 可自定义transform传入特征归一化、标准化、数据增强等预处理逻辑,所有输入为特征数组、输出为处理后特征的可调用对象都可直接使用

内容的提问来源于stack exchange,提问作者onni white

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 16:54:01