Python-pytorch-数据加载

PyTorch 数据加载

📦 概述

PyTorch 提供了两个核心抽象用于数据加载:

组件 职责
torch.utils.data.Dataset 定义"如何取一条数据"
torch.utils.data.DataLoader 定义"如何批量加载、打乱、并行取数据"

🗂️ Dataset --- 自定义数据集

标准模板

python 复制代码
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 数据集

python 复制代码
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]))

实际示例:图像文件夹

python 复制代码
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 --- 批量加载器 ⭐

python 复制代码
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)

python 复制代码
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: 目标检测(不同数量的框)

python 复制代码
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 内置数据集

python 复制代码
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() 自定义文件夹图片

🔧 数据增强(Transforms)

常用图像增强

python 复制代码
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]),
])

常用 Transform 列表

操作 代码 说明
转 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(更快的增强库)

python 复制代码
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']

📊 数据集划分与采样

随机划分训练/验证集

python 复制代码
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)

分布式采样

python 复制代码
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:
        ...

加权采样(处理不均衡数据)

python 复制代码
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 --- 流式数据

python 复制代码
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 --- 拼接多个数据集

python 复制代码
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-总览\|← 返回总览\]

相关推荐
小刘快学习2 小时前
印刷包装企业的工艺问答与报价,为什么从直连模型改成了聚合网关
大数据·人工智能
you鬰2 小时前
datawhale--llm-algo-leetcod1️⃣
python·datawhale
heimeiyingwang2 小时前
【架构实战】可观测性三支柱实战:Metrics、Logging、Tracing 如何统一落地
开发语言·架构·php
40岁资深老架构师尼恩2 小时前
RAGFlow 三大引擎 详解:DeepDoc、RAPTOR、GraphRAG
人工智能
水境传感 李兆栋2 小时前
林地农田生态观测,多光谱植被监测仪发挥哪些作用
人工智能
_風箏2 小时前
TRAE WORK 复杂任务放心交给ta
人工智能
MobotStone2 小时前
AI 评测:别瞎猜,用评估模型说话
人工智能
大模型探索者2 小时前
2026零售与电子商务大模型训推平台选型指南:大促高并发与极致算力降本破局
大数据·人工智能·零售
叶辞树2 小时前
使用AI分析trace
人工智能