在这里插入图片描述

文章目录


📖 课前导读

为什么数据加载需要专门学习?

初学者很容易写出这样的代码:

# 新手写法(反面教材)
images = []
labels = []
for img_path in os.listdir("data/train"):
    img = cv2.imread(img_path)
    images.append(img)  # 把所有数据一次性加载到内存
    labels.append(label)
# 训练时每次手动取batch
for i in range(0, len(images), batch_size):
    batch = images[i:i+batch_size]
    # 训练...

这种做法有三个致命问题:

  1. 内存爆炸:ImageNet数据量150GB,一次性加载不现实
  2. 效率低下:没有利用多线程预加载,GPU在等待CPU读数据
  3. 缺乏灵活性:数据打乱、增强、批处理都要自己实现

PyTorch的DatasetDataLoader就是解决这些问题的标准方案。它们提供了:

  • 惰性加载(按需读取,不占满内存)
  • 自动批处理与打乱
  • 多线程并行加载
  • 无缝集成数据增强

💡 类比Dataset就像超市的仓库(存放原始数据),DataLoader就像自动补货机器人(按需批量取货、打乱顺序、多线程搬运)。

学完这一课,你将能够:

  • ✅ 使用torchvision.datasets调用MNIST、CIFAR-10等经典数据集
  • ✅ 构建自定义Dataset加载自己的图片、CSV、文本数据
  • ✅ 配置DataLoader的批量大小、打乱、多线程等参数
  • ✅ 实现数据归一化(Normalization)和标准化(Standardization)
  • ✅ 使用torchvision.transforms进行基础图像增强
  • ✅ 正确划分训练集、验证集、测试集

一、知识原理:数据加载的核心组件

1.1 Dataset:数据集的抽象

Dataset是PyTortch中表示数据集的抽象类。你需要继承它并实现两个方法:

  • __len__():返回数据集的总样本数
  • __getitem__(idx):根据索引idx返回一个样本(通常是(数据, 标签)对)

内部机制:DataLoader会在内部调用dataset[idx]来获取第idx个样本,然后自动组装成batch。

1.2 DataLoader:数据加载器

DataLoader包装了一个Dataset,提供:

  • 批处理:将多个样本打包成一个batch(通常形状为[batch_size, ...]
  • 打乱:每个epoch重新洗牌,避免顺序学习导致的偏差
  • 多线程加载:使用多个子进程并行读取数据,隐藏I/O延迟
  • 自动内存管理:支持pin_memory加速CPU→GPU传输

1.3 数据预处理流水线:Transforms

原始数据(图片、文本)通常需要转换为张量,并进行归一化、尺寸调整等操作。torchvision.transforms提供了一系列可组合的预处理函数。

transform = transforms.Compose([
    transforms.Resize((224, 224)),   # 调整大小
    transforms.RandomHorizontalFlip(), # 随机水平翻转(增强)
    transforms.ToTensor(),            # 转为Tensor (C,H,W),值缩放到[0,1]
    transforms.Normalize(mean=[0.485, 0.456, 0.406],  # 标准化
                         std=[0.229, 0.224, 0.225])
])

二、环境搭建与准备

import torch
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import Dataset, DataLoader, random_split
import numpy as np
import matplotlib.pyplot as plt
from PIL import Image
import os
import pandas as pd

print(f"PyTorch版本: {torch.__version__}")
print(f"torchvision版本: {torchvision.__version__}")

# 设置随机种子
torch.manual_seed(42)
np.random.seed(42)

# 创建数据目录
os.makedirs("./data", exist_ok=True)

三、代码实战:内置数据集调用

3.1 MNIST手写数字数据集

MNIST是深度学习界的“Hello World”,包含6万张28x28的手写数字灰度图。

# 定义预处理:转为Tensor,并归一化到[-1, 1](可选)
transform_mnist = transforms.Compose([
    transforms.ToTensor(),  # 将PIL图像或numpy数组转为Tensor,自动除以255
    transforms.Normalize((0.1307,), (0.3081,))  # MNIST的均值和标准差
])

# 下载并加载训练集
train_set = torchvision.datasets.MNIST(
    root='./data',          # 保存路径
    train=True,             # 训练集
    download=True,          # 如果本地没有则下载
    transform=transform_mnist
)

# 测试集
test_set = torchvision.datasets.MNIST(
    root='./data',
    train=False,
    download=True,
    transform=transform_mnist
)

print(f"训练集大小: {len(train_set)}")
print(f"测试集大小: {len(test_set)}")

# 查看一个样本
image, label = train_set[0]
print(f"图片形状: {image.shape}")  # torch.Size([1, 28, 28]),单通道
print(f"标签: {label}")

# 可视化
plt.imshow(image.squeeze(), cmap='gray')
plt.title(f"Label: {label}")
plt.show()

3.2 CIFAR-10彩色图像数据集

CIFAR-10包含10类32x32的彩色图片。

transform_cifar = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.4914, 0.4822, 0.4465],  # 每通道均值
                         std=[0.2023, 0.1994, 0.2010])   # 每通道标准差
])

train_cifar = torchvision.datasets.CIFAR10(
    root='./data', train=True, download=True,
    transform=transform_cifar
)

test_cifar = torchvision.datasets.CIFAR10(
    root='./data', train=False, download=True,
    transform=transform_cifar
)

# 查看样本
image, label = train_cifar[0]
print(f"CIFAR图像形状: {image.shape}")  # torch.Size([3, 32, 32])
classes = ['airplane', 'automobile', 'bird', 'cat', 'deer',
           'dog', 'frog', 'horse', 'ship', 'truck']
print(f"类别: {classes[label]}")

# 反归一化并显示(辅助函数)
def imshow(img):
    img = img * torch.tensor(transform_cifar.transforms[-1].std).view(3,1,1) + \
          torch.tensor(transform_cifar.transforms[-1].mean).view(3,1,1)
    img = img.permute(1, 2, 0).clamp(0, 1).numpy()
    plt.imshow(img)
    plt.axis('off')

plt.figure()
imshow(image)
plt.title(classes[label])
plt.show()

3.3 Fashion-MNIST:更难的服装分类

fashion_mnist = torchvision.datasets.FashionMNIST(
    root='./data', train=True, download=True,
    transform=transforms.ToTensor()
)
# 类别:T-shirt/top, Trouser, Pullover, Dress, Coat, Sandal, Shirt, Sneaker, Bag, Ankle boot
print(f"Fashion-MNIST类别数: {len(fashion_mnist.classes)}")

四、DataLoader核心用法

4.1 基本配置

from torch.utils.data import DataLoader

# 创建DataLoader
train_loader = DataLoader(
    dataset=train_set,      # Dataset实例
    batch_size=64,          # 每个batch的样本数
    shuffle=True,           # 每个epoch重新打乱
    num_workers=2,          # 使用2个子进程加载数据
    pin_memory=True,        # 如果使用GPU,建议开启加速传输
    drop_last=False         # 是否丢弃最后不足batch_size的batch
)

test_loader = DataLoader(
    dataset=test_set,
    batch_size=64,
    shuffle=False,          # 测试集不需要打乱
    num_workers=2
)

print(f"训练集batch数: {len(train_loader)}")  # 60000/64 ≈ 938
print(f"测试集batch数: {len(test_loader)}")   # 10000/64 = 157

4.2 迭代DataLoader

# 取出一个batch
for batch_idx, (data, targets) in enumerate(train_loader):
    print(f"Batch {batch_idx}: data shape={data.shape}, targets shape={targets.shape}")
    # data: [64, 1, 28, 28]
    # targets: [64]
    if batch_idx == 0:
        break

# 训练循环的标准写法
def train_one_epoch(model, train_loader, optimizer, criterion, device):
    model.train()
    total_loss = 0
    for batch_idx, (data, targets) in enumerate(train_loader):
        data, targets = data.to(device), targets.to(device)
        
        # 前向传播
        outputs = model(data)
        loss = criterion(outputs, targets)
        
        # 反向传播
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item()
    
    return total_loss / len(train_loader)

4.3 num_workers与加载速度

import time

def test_loader_speed(num_workers):
    loader = DataLoader(train_set, batch_size=64, shuffle=False, 
                        num_workers=num_workers, pin_memory=False)
    start = time.time()
    total = 0
    for _ in loader:
        total += 1
    elapsed = time.time() - start
    print(f"num_workers={num_workers}: {elapsed:.2f}秒, 速度={total/elapsed:.1f} batch/s")

# 测试不同worker数(注意:Windows下num_workers>0需要if __name__=='__main__'保护)
# test_loader_speed(0)
# test_loader_speed(2)
# test_loader_speed(4)

# 通常建议:CPU核心数,但过大会导致内存开销

⚠️ Windows注意事项:在Windows上使用num_workers>0时,必须将训练代码放在if __name__ == '__main__':块中,否则会报错。


五、自定义数据集

5.1 图片分类数据集(按文件夹结构)

假设数据目录结构如下:

data/
    train/
        cat/
            cat001.jpg
            cat002.jpg
        dog/
            dog001.jpg
            dog002.jpg
    test/
        cat/
        dog/
方法1:使用ImageFolder(最简单)
# ImageFolder自动根据子文件夹名生成标签
from torchvision.datasets import ImageFolder

train_dataset = ImageFolder(
    root='./data/train',
    transform=transforms.Compose([
        transforms.Resize((224, 224)),
        transforms.RandomHorizontalFlip(),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406],
                             std=[0.229, 0.224, 0.225])
    ])
)

print(f"类别映射: {train_dataset.class_to_idx}")  # {'cat':0, 'dog':1}
方法2:自定义Dataset(完全控制)

当数据不是简单的文件夹结构时(如CSV中存储图片路径和标签),需要自定义。

class CustomImageDataset(Dataset):
    """自定义图片数据集"""
    def __init__(self, annotations_file, img_dir, transform=None):
        """
        参数:
            annotations_file: CSV文件路径,包含两列 'image_name', 'label'
            img_dir: 图片所在目录
            transform: 可选的数据预处理
        """
        self.img_labels = pd.read_csv(annotations_file)
        self.img_dir = img_dir
        self.transform = transform
    
    def __len__(self):
        return len(self.img_labels)
    
    def __getitem__(self, idx):
        # 获取图片路径和标签
        img_path = os.path.join(self.img_dir, self.img_labels.iloc[idx, 0])
        label = self.img_labels.iloc[idx, 1]
        
        # 读取图片(使用PIL)
        image = Image.open(img_path).convert('RGB')  # 统一转为RGB
        
        # 应用预处理
        if self.transform:
            image = self.transform(image)
        
        return image, label

# 示例:创建一个假的CSV(实际使用时换成真实数据)
# df = pd.DataFrame({'image_name': ['img1.jpg', 'img2.jpg'], 'label': [0, 1]})
# df.to_csv('data/annotations.csv', index=False)
# dataset = CustomImageDataset('data/annotations.csv', 'data/images', transform=...)

5.2 CSV表格数据(结构化数据)

对于回归或表格分类任务(如房价预测),数据存储在CSV中。

class CSVDataset(Dataset):
    """加载CSV文件中的表格数据"""
    def __init__(self, csv_file, target_column, transform=None):
        self.data = pd.read_csv(csv_file)
        self.target_column = target_column
        self.transform = transform
        
        # 分离特征和标签
        self.features = self.data.drop(columns=[target_column]).values.astype(np.float32)
        self.labels = self.data[target_column].values.astype(np.float32)
        
        # 可选:保存特征的均值和标准差用于标准化
        self.feature_mean = self.features.mean(axis=0)
        self.feature_std = self.features.std(axis=0)
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        x = self.features[idx]
        y = self.labels[idx]
        
        # 标准化(可选)
        if self.transform == 'standardize':
            x = (x - self.feature_mean) / (self.feature_std + 1e-8)
        
        # 转为Tensor
        x = torch.tensor(x, dtype=torch.float32)
        y = torch.tensor(y, dtype=torch.float32).view(-1)  # 确保形状一致
        
        return x, y

# 使用示例(生成假数据)
# df = pd.DataFrame({'feat1': np.random.randn(1000), 'feat2': np.random.randn(1000), 'target': np.random.randn(1000)})
# df.to_csv('data/regression.csv', index=False)
# dataset = CSVDataset('data/regression.csv', 'target', transform='standardize')

5.3 文本数据(单词级)

class TextDataset(Dataset):
    """简单的文本分类数据集"""
    def __init__(self, texts, labels, vocab=None, max_length=100):
        self.texts = texts
        self.labels = labels
        self.max_length = max_length
        
        # 构建词汇表(如果未提供)
        if vocab is None:
            vocab = {'<PAD>': 0, '<UNK>': 1}
            for text in texts:
                for word in text.split():
                    if word not in vocab:
                        vocab[word] = len(vocab)
        self.vocab = vocab
    
    def text_to_indices(self, text):
        indices = [self.vocab.get(word, self.vocab['<UNK>']) for word in text.split()]
        # 填充或截断至固定长度
        if len(indices) < self.max_length:
            indices += [self.vocab['<PAD>']] * (self.max_length - len(indices))
        else:
            indices = indices[:self.max_length]
        return torch.tensor(indices, dtype=torch.long)
    
    def __len__(self):
        return len(self.texts)
    
    def __getitem__(self, idx):
        x = self.text_to_indices(self.texts[idx])
        y = torch.tensor(self.labels[idx], dtype=torch.long)
        return x, y

# 使用示例
# texts = ["hello world", "deep learning is fun", "pytorch rocks"]
# labels = [0, 1, 0]
# dataset = TextDataset(texts, labels)
# loader = DataLoader(dataset, batch_size=2)

六、数据预处理与增强

6.1 归一化 vs 标准化

方法公式作用
归一化(x - min) / (max - min)缩放到[0,1]区间
标准化(x - mean) / std转换为均值为0,方差为1

在深度学习中,标准化更常用,因为它保留了数据的分布形状,且有利于梯度下降收敛。

# 计算数据集的均值和标准差(针对图像)
def compute_mean_std(dataset):
    """遍历数据集,计算每个通道的均值和标准差"""
    loader = DataLoader(dataset, batch_size=64, shuffle=False, num_workers=2)
    mean = 0.
    std = 0.
    total = 0
    for images, _ in loader:
        # images: [B, C, H, W]
        batch_samples = images.size(0)
        images = images.view(batch_samples, images.size(1), -1)  # [B, C, H*W]
        mean += images.mean(2).sum(0)
        std += images.std(2).sum(0)
        total += batch_samples
    mean /= total
    std /= total
    return mean, std

# 示例:计算CIFAR-10的均值标准差(官方已经提供,但自己验证一下)
temp_dataset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, 
                                             transform=transforms.ToTensor())
# mean, std = compute_mean_std(temp_dataset)
# print(f"计算得到的均值: {mean}, 标准差: {std}")
# 应该接近 (0.4914, 0.4822, 0.4465) 和 (0.2023, 0.1994, 0.2010)

6.2 torchvision.transforms 常用图像增强

from torchvision import transforms

# 训练集增强(更丰富)
train_transforms = transforms.Compose([
    transforms.RandomResizedCrop(224),        # 随机裁剪并缩放到224
    transforms.RandomHorizontalFlip(p=0.5),   # 随机水平翻转
    transforms.RandomRotation(15),            # 随机旋转±15度
    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),  # 色彩抖动
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# 验证/测试集(仅调整大小和标准化)
val_transforms = 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])
])

增强效果可视化

def visualize_augmentations(dataset, num_samples=5):
    """显示数据增强的效果"""
    fig, axes = plt.subplots(num_samples, 4, figsize=(12, num_samples*3))
    for i in range(num_samples):
        img, label = dataset[i]
        # 原图(反归一化)
        img_orig = img * torch.tensor([0.229, 0.224, 0.225]).view(3,1,1) + torch.tensor([0.485, 0.456, 0.406]).view(3,1,1)
        axes[i, 0].imshow(img_orig.permute(1,2,0).clamp(0,1).numpy())
        axes[i, 0].set_title("Original")
        axes[i, 0].axis('off')
        
        # 再应用几次增强显示不同效果
        for j, aug in enumerate([transforms.RandomHorizontalFlip(p=1), 
                                 transforms.RandomRotation(30),
                                 transforms.ColorJitter(brightness=0.5)]):
            aug_img = aug(img_orig)
            axes[i, j+1].imshow(aug_img.permute(1,2,0).clamp(0,1).numpy())
            axes[i, j+1].set_title(["Horizontal Flip", "Rotation", "Color Jitter"][j])
            axes[i, j+1].axis('off')
    plt.tight_layout()
    plt.show()

# 使用一个不带归一化的数据集来可视化
# demo_dataset = torchvision.datasets.CIFAR10(root='./data', train=True, transform=transforms.ToTensor())
# visualize_augmentations(demo_dataset)

6.3 自定义Transform

class AddGaussianNoise(object):
    """添加高斯噪声"""
    def __init__(self, mean=0., std=0.05):
        self.mean = mean
        self.std = std
    
    def __call__(self, tensor):
        # tensor 是 [C,H,W] 范围的Tensor
        noise = torch.randn(tensor.size()) * self.std + self.mean
        return tensor + noise

# 使用
custom_transforms = transforms.Compose([
    transforms.ToTensor(),
    AddGaussianNoise(std=0.02),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

七、训练/验证/测试集划分实战

7.1 使用random_split划分

from torch.utils.data import random_split

# 假设我们有一个完整的数据集
full_dataset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True,
                                            transform=train_transforms)

# 划分比例: 80%训练, 10%验证, 10%测试(但原始测试集已经单独存在)
train_size = int(0.8 * len(full_dataset))
val_size = len(full_dataset) - train_size
train_subset, val_subset = random_split(full_dataset, [train_size, val_size])

# 注意:random_split返回的是Subset对象,不是Dataset,但可以正常用于DataLoader
print(f"训练子集大小: {len(train_subset)}")
print(f"验证子集大小: {len(val_subset)}")

train_loader = DataLoader(train_subset, batch_size=64, shuffle=True, num_workers=2)
val_loader = DataLoader(val_subset, batch_size=64, shuffle=False, num_workers=2)

7.2 划分并保持类别分布(分层采样)

使用StratifiedShuffleSplit需要借助sklearn,但可以手动实现:

from collections import defaultdict

def stratified_split(dataset, val_ratio=0.1):
    """按标签分层划分数据集"""
    # 获取所有标签
    labels = [dataset[i][1] for i in range(len(dataset))]
    # 按标签索引分组
    indices_by_label = defaultdict(list)
    for idx, label in enumerate(labels):
        indices_by_label[label].append(idx)
    
    train_indices = []
    val_indices = []
    for label, indices in indices_by_label.items():
        n_val = int(len(indices) * val_ratio)
        # 随机打乱后取前n_val作为验证集
        np.random.shuffle(indices)
        val_indices.extend(indices[:n_val])
        train_indices.extend(indices[n_val:])
    
    # 再整体打乱一下
    np.random.shuffle(train_indices)
    np.random.shuffle(val_indices)
    
    # 使用Subset
    train_subset = torch.utils.data.Subset(dataset, train_indices)
    val_subset = torch.utils.data.Subset(dataset, val_indices)
    return train_subset, val_subset

# 使用
# train_subset, val_subset = stratified_split(full_dataset, val_ratio=0.1)

7.3 完整的数据管道示例(端到端)

# 1. 定义变换
train_transform = transforms.Compose([
    transforms.Resize((128, 128)),
    transforms.RandomHorizontalFlip(),
    transforms.RandomRotation(10),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

val_transform = transforms.Compose([
    transforms.Resize((128, 128)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# 2. 加载数据集(假设使用ImageFolder)
full_train_dataset = ImageFolder(root='./data/train', transform=train_transform)
test_dataset = ImageFolder(root='./data/test', transform=val_transform)

# 3. 划分训练集和验证集(从训练集中分出10%)
train_size = int(0.9 * len(full_train_dataset))
val_size = len(full_train_dataset) - train_size
train_dataset, val_dataset = random_split(full_train_dataset, [train_size, val_size])

# 4. 创建DataLoader
batch_size = 32
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4, pin_memory=True)
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=4, pin_memory=True)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=4, pin_memory=True)

# 5. 打印信息
print(f"训练集batch数: {len(train_loader)}")
print(f"验证集batch数: {len(val_loader)}")
print(f"测试集batch数: {len(test_loader)}")

八、难点解析:常见问题与优化

8.1 RuntimeError: DataLoader worker (pid xxx) is killed by signal: Bus error(Mac)

原因:macOS上num_workers设置过高导致内存问题。

解决:设置num_workers=0,或降低到1

8.2 训练时数据加载成为瓶颈

症状:GPU利用率很低(nvidia-smi显示Volatile GPU-Util < 50%),但CPU占用高。

解决方案

  • 增加num_workers(不超过CPU核心数)
  • 开启pin_memory=True(GPU训练时)
  • 将数据预处理(如解码JPEG)移到GPU上(使用torchvisiondecode
  • 使用更快的存储(SSD)
  • 增加batch_size减少迭代次数

8.3 数据不平衡问题

当类别样本数量差异大时,需要采用加权采样。

from sklearn.utils.class_weight import compute_class_weight
from torch.utils.data import WeightedRandomSampler

# 计算类别权重
labels = [dataset[i][1] for i in range(len(dataset))]
class_weights = compute_class_weight('balanced', classes=np.unique(labels), y=labels)
class_weights = torch.tensor(class_weights, dtype=torch.float32)

# 为每个样本分配权重
sample_weights = [class_weights[label] for label in labels]
sampler = WeightedRandomSampler(sample_weights, num_samples=len(sample_weights), replacement=True)

# 使用sampler替代shuffle
train_loader = DataLoader(dataset, batch_size=64, sampler=sampler, num_workers=2)

8.4 大图像数据集的懒加载

对于很大的图像,推荐在__getitem__中即时读取,而不是预加载所有数据到列表。上面自定义Dataset的做法就是懒加载。

8.5 内存泄漏

如果训练中内存持续增长,可能是num_workers太高且数据集返回了大数据对象。尝试降低num_workers或在__getitem__中确保释放资源(如关闭文件句柄)。

8.6 确保数据增强只在训练集上使用

这是常见错误:验证集也应用了随机增强,导致评估指标波动。

# 正确做法:分别定义变换
train_dataset = ImageFolder(root='train', transform=train_transform)
val_dataset = ImageFolder(root='val', transform=val_transform)   # 不含随机增强

九、课后总结

核心API速查表

API作用关键参数
Dataset数据集的抽象基类__len__, __getitem__
DataLoader批量加载器batch_size, shuffle, num_workers
ImageFolder按文件夹结构的图片数据集root, transform
random_split随机划分数据集dataset, lengths
transforms.Compose组合多个预处理transforms列表
transforms.ToTensor()图片→Tensor,缩放到[0,1]
transforms.Normalize标准化mean, std
WeightedRandomSampler不平衡数据采样weights, num_samples

检查清单

  • 能调用内置数据集(MNIST、CIFAR-10)
  • 会自定义Dataset加载自己的数据
  • 能配置DataLoader的batch、shuffle、num_workers
  • 理解训练集和验证集的变换区别
  • 能完成训练/验证/测试集划分
  • 会使用transforms进行基础图像增强
  • 知道如何排查数据加载速度瓶颈

十、课后作业

作业1:内置数据集探索

使用torchvision.datasets.FashionMNIST,分别用transforms.ToTensor()transforms.Normalize预处理,可视化前5个样本并打印标签。

作业2:自定义图片数据集

从网上下载10张猫和10张狗的图片,按照文件夹结构组织,使用ImageFolder加载。如果没有图片,可以用torchvision.datasets.ImageFolder加载一个虚拟的测试目录(创建一个临时目录结构)。

作业3:DataLoader性能对比

对比num_workers=0, 2, 4时的数据加载速度(迭代完整数据集耗时),记录结果并分析瓶颈。

作业4:数据增强效果验证

在CIFAR-10上,分别用“仅ToTensor()”和“包含RandomHorizontalFlip+RandomRotation”两种增强方式,训练一个简单的CNN(可以用后面课程的网络,或用torchvision.models),对比验证集准确率,观察增强带来的提升。

作业5:完整数据管道搭建

选择一个你喜欢的数据集(如通过sklearn.datasets.make_classification生成合成数据,或从Kaggle下载一个小型CSV),自己实现CustomDataset,并完成训练/验证/测试划分,创建DataLoader。不需要训练模型,只要能够正确迭代并打印形状即可。

作业6:实现一个缓存机制

在自定义Dataset中,添加一个缓存字典,将已读取的图片数据存储在内存中,避免重复I/O。注意控制缓存大小,防止内存溢出。


十一、下一课预告

第6课我们将深入神经网络基础层与nn.Module核心架构,学习:

  • nn.Module的底层原理
  • 自定义网络的标准写法
  • 网络参数初始化
  • 模型保存与加载
  • Sequential容器用法

学完第6课,你就能搭建自己的第一个神经网络了!


🔗《精讲25课|PyTorch 从入门到精通》系列课程导航

去订阅

🌟 感谢您耐心阅读到这里!
💡 如果本文对您有所启发欢迎:
👍 点赞📌 收藏 📤 分享给更多需要的伙伴。
🗣️ 期待在评论区看到您的想法, 共同进步。
🔔 关注我,持续获取更多干货内容~
🤗 我们下篇文章见~

更多推荐