343 lines
12 KiB
Python
343 lines
12 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
|
||
import sys
|
||
import time
|
||
|
||
# 添加net目录到路径
|
||
# sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'net'))
|
||
sys.path.append(os.path.join(os.path.dirname(__file__), '..'))
|
||
from 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}")
|
||
|
||
|
||
|
||
# 数据预处理 - 必须包含Resize以保证batch中tensor尺寸一致,只有Normalize由模型内部完成
|
||
# 把32,32换成224,224
|
||
# transform_train = transforms.Compose([
|
||
# transforms.Resize((224, 224)),
|
||
# transforms.RandomHorizontalFlip(p=0.5),
|
||
# transforms.RandomRotation(10),
|
||
# transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),
|
||
# transforms.ToTensor(),
|
||
# # 注意:只有Normalize由模型内部处理
|
||
# ])
|
||
transform_train = transforms.Compose([
|
||
transforms.Resize((224, 224)),
|
||
transforms.RandomHorizontalFlip(p=0.5),
|
||
transforms.RandomRotation(10),
|
||
transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.02),
|
||
transforms.ToTensor(),
|
||
# 注意:只有Normalize由模型内部处理
|
||
])
|
||
|
||
# 把缩放成32*32,修改为224*224
|
||
transform_test = transforms.Compose([
|
||
transforms.Resize((224, 224)), # 必须保留,确保batch中tensor尺寸一致
|
||
transforms.ToTensor(),
|
||
# 注意:只有Normalize由模型内部处理
|
||
])
|
||
|
||
# 训练函数
|
||
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(settings.NUM_CLASSES))
|
||
class_total = list(0. for i in range(settings.NUM_CLASSES))
|
||
|
||
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(settings.NUM_CLASSES):
|
||
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(use_internal_preprocess=True).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
|
||
|
||
# 检查并创建模型保存目录
|
||
model_dir = os.path.dirname(best_model_path)
|
||
if not os.path.exists(model_dir):
|
||
os.makedirs(model_dir)
|
||
print(f"创建模型保存目录: {model_dir}")
|
||
|
||
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(settings.TRAINING_CURVES_PATH, 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}秒')
|
||
|
||
# 保存训练结果到文件
|
||
import datetime
|
||
os.makedirs(settings.MODEL_DIR, exist_ok=True)
|
||
|
||
with open(settings.TRAINING_RESULTS_PATH, 'w', encoding='utf-8') as f:
|
||
f.write("=" * 60 + "\n")
|
||
f.write("食物分类器训练结果报告\n")
|
||
f.write("=" * 60 + "\n\n")
|
||
|
||
# 训练基本信息
|
||
f.write(f"训练完成时间: {datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n")
|
||
f.write(f"使用设备: {device}\n")
|
||
f.write(f"训练时长: {hours}小时 {minutes}分钟 {seconds}秒\n\n")
|
||
|
||
# 数据集信息
|
||
f.write("数据集信息:\n")
|
||
f.write("-" * 30 + "\n")
|
||
f.write(f"类别: {class_names}\n")
|
||
f.write(f"训练集大小: {len(train_dataset)}\n")
|
||
f.write(f"验证集大小: {len(val_dataset)}\n")
|
||
f.write(f"测试集大小: {len(test_dataset)}\n\n")
|
||
|
||
# 训练参数
|
||
f.write("训练参数:\n")
|
||
f.write("-" * 30 + "\n")
|
||
f.write(f"训练轮数: {settings.NUM_EPOCHS}\n")
|
||
f.write(f"批次大小: {settings.BATCH_SIZE}\n")
|
||
f.write(f"学习率: {settings.LEARNING_RATE}\n")
|
||
f.write(f"权重衰减: {settings.WEIGHT_DECAY}\n")
|
||
f.write(f"学习率调度器步长: {settings.SCHEDULER_STEP_SIZE}\n")
|
||
f.write(f"学习率衰减因子: {settings.SCHEDULER_GAMMA}\n\n")
|
||
|
||
# 模型信息
|
||
f.write("模型信息:\n")
|
||
f.write("-" * 30 + "\n")
|
||
f.write(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad)}\n")
|
||
f.write(f"模型保存路径: {settings.BEST_MODEL_PATH}\n\n")
|
||
|
||
# 训练结果
|
||
f.write("训练结果:\n")
|
||
f.write("-" * 30 + "\n")
|
||
f.write(f"最佳验证准确率: {best_val_acc:.2f}%\n")
|
||
f.write(f"最终测试准确率: {test_acc:.2f}%\n\n")
|
||
|
||
# 各类别准确率详情
|
||
f.write("各类别测试准确率:\n")
|
||
f.write("-" * 30 + "\n")
|
||
model.eval()
|
||
class_correct = list(0. for i in range(settings.NUM_CLASSES))
|
||
class_total = list(0. for i in range(settings.NUM_CLASSES))
|
||
|
||
with torch.no_grad():
|
||
for data, target in test_loader:
|
||
data, target = data.to(device), target.to(device)
|
||
output = model(data)
|
||
_, predicted = output.max(1)
|
||
|
||
c = (predicted == target).squeeze()
|
||
for i in range(target.size(0)):
|
||
label = target[i]
|
||
class_correct[label] += c[i].item()
|
||
class_total[label] += 1
|
||
|
||
for i in range(3):
|
||
if class_total[i] > 0:
|
||
acc = 100. * class_correct[i] / class_total[i]
|
||
f.write(f"{class_names[i]}: {acc:.2f}% ({int(class_correct[i])}/{int(class_total[i])})\n")
|
||
|
||
f.write("\n" + "=" * 60 + "\n")
|
||
f.write("训练完成!\n")
|
||
f.write("=" * 60 + "\n")
|
||
|
||
print(f'训练结果已保存为: {settings.TRAINING_RESULTS_PATH}')
|