本节你会学到
- 说清楚「自定义 Dataset:怎样封装自己的训练数据」解决的核心问题
- 知道它在「PyTorch 深度学习基础(43~58)」中的位置
- 把 DataLoader、Linear 层、Tensor 这些关键词联系起来
PyTorch 内置数据集只够入门练手。真实项目里你的数据可能是 CSV、文件夹图片、JSON 文本——这时候必须自己写 Dataset。这篇把 __len__ 和 __getitem__ 讲透,附带三个真实场景。
学习路线:PyTorch 深度学习基础(43~58) · 第 48 课
在前 20 篇 PyTorch 入门上继续深化,补齐张量操作、自动微分、优化器家族和训练排错能力。
学完本阶段你能做到:独立搭建一个 PyTorch 训练脚本,正确处理 Shape、device、梯度清空、优化器选择和学习率策略,并能根据 loss/accuracy 曲线诊断训练问题。
推荐读法:当前更新主线,建议跟更;每篇都要配合代码自检。
查看完整阶段 · 16 篇先看目标,再带着问题读正文。读完后用练习确认自己真的理解了。
先检查模型参数、输入 Tensor、标签 Tensor 是否在同一个设备上。常见修正是把它们统一 `.to(device)`。
因为训练是否正确经常取决于 shape、dtype、device 和梯度流,API 名字只能告诉你工具,不会保证数据契约正确。
`输入数据 → forward → loss → backward → optimizer.step()`,并知道本文主题位于这条链路的哪一环。
你跟着教程跑 MNIST、CIFAR-10,一切顺利——因为 torchvision.datasets 帮你把脏活干完了。但真到了自己的项目:数据是文件夹里的图片、是带脏数据的 CSV、是几十个 JSON 文件,你才发现根本不知道怎么把数据喂进模型。
这一篇就解决这个问题:自己写 Dataset。
第 17 篇我们第一次接触 DataLoader,它负责批量读取、打乱、并行加载。第 47 篇(上一篇)我们讲了张量形状转换。今天往数据上游走一步——DataLoader 吃的是什么?是 Dataset。Dataset 定义“数据长什么样、怎么取一条”,DataLoader 定义“怎么批量地取”。
自定义 Dataset 就是继承
torch.utils.data.Dataset,实现两个方法:__len__告诉 PyTorch 数据有多少条,__getitem__告诉它第 i 条数据是什么。
from torch.utils.data import Dataset
class MyDataset(Dataset):
def __init__(self, features, labels):
"""初始化时把数据存进来"""
self.features = features
self.labels = labels
def __len__(self):
"""返回数据总量"""
return len(self.features)
def __getitem__(self, index):
"""返回第 index 条数据 (x, y)"""
x = self.features[index]
y = self.labels[index]
return x, y
就这么多。__init__ 存数据,__len__ 报数量,__getitem__ 按索引取一条。然后就能丢给 DataLoader:
from torch.utils.data import DataLoader
import torch
features = torch.randn(1000, 20) # 1000 条,每条 20 特征
labels = torch.randint(0, 3, (1000,)) # 3 分类
dataset = MyDataset(features, labels)
loader = DataLoader(dataset, batch_size=32, shuffle=True)
for x, y in loader:
print(x.shape, y.shape) # [32, 20] [32]
break
__len__(self) → int告诉 PyTorch 这个数据集有多少条数据。DataLoader 要靠它算“一个 epoch 有多少个 batch”。
def __len__(self):
return len(self.features) # 数据量
__getitem__(self, index) → (x, y)返回第 index 条数据。这是最灵活的地方——你可以在这里做任何预处理。
| 输入 | 输出 | 含义 |
|---|---|---|
index(int) |
(x, y) 元组 |
第 index 条的特征和标签 |
返回的 x 和 y 最好是 Tensor,DataLoader 会自动把它们拼成 batch。
import pandas as pd
import torch
from torch.utils.data import Dataset
class CSVDataset(Dataset):
def __init__(self, csv_path, has_label=True):
df = pd.read_csv(csv_path)
if has_label:
self.features = torch.tensor(df.iloc[:, :-1].values).float()
self.labels = torch.tensor(df.iloc[:, -1].values).long()
else:
self.features = torch.tensor(df.values).float()
def __len__(self):
return len(self.features)
def __getitem__(self, index):
x = self.features[index]
y = self.labels[index]
return x, y
dataset = CSVDataset("data.csv")
print(len(dataset)) # 数据条数
print(dataset[0]) # 第一条 (x, y)
from PIL import Image
from pathlib import Path
import torch
from torch.utils.data import Dataset
from torchvision import transforms
class ImageFolderDataset(Dataset):
def __init__(self, root, transform=None):
self.root = Path(root)
# 假设结构:root/类别名/图片.jpg
self.samples = []
self.class_to_idx = {}
for idx, class_dir in enumerate(sorted(self.root.iterdir())):
if class_dir.is_dir():
self.class_to_idx[class_dir.name] = idx
for img_path in class_dir.iterdir():
if img_path.suffix in ('.jpg', '.png'):
self.samples.append((img_path, idx))
self.transform = transform or transforms.ToTensor()
def __len__(self):
return len(self.samples)
def __getitem__(self, index):
img_path, label = self.samples[index]
img = Image.open(img_path).convert('RGB')
img = self.transform(img) # 在这里做预处理
return img, label
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
])
dataset = ImageFolderDataset("data/images", transform=transform)
关键点:图片在 __getitem__ 里才读取——不要在 __init__ 里把所有图片读进内存,会爆。
import jieba
import torch
from torch.utils.data import Dataset
class TextDataset(Dataset):
def __init__(self, texts, labels, vocab, max_len=100):
self.texts = texts
self.labels = labels
self.vocab = vocab # 词到 ID 的映射
self.max_len = max_len
def __len__(self):
return len(self.texts)
def __getitem__(self, index):
# 在这里分词、转 ID、截断/填充
words = jieba.lcut(self.texts[index])
ids = [self.vocab.get(w, 0) for w in words][:self.max_len]
# 填充到固定长度
if len(ids) < self.max_len:
ids += [0] * (self.max_len - len(ids))
x = torch.tensor(ids).long()
y = self.labels[index]
return x, y
关键点:分词、转 ID、Padding 这些预处理放在 __getitem__ 里,DataLoader 的多进程能并行加速。
Dataset.__getitem__(index)
输入:index (int)
输出:(x, y)
x: 单个样本的特征,如 [F] 或 [C, H, W]
y: 单个标签,如标量
经过 DataLoader(batch_size=B) 后:
x: [B, F] 或 [B, C, H, W] ← 自动加了 batch 维
y: [B]
上一篇我们说过“DataLoader 后最前面一定多一个 B”——这就是 Dataset 和 DataLoader 的分工:Dataset 吐单条,DataLoader 拼成 batch。
# ❌ 10 万张图片直接读进内存
def __init__(self, root):
self.images = [Image.open(p) for p in all_paths] # 内存爆炸
修复:__init__ 只存路径,__getitem__ 里才读图。
# ❌ 有的返回 numpy,有的返回 tensor
def __getitem__(self, index):
if random.random() > 0.5:
return np.array(x), y # numpy
return torch.tensor(x), y # tensor
DataLoader 拼批量时类型不一致会报错。修复:统一返回 Tensor。
# ❌ 返回 Python int
def __getitem__(self, index):
return self.features[index], int(self.labels[index])
DataLoader 拼出来的 batch 标签可能是 list 而不是 Tensor,下游算 loss 报错。修复:返回 torch.tensor(label)。
练习 1:写一个 Dataset,从 JSON 文件读取数据(每行一个 {"text": "...", "label": 1}),__getitem__ 返回文本长度和标签。
练习 2:下面的 Dataset 哪里有问题?
class BadDataset(Dataset):
def __init__(self, data):
self.data = data
def __getitem__(self, index):
return self.data[index]
练习 3:训练集和测试集怎么用同一个 Dataset 类分开?写出代码思路。
练习 1:
import json
import torch
from torch.utils.data import Dataset
class JsonDataset(Dataset):
def __init__(self, json_path):
with open(json_path, 'r', encoding='utf-8') as f:
lines = f.readlines()
self.data = [json.loads(line) for line in lines]
def __len__(self):
return len(self.data)
def __getitem__(self, index):
item = self.data[index]
x = torch.tensor(len(item["text"])).float()
y = torch.tensor(item["label"]).long()
return x, y
练习 2:缺 __len__ 方法。DataLoader 需要它来计算 epoch 长度和 batch 数量。补上:
def __len__(self):
return len(self.data)
练习 3:给 Dataset 加一个参数区分训练/测试,或者分别实例化两个 Dataset:
train_dataset = CSVDataset("train.csv")
test_dataset = CSVDataset("test.csv")
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)
测试集 shuffle=False 是惯例,方便对照预测和真实标签。
__len__ 和 __getitem__ 两个方法__init__ 存路径或索引,__getitem__ 做实际的数据读取和预处理__getitem__ 里才读,别在 __init__ 全读进内存__getitem__ 返回的 x 和 y 都应是 Tensor,类型要一致下一篇进阶:DataLoader 的 batch、shuffle、num_workers 到底怎么影响训练效率。
留言讨论