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 import sys import time # 添加net目录到路径 # sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'net')) # from food_net import create_food_cnn from net import create_food_cnn from settings import settings # 设置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}") # 数据预处理 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(3)) class_total = list(0. for i in range(3)) 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(3): 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(settings.TRAIN_DATA_DIR, transform=transform_train) val_dataset = datasets.ImageFolder(settings.VAL_DATA_DIR, transform=transform_test) test_dataset = datasets.ImageFolder(settings.TEST_DATA_DIR, transform=transform_test) # 创建数据加载器 train_loader = DataLoader(train_dataset, batch_size=settings.BATCH_SIZE, shuffle=True, num_workers=settings.NUM_WORKERS) val_loader = DataLoader(val_dataset, batch_size=settings.BATCH_SIZE, shuffle=False, num_workers=settings.NUM_WORKERS) test_loader = DataLoader(test_dataset, batch_size=settings.BATCH_SIZE, shuffle=False, num_workers=settings.NUM_WORKERS) # 类别名称 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 = create_food_cnn().to(device) print(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad)}") # 定义损失函数和优化器 criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=settings.LEARNING_RATE, weight_decay=settings.WEIGHT_DECAY) scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=settings.SCHEDULER_STEP_SIZE, gamma=settings.SCHEDULER_GAMMA) # 训练模型 train_losses = [] train_accuracies = [] val_losses = [] val_accuracies = [] best_val_acc = 0.0 print("开始训练...") start_time = time.time() # 记录训练开始时间 num_epochs = settings.NUM_EPOCHS 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 best_model_path = settings.BEST_MODEL_PATH torch.save(model.state_dict(), settings.BEST_MODEL_PATH) print(f'保存最佳模型,验证准确率: {best_val_acc:.2f}%') end_time = time.time() # 记录训练结束时间 training_duration = end_time - start_time # 计算训练时长 # 将秒转换为小时、分钟和秒 hours = int(training_duration // 3600) minutes = int((training_duration % 3600) // 60) seconds = int(training_duration % 60) 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('../model/02/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') print(f'训练时长: {hours}小时 {minutes}分钟 {seconds}秒')