📁 数据准备
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)
)
])
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
)
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
model = models.resnet18(pretrained=False)
model.conv1 = nn.Conv2d(3, 64, kernel_size=3,
stride=1, padding=1, bias=False)
model.fc = nn.Linear(model.fc.in_features, 10)
model.maxpool = nn.Identity()
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()
optimizer = optim.SGD(
model.parameters(),
lr=0.1,
momentum=0.9,
weight_decay=5e-4
)
scheduler = optim.lr_scheduler.MultiStepLR(
optimizer,
milestones=[60, 120, 160],
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
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)))