如何为本地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
相关产品推荐
相关产品推荐

