PyTorch 数据加载
📦 概述
PyTorch 提供了两个核心抽象用于数据加载:
| 组件 |
职责 |
torch.utils.data.Dataset |
定义"如何取一条数据" |
torch.utils.data.DataLoader |
定义"如何批量加载、打乱、并行取数据" |
🗂️ Dataset --- 自定义数据集
标准模板
import torch
from torch.utils.data import Dataset
class MyDataset(Dataset):
def __init__(self, data, labels, transform=None):
self.data = data
self.labels = labels
self.transform = transform
def __len__(self):
"""返回数据集大小"""
return len(self.data)
def __getitem__(self, idx):
"""返回第 idx 条数据"""
x = self.data[idx]
y = self.labels[idx]
if self.transform:
x = self.transform(x)
return x, y
实际示例:CSV 数据集
import pandas as pd
from torch.utils.data import Dataset
class CSVDataset(Dataset):
def __init__(self, csv_path):
df = pd.read_csv(csv_path)
self.features = df.drop('label', axis=1).values.astype('float32')
self.labels = df['label'].values.astype('int64')
def __len__(self):
return len(self.features)
def __getitem__(self, idx):
return (torch.tensor(self.features[idx]),
torch.tensor(self.labels[idx]))
实际示例:图像文件夹
from pathlib import Path
from PIL import Image
class ImageFolderDataset(Dataset):
def __init__(self, root_dir, transform=None):
self.root = Path(root_dir)
self.classes = sorted(d.name for d in self.root.iterdir() if d.is_dir())
self.class_to_idx = {c: i for i, c in enumerate(self.classes)}
self.samples = []
for class_dir in self.root.iterdir():
if class_dir.is_dir():
for img_path in class_dir.glob('*.jpg'):
self.samples.append((str(img_path), self.class_to_idx[class_dir.name]))
self.transform = transform
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
img_path, label = self.samples[idx]
image = Image.open(img_path).convert('RGB')
if self.transform:
image = self.transform(image)
return image, label
🔄 DataLoader --- 批量加载器 ⭐
from torch.utils.data import DataLoader
dataset = MyDataset(data, labels)
loader = DataLoader(
dataset,
batch_size=32, # 每批样本数
shuffle=True, # 每个 epoch 是否打乱
num_workers=4, # 子进程数(数据加载并行)
pin_memory=True, # 固定内存(加速 CPU→GPU 传输)
drop_last=True, # 丢弃最后不完整的 batch
prefetch_factor=2, # 预取倍数(num_workers > 0 时有效)
persistent_workers=True, # 保持 worker 进程存活(避免重复创建)
collate_fn=None, # 自定义 batch 组装方式
)
for batch_X, batch_y in loader:
print(batch_X.shape) # (batch_size, ...)
print(batch_y.shape) # (batch_size,)
# 训练...
核心参数详解
| 参数 |
默认 |
说明 |
batch_size |
1 |
每批加载样本数 |
shuffle |
False |
训练设为 True,验证设为 False |
num_workers |
0 |
Windows 上建议设 0(多进程有 Bug),Linux 上设 4~8 |
pin_memory |
False |
有 GPU 时设为 True |
drop_last |
False |
训练时通常设为 True 避免 batch 大小不一致 |
collate_fn |
默认堆叠 |
变长数据时需要自定义 |
🎨 collate_fn --- 自定义批组装
场景 1: 变长序列(NLP)
def collate_fn(batch):
"""每个样本: (token_ids, label),长度可能不同"""
texts, labels = zip(*batch)
# Padding 到相同长度
max_len = max(len(t) for t in texts)
padded = torch.zeros(len(texts), max_len, dtype=torch.long)
for i, t in enumerate(texts):
padded[i, :len(t)] = torch.tensor(t)
return padded, torch.tensor(labels)
loader = DataLoader(dataset, batch_size=32, collate_fn=collate_fn)
场景 2: 目标检测(不同数量的框)
def detection_collate(batch):
"""每个样本: (image, boxes, labels),boxes 数量可能不同"""
images, boxes, labels = zip(*batch)
return (torch.stack(images, 0),
list(boxes), # 保持为 list,长度不统一
list(labels))
loader = DataLoader(dataset, batch_size=16, collate_fn=detection_collate)
🖼️ TorchVision 内置数据集
import torchvision
import torchvision.transforms as transforms
# CIFAR-10
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465),
(0.2470, 0.2435, 0.2616))
])
train_set = torchvision.datasets.CIFAR10(
root='./data', train=True, download=True, transform=transform
)
test_set = torchvision.datasets.CIFAR10(
root='./data', train=False, download=True, transform=transform
)
# ImageNet 风格的数据集
dataset = torchvision.datasets.ImageFolder(
root='path/to/images',
transform=transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
)
常用内置数据集
| 数据集 |
代码 |
用途 |
| MNIST |
datasets.MNIST() |
手写数字分类 |
| Fashion-MNIST |
datasets.FashionMNIST() |
服装分类 |
| CIFAR-10/100 |
datasets.CIFAR10() / CIFAR100() |
小图分类 |
| ImageNet |
datasets.ImageNet() |
大规模分类 |
| SVHN |
datasets.SVHN() |
街景门牌号 |
| ImageFolder |
datasets.ImageFolder() |
自定义文件夹图片 |
常用图像增强
from torchvision import transforms
# 训练时的增强
train_transform = transforms.Compose([
transforms.RandomResizedCrop(224), # 随机裁剪
transforms.RandomHorizontalFlip(p=0.5), # 随机水平翻转
transforms.RandomRotation(15), # 随机旋转 ±15°
transforms.ColorJitter(brightness=0.2, # 颜色抖动
contrast=0.2,
saturation=0.2,
hue=0.1),
transforms.ToTensor(), # 转为 Tensor (HWC→CHW, 0-255→0-1)
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]), # 标准化
])
# 验证/测试时的增强(不做数据增强)
val_transform = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
| 操作 |
代码 |
说明 |
| 转 Tensor |
transforms.ToTensor() |
HWC→CHW, 归一化到 0,1 |
| 标准化 |
transforms.Normalize(mean, std) |
(x-mean)/std |
| 缩放 |
transforms.Resize((H, W)) |
缩放到固定尺寸 |
| 中心裁剪 |
transforms.CenterCrop(size) |
从中心裁 |
| 随机裁剪 |
transforms.RandomCrop(size) |
随机位置裁 |
| 随机缩放裁剪 |
transforms.RandomResizedCrop(224) |
先随机缩放再裁 |
| 随机翻转 |
transforms.RandomHorizontalFlip() |
水平翻转 |
| 随机旋转 |
transforms.RandomRotation(degrees) |
旋转 |
| 颜色抖动 |
transforms.ColorJitter(b, c, s, h) |
亮度/对比度/饱和度/色调 |
| 转 PIL |
transforms.ToPILImage() |
Tensor→PIL |
Albumentations(更快的增强库)
import albumentations as A
from albumentations.pytorch import ToTensorV2
transform = A.Compose([
A.RandomResizedCrop(224, 224),
A.HorizontalFlip(p=0.5),
A.ColorJitter(brightness=0.2, contrast=0.2, p=0.5),
A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
ToTensorV2(),
])
# 注意 Albumentations 需要解包:
# augmented = transform(image=image)
# image = augmented['image']
📊 数据集划分与采样
随机划分训练/验证集
from torch.utils.data import random_split
dataset = MyDataset(...)
train_size = int(0.8 * len(dataset))
val_size = len(dataset) - train_size
train_set, val_set = random_split(dataset, [train_size, val_size])
train_loader = DataLoader(train_set, batch_size=32, shuffle=True)
val_loader = DataLoader(val_set, batch_size=32, shuffle=False)
分布式采样
from torch.utils.data.distributed import DistributedSampler
sampler = DistributedSampler(dataset, shuffle=True)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
# 每个 epoch 开始时 shuffle epoch
for epoch in range(num_epochs):
sampler.set_epoch(epoch)
for batch in loader:
...
加权采样(处理不均衡数据)
from torch.utils.data import WeightedRandomSampler
# 每个类的样本数
class_counts = [1000, 300, 200, 50]
class_weights = 1.0 / torch.tensor(class_counts, dtype=torch.float)
sample_weights = class_weights[labels] # 每个样本的权重
sampler = WeightedRandomSampler(
weights=sample_weights,
num_samples=len(sample_weights),
replacement=True
)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
🧪 PyTorch 2.0+ torch.utils.data 新特性
IterableDataset --- 流式数据
from torch.utils.data import IterableDataset
class StreamDataset(IterableDataset):
def __init__(self, file_path):
self.file_path = file_path
def __iter__(self):
with open(self.file_path) as f:
for line in f:
# 处理每一行...
yield torch.tensor(parse(line))
ChainDataset --- 拼接多个数据集
from torch.utils.data import ChainDataset
combined = ChainDataset([dataset1, dataset2, dataset3])
loader = DataLoader(combined, batch_size=32)
📝 速查表
| 需求 |
代码 |
| 自定义数据集 |
class MyDataset(Dataset): |
| 数据集长度 |
def __len__(self): |
| 取一条数据 |
def __getitem__(self, idx): |
| 创建加载器 |
DataLoader(ds, batch_size=32, shuffle=True) |
| 多进程加载 |
DataLoader(ds, num_workers=4) |
| 锁页内存 |
DataLoader(ds, pin_memory=True) |
| 丢弃最后批次 |
DataLoader(ds, drop_last=True) |
| 随机划分 |
random_split(dataset, [n1, n2]) |
| 分布式采样 |
DistributedSampler(dataset) |
| 加权采样 |
WeightedRandomSampler(weights, n) |
| 自定义批次 |
DataLoader(ds, collate_fn=my_fn) |
\[pytorch-总览\|← 返回总览\]