把 CSV 交给 PyTorch 训练,不能只把每行转成张量:还要保证特征顺序固定、标签类型正确,并确认每个批次里特征、标签和样本编号仍一一对应。下面用自定义 Dataset 读取三条表格样本,再用 DataLoader 检查完整批次和最后不足一批的样本。
下方数据完全是人为构造的接口演示,已在 Windows、Python 3.11、PyTorch 2.14.1 CPU 环境运行,不涉及真实个人信息或模型效果评估。读取器会把整表保存在内存中,适合入门和小表验证;它不是海量 CSV 的流式读取方案。

约定样本结构,先检查单行
输入列为 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