PyTorch 怎么读取 CSV 训练数据?自定义 Dataset 并检查 DataLoader 批次

把本地 CSV 数值特征和类别标签转为 PyTorch Dataset,检查缺列、异常值与重复编号,并用 DataLoader 核对末批和样本数量。

把 CSV 交给 PyTorch 训练,不能只把每行转成张量:还要保证特征顺序固定、标签类型正确,并确认每个批次里特征、标签和样本编号仍一一对应。下面用自定义 Dataset 读取三条表格样本,再用 DataLoader 检查完整批次和最后不足一批的样本。

下方数据完全是人为构造的接口演示,已在 Windows、Python 3.11、PyTorch 2.14.1 CPU 环境运行,不涉及真实个人信息或模型效果评估。读取器会把整表保存在内存中,适合入门和小表验证;它不是海量 CSV 的流式读取方案。

PyTorch 怎么读取 CSV 训练数据?自定义 Dataset 并检查 DataLoader 批次

约定样本结构,先检查单行

输入列为 id,f1,f2,label。f1、f2 是按这个顺序输入模型的两个数值特征;label 只能是 0 或 1;id 只作核对用途,不送入模型。演示会创建专用 demo_samples.csv,再次运行会覆盖它。

根据 PyTorch 自定义数据集教程,初始化负责准备数据,__len__() 返回样本数量,__getitem__() 按索引返回一个样本。我们额外检查缺列、空编号、重复编号、未知标签及非有限特征,让异常在读取时暴露。

完整代码:从 CSV 到可核对批次

在独立环境安装下列 CPU 依赖,把代码保存为 csv_loader_demo.py,执行 python csv_loader_demo.py。

python -m pip install torch==2.14.1 numpy --index-url https://download.pytorch.org/whl/cpu
import csv
import math
from pathlib import Path
import torch
from torch.utils.data import Dataset, DataLoader

class CsvDataset(Dataset):
    def __init__(self, path):
        self.samples = []
        seen = set()
        with open(path, encoding='utf-8-sig', newline='') as f:
            reader = csv.DictReader(f)
            required = {'id', 'f1', 'f2', 'label'}
            if not required.issubset(set(reader.fieldnames or [])):
                raise ValueError('CSV 缺少 id/f1/f2/label 列')
            for line, row in enumerate(reader, start=2):
                sample_id = (row['id'] or '').strip()
                label = (row['label'] or '').strip()
                if not sample_id or sample_id in seen:
                    raise ValueError(f'第 {line} 行 ID 为空或重复')
                if label not in {'0', '1'}:
                    raise ValueError(f'第 {line} 行标签必须为 0 或 1')
                try:
                    values = [float(row['f1']), float(row['f2'])]
                except (ValueError, TypeError) as exc:
                    raise ValueError(f'第 {line} 行特征不是数字') from exc
                features = torch.tensor(values, dtype=torch.float32)
                if not all(math.isfinite(v) for v in values) or not torch.isfinite(features).all():
                    raise ValueError(f'第 {line} 行特征含非有限值')
                seen.add(sample_id)
                self.samples.append((features, torch.tensor(int(label), dtype=torch.long), sample_id))
        if not self.samples:
            raise ValueError('CSV 没有有效样本')

    def __len__(self):
        return len(self.samples)

    def __getitem__(self, index):
        return self.samples[index]

def main():
    # 只创建专用演示文件,换真实数据时删去本段写文件代码。
    path = Path('demo_samples.csv')
    path.write_text('id,f1,f2,label\ns01,1.0,2.0,0\ns02,2.0,1.0,1\ns03,3.0,4.0,0\n', encoding='utf-8')
    data = CsvDataset(path)
    loader = DataLoader(data, batch_size=2, shuffle=False, num_workers=0, drop_last=False)
    all_ids = []
    for features, labels, ids in loader:
        assert features.ndim == 2 and features.shape[1] == 2
        assert labels.ndim == 1 and labels.dtype == torch.long
        assert features.shape[0] == labels.shape[0] == len(ids)
        print('batch:', tuple(features.shape), tuple(labels.shape), list(ids))
        all_ids.extend(ids)
    assert all_ids == ['s01', 's02', 's03']
    print('rows:', len(data), 'seen:', len(all_ids))

if __name__ == '__main__':
    main()

第一个返回值是长度为 2 的 float32 特征;第二个是 long 类型类别编号;第三个保留字符串编号。DataLoader 官方文档说明默认批处理会给张量增加批次维度,字符串则保留为可逐个检查的编号序列。

先设 shuffle=False 便于核对原始行顺序,num_workers=0 使用同一进程读取,报错更容易定位。drop_last=False 保留最后不足一批的样本。正式训练是否随机打乱,要结合任务和采样规则决定;不要为了观察固定顺序而永久关闭本应需要的训练随机采样。

哪些输出证明读取正确

这组 CPU 示例的实际输出如下:

batch: (2, 2) (2,) ['s01', 's02']
batch: (1, 2) (1,) ['s03']
rows: 3 seen: 3

第一批两行、第二批一行,每行都有两个特征;标签数量和编号数量与批次大小相同;三个编号按原始顺序出现一次。把真实表接入时,先打印 data[0] 并与原 CSV 第一条有效记录人工核对,再遍历一个完整轮次,检查总量与编号,最后才传入训练循环。

替换真实数据前,删掉 main() 中创建演示文件的两行,将 path 指向自己的 CSV,并按实际字段调整列名与标签规则。先用少量已知答案的记录验证;不要只改文件路径,就默认不同业务的特征、类别和损失函数都适配本例。

常见失败与处理边界

  • 缺列或标签不是 0/1:读取器会直接报错并尽量定位记录行号。若业务有三类以上标签,应建立明确的类别到编号映射,并同步修改模型输出数量。
  • 空值、NaN、Infinity:本例拒绝,不默默变成零。先根据业务决定缺失值处理规则;浮点值还须落在 float32 可表示范围内。
  • 最后少了一个样本:先查 drop_last,再查自定义采样器和读取过滤。当前演示明确保留末批。
  • 同一批不同样本长度不一:默认堆叠可能失败。文本变长序列应设计填充或自定义 collate_fn,不能直接套用固定两维特征例子。
  • Windows 调高工作进程后重复执行或报启动错误:官方文档要求兼容多进程的主入口保护与顶层数据集定义。本例保留主入口,但未测试多进程加速;先在 num_workers=0 跑通,再单独测试工作进程配置。

Dataset 读取成功只是训练输入链路的验证,不能说明数据没有泄漏或模型有用。真实项目仍需独立划分训练与验证数据,明确特征归一化、标签质量与数据许可;这些条件应在进入模型训练前检查。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30463.html

赞 (0)
AI小管家的头像AI小管家
Netron 怎么查看模型结构?打开 ONNX 并核对输入输出
上一篇 4小时前
scikit-learn 模型怎么转 ONNX?转换后比较预测结果
下一篇 4小时前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部