第5课:PyTorch|数据流与数据预处理基础【构建数据管道】

文章目录
📖 课前导读
为什么数据加载需要专门学习?
初学者很容易写出这样的代码:
# 新手写法(反面教材)
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]
# 训练...
这种做法有三个致命问题:
- 内存爆炸:ImageNet数据量150GB,一次性加载不现实
- 效率低下:没有利用多线程预加载,GPU在等待CPU读数据
- 缺乏灵活性:数据打乱、增强、批处理都要自己实现
PyTorch的Dataset和DataLoader就是解决这些问题的标准方案。它们提供了:
- 惰性加载(按需读取,不占满内存)
- 自动批处理与打乱
- 多线程并行加载
- 无缝集成数据增强
💡 类比:
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上(使用
torchvision的decode) - 使用更快的存储(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 从入门到精通》系列课程导航
🌟 感谢您耐心阅读到这里!
💡 如果本文对您有所启发欢迎:
👍 点赞📌 收藏 📤 分享给更多需要的伙伴。
🗣️ 期待在评论区看到您的想法, 共同进步。
🔔 关注我,持续获取更多干货内容~
🤗 我们下篇文章见~
更多推荐


所有评论(0)