287 lines
9.4 KiB
Python
287 lines
9.4 KiB
Python
import torch
|
|
import torch.nn as nn
|
|
import torch.optim as optim
|
|
import torch.nn.functional as F
|
|
from torch.utils.data import DataLoader
|
|
from torchvision import datasets, transforms
|
|
import matplotlib.pyplot as plt
|
|
import matplotlib
|
|
import numpy as np
|
|
from tqdm import tqdm
|
|
import os
|
|
|
|
# 设置matplotlib支持中文显示
|
|
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans'] # 指定默认字体
|
|
plt.rcParams['axes.unicode_minus'] = False # 解决保存图像是负号'-'显示为方块的问题
|
|
|
|
# 设置设备
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
print(f"使用设备: {device}")
|
|
|
|
# 定义CNN模型(基于CIFAR10结构,输出层改为2分类)
|
|
class FoodCNN(nn.Module):
|
|
def __init__(self):
|
|
super(FoodCNN, self).__init__()
|
|
# 第一个卷积块
|
|
self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
|
|
self.conv2 = nn.Conv2d(32, 32, 3, padding=1)
|
|
self.pool1 = nn.MaxPool2d(2, 2)
|
|
self.dropout1 = nn.Dropout2d(0.25)
|
|
|
|
# 第二个卷积块
|
|
self.conv3 = nn.Conv2d(32, 64, 3, padding=1)
|
|
self.conv4 = nn.Conv2d(64, 64, 3, padding=1)
|
|
self.pool2 = nn.MaxPool2d(2, 2)
|
|
self.dropout2 = nn.Dropout2d(0.25)
|
|
|
|
# 第三个卷积块
|
|
self.conv5 = nn.Conv2d(64, 128, 3, padding=1)
|
|
self.conv6 = nn.Conv2d(128, 128, 3, padding=1)
|
|
self.pool3 = nn.MaxPool2d(2, 2)
|
|
self.dropout3 = nn.Dropout2d(0.25)
|
|
|
|
# 全连接层
|
|
self.fc1 = nn.Linear(128 * 4 * 4, 512)
|
|
self.dropout4 = nn.Dropout(0.5)
|
|
self.fc2 = nn.Linear(512, 2) # 改为2分类
|
|
|
|
def forward(self, x):
|
|
# 第一个卷积块
|
|
x = F.relu(self.conv1(x))
|
|
x = F.relu(self.conv2(x))
|
|
x = self.pool1(x)
|
|
x = self.dropout1(x)
|
|
|
|
# 第二个卷积块
|
|
x = F.relu(self.conv3(x))
|
|
x = F.relu(self.conv4(x))
|
|
x = self.pool2(x)
|
|
x = self.dropout2(x)
|
|
|
|
# 第三个卷积块
|
|
x = F.relu(self.conv5(x))
|
|
x = F.relu(self.conv6(x))
|
|
x = self.pool3(x)
|
|
x = self.dropout3(x)
|
|
|
|
# 展平
|
|
x = x.view(-1, 128 * 4 * 4)
|
|
|
|
# 全连接层
|
|
x = F.relu(self.fc1(x))
|
|
x = self.dropout4(x)
|
|
x = self.fc2(x)
|
|
|
|
return x
|
|
|
|
# 数据预处理
|
|
transform_train = transforms.Compose([
|
|
transforms.Resize((32, 32)),
|
|
transforms.RandomHorizontalFlip(p=0.5),
|
|
transforms.RandomRotation(10),
|
|
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
|
|
])
|
|
|
|
transform_test = transforms.Compose([
|
|
transforms.Resize((32, 32)),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
|
|
])
|
|
|
|
# 训练函数
|
|
def train_epoch(model, train_loader, criterion, optimizer, device):
|
|
model.train()
|
|
running_loss = 0.0
|
|
correct = 0
|
|
total = 0
|
|
|
|
train_bar = tqdm(train_loader, desc='训练中')
|
|
for batch_idx, (data, target) in enumerate(train_bar):
|
|
data, target = data.to(device), target.to(device)
|
|
|
|
optimizer.zero_grad()
|
|
output = model(data)
|
|
loss = criterion(output, target)
|
|
loss.backward()
|
|
optimizer.step()
|
|
|
|
running_loss += loss.item()
|
|
_, predicted = output.max(1)
|
|
total += target.size(0)
|
|
correct += predicted.eq(target).sum().item()
|
|
|
|
# 更新进度条
|
|
train_bar.set_postfix({
|
|
'Loss': f'{running_loss/(batch_idx+1):.4f}',
|
|
'Acc': f'{100.*correct/total:.2f}%'
|
|
})
|
|
|
|
return running_loss/len(train_loader), 100.*correct/total
|
|
|
|
# 验证函数
|
|
def validate(model, val_loader, criterion, device):
|
|
model.eval()
|
|
val_loss = 0.0
|
|
correct = 0
|
|
total = 0
|
|
|
|
with torch.no_grad():
|
|
val_bar = tqdm(val_loader, desc='验证中')
|
|
for data, target in val_bar:
|
|
data, target = data.to(device), target.to(device)
|
|
output = model(data)
|
|
val_loss += criterion(output, target).item()
|
|
|
|
_, predicted = output.max(1)
|
|
total += target.size(0)
|
|
correct += predicted.eq(target).sum().item()
|
|
|
|
val_bar.set_postfix({
|
|
'Loss': f'{val_loss/len(val_loader):.4f}',
|
|
'Acc': f'{100.*correct/total:.2f}%'
|
|
})
|
|
|
|
return val_loss/len(val_loader), 100.*correct/total
|
|
|
|
# 测试函数
|
|
def test(model, test_loader, device, class_names):
|
|
model.eval()
|
|
correct = 0
|
|
total = 0
|
|
class_correct = list(0. for i in range(2))
|
|
class_total = list(0. for i in range(2))
|
|
|
|
with torch.no_grad():
|
|
test_bar = tqdm(test_loader, desc='测试中')
|
|
for data, target in test_bar:
|
|
data, target = data.to(device), target.to(device)
|
|
output = model(data)
|
|
_, predicted = output.max(1)
|
|
total += target.size(0)
|
|
correct += predicted.eq(target).sum().item()
|
|
|
|
# 计算每个类别的准确率
|
|
c = (predicted == target).squeeze()
|
|
for i in range(target.size(0)):
|
|
label = target[i]
|
|
class_correct[label] += c[i].item()
|
|
class_total[label] += 1
|
|
|
|
test_bar.set_postfix({
|
|
'Acc': f'{100.*correct/total:.2f}%'
|
|
})
|
|
|
|
print(f'\n测试集总体准确率: {100.*correct/total:.2f}%')
|
|
for i in range(2):
|
|
if class_total[i] > 0:
|
|
print(f'{class_names[i]} 准确率: {100.*class_correct[i]/class_total[i]:.2f}%')
|
|
|
|
return 100.*correct/total
|
|
|
|
if __name__ == '__main__':
|
|
# 加载数据集
|
|
train_dataset = datasets.ImageFolder('../dataset/train', transform=transform_train)
|
|
val_dataset = datasets.ImageFolder('../dataset/val', transform=transform_test)
|
|
test_dataset = datasets.ImageFolder('../dataset/test', transform=transform_test)
|
|
|
|
# 创建数据加载器
|
|
batch_size = 32
|
|
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=0)
|
|
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
|
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
|
|
|
# 类别名称
|
|
class_names = train_dataset.classes
|
|
print(f"类别: {class_names}")
|
|
print(f"训练集大小: {len(train_dataset)}")
|
|
print(f"验证集大小: {len(val_dataset)}")
|
|
print(f"测试集大小: {len(test_dataset)}")
|
|
|
|
# 创建模型
|
|
model = FoodCNN().to(device)
|
|
print(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad)}")
|
|
|
|
# 定义损失函数和优化器(使用与CIFAR10相同的超参数)
|
|
criterion = nn.CrossEntropyLoss()
|
|
optimizer = optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4)
|
|
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
|
|
|
|
# 训练模型
|
|
num_epochs = 100
|
|
train_losses = []
|
|
train_accuracies = []
|
|
val_losses = []
|
|
val_accuracies = []
|
|
|
|
best_val_acc = 0.0
|
|
best_model_path = '../model/01/best_food_model.pth'
|
|
|
|
print("开始训练...")
|
|
for epoch in range(num_epochs):
|
|
print(f'\nEpoch {epoch+1}/{num_epochs}')
|
|
print('-' * 50)
|
|
|
|
# 训练
|
|
train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)
|
|
|
|
# 验证
|
|
val_loss, val_acc = validate(model, val_loader, criterion, device)
|
|
|
|
# 更新学习率
|
|
scheduler.step()
|
|
|
|
# 记录结果
|
|
train_losses.append(train_loss)
|
|
train_accuracies.append(train_acc)
|
|
val_losses.append(val_loss)
|
|
val_accuracies.append(val_acc)
|
|
|
|
print(f'训练损失: {train_loss:.4f}, 训练准确率: {train_acc:.2f}%')
|
|
print(f'验证损失: {val_loss:.4f}, 验证准确率: {val_acc:.2f}%')
|
|
print(f'当前学习率: {optimizer.param_groups[0]["lr"]:.6f}')
|
|
|
|
# 保存最佳模型
|
|
if val_acc > best_val_acc:
|
|
best_val_acc = val_acc
|
|
torch.save(model.state_dict(), best_model_path)
|
|
print(f'保存最佳模型,验证准确率: {best_val_acc:.2f}%')
|
|
|
|
print(f'\n训练完成!最佳验证准确率: {best_val_acc:.2f}%')
|
|
|
|
# 加载最佳模型进行测试
|
|
print('\n加载最佳模型进行测试...')
|
|
model.load_state_dict(torch.load(best_model_path))
|
|
test_acc = test(model, test_loader, device, class_names)
|
|
|
|
# 绘制训练曲线
|
|
plt.figure(figsize=(12, 4))
|
|
|
|
plt.subplot(1, 2, 1)
|
|
plt.plot(train_losses, label='Train Loss')
|
|
plt.plot(val_losses, label='Val Loss')
|
|
plt.title('Loss Curve')
|
|
plt.xlabel('Epoch')
|
|
plt.ylabel('Loss')
|
|
plt.legend()
|
|
plt.grid(True)
|
|
|
|
plt.subplot(1, 2, 2)
|
|
plt.plot(train_accuracies, label='Train Accuracy')
|
|
plt.plot(val_accuracies, label='Val Accuracy')
|
|
plt.title('Accuracy Curve')
|
|
plt.xlabel('Epoch')
|
|
plt.ylabel('Accuracy (%)')
|
|
plt.legend()
|
|
plt.grid(True)
|
|
|
|
plt.tight_layout()
|
|
plt.savefig('training_curves.png', dpi=300, bbox_inches='tight')
|
|
plt.show()
|
|
|
|
print(f'\n最终结果:')
|
|
print(f'最佳验证准确率: {best_val_acc:.2f}%')
|
|
print(f'测试准确率: {test_acc:.2f}%')
|
|
print(f'模型已保存为: {best_model_path}')
|
|
print(f'训练曲线已保存为: training_curves.png') |