🎯 学习目标

  • 掌握图像分类任务的完整流程
  • 学会数据预处理和数据增强技术
  • 实现完整的训练和评估流程
  • 掌握模型调优和结果可视化方法
图像分类流程

图像分类流程

构建一个完整的图像分类系统需要经历数据准备、模型构建、训练优化、评估测试四个阶段。 本节将基于CIFAR-10数据集,使用ResNet架构实现一个端到端的图像分类项目。

📊 项目流程概览

1
数据准备
2
模型构建
3
训练优化
4
评估部署

📁 数据准备

import torch import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader # 训练数据增强 train_transform = transforms.Compose([ transforms.RandomCrop(32, padding=4), # 随机裁剪 transforms.RandomHorizontalFlip(), # 随机水平翻转 transforms.RandomRotation(15), # 随机旋转 transforms.ColorJitter( # 颜色抖动 brightness=0.2, contrast=0.2, saturation=0.2 ), transforms.ToTensor(), transforms.Normalize( mean=(0.4914, 0.4822, 0.4465), std=(0.2023, 0.1994, 0.2010) ) ]) # 测试数据预处理(不增强) test_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize( mean=(0.4914, 0.4822, 0.4465), std=(0.2023, 0.1994, 0.2010) ) ]) # 加载CIFAR-10数据集 train_dataset = torchvision.datasets.CIFAR10( root='./data', train=True, download=True, transform=train_transform ) test_dataset = torchvision.datasets.CIFAR10( root='./data', train=False, download=True, transform=test_transform ) # 创建DataLoader train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=4) test_loader = DataLoader(test_dataset, batch_size=100, shuffle=False, num_workers=4)

🏗️ 模型构建

import torch.nn as nn import torchvision.models as models # 使用预训练ResNet-18 model = models.resnet18(pretrained=False) # 修改第一层适应CIFAR-10的32x32输入 model.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False) # 修改最后一层为10分类 model.fc = nn.Linear(model.fc.in_features, 10) # 移除maxpool层(CIFAR图像太小) model.maxpool = nn.Identity() # 将模型移到GPU device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = model.to(device) print(model)

🔧 训练配置

import torch.optim as optim # 损失函数 criterion = nn.CrossEntropyLoss() # 优化器:带动量的SGD optimizer = optim.SGD( model.parameters(), lr=0.1, momentum=0.9, weight_decay=5e-4 # L2正则化 ) # 学习率调度器 scheduler = optim.lr_scheduler.MultiStepLR( optimizer, milestones=[60, 120, 160], # 在这些epoch降低学习率 gamma=0.2 # 学习率衰减因子 ) # 训练参数 num_epochs = 200 best_acc = 0.0

🔄 训练与评估函数

def train(epoch): model.train() train_loss = 0 correct = 0 total = 0 for batch_idx, (inputs, targets) in enumerate(train_loader): inputs, targets = inputs.to(device), targets.to(device) optimizer.zero_grad() outputs = model(inputs) loss = criterion(outputs, targets) loss.backward() optimizer.step() train_loss += loss.item() _, predicted = outputs.max(1) total += targets.size(0) correct += predicted.eq(targets).sum().item() print(f'Epoch: {epoch} | Train Loss: {train_loss/(batch_idx+1):.3f} | ' f'Acc: {100.*correct/total:.3f}%') def test(epoch): global best_acc model.eval() test_loss = 0 correct = 0 total = 0 with torch.no_grad(): for inputs, targets in test_loader: inputs, targets = inputs.to(device), targets.to(device) outputs = model(inputs) loss = criterion(outputs, targets) test_loss += loss.item() _, predicted = outputs.max(1) total += targets.size(0) correct += predicted.eq(targets).sum().item() acc = 100. * correct / total print(f'Test Loss: {test_loss/(len(test_loader)):.3f} | Acc: {acc:.3f}%') # 保存最佳模型 if acc > best_acc: best_acc = acc torch.save(model.state_dict(), 'best_model.pth') # 主训练循环 for epoch in range(num_epochs): train(epoch) test(epoch) scheduler.step()

📊 结果可视化

import matplotlib.pyplot as plt import numpy as np # CIFAR-10类别名称 classes = ('plane', 'car', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck') def imshow(img): img = img / 2 + 0.5 # 反归一化 npimg = img.numpy() plt.imshow(np.transpose(npimg, (1, 2, 0))) plt.show() # 显示预测结果 dataiter = iter(test_loader) images, labels = next(dataiter) outputs = model(images.to(device)) _, predicted = torch.max(outputs, 1) # 打印预测结果 print('GroundTruth: ', ' '.join(f'{classes[labels[j]]}' for j in range(4))) print('Predicted: ', ' '.join(f'{classes[predicted[j]]}' for j in range(4)))

⚖️ 训练技巧总结

技巧 作用 推荐设置
数据增强 增加数据多样性,防止过拟合 RandomCrop, RandomFlip
学习率衰减 精细调优,提高最终精度 MultiStepLR或CosineAnnealing
权重衰减 L2正则化,防止过拟合 1e-4 ~ 5e-4
动量SGD 加速收敛,跳出局部最优 momentum=0.9
Label Smoothing 防止过度自信 smoothing=0.1
⚠️
常见问题

训练不收敛时检查:1) 学习率是否过大;2) 数据预处理是否正确;3) 模型结构是否匹配输入尺寸;4) 标签是否正确。

图像分类应用
图:图像分类是计算机视觉的基础任务

📝 本节小结

  • • 数据增强是提升模型泛化能力的关键
  • • 使用预训练模型可加速收敛
  • • 学习率调度对最终精度影响显著
  • • 权重衰减防止过拟合,提升泛化
  • • 训练完成后需在测试集上验证最终效果