Files
FoodClassifier/train/train_food_classifier.py

343 lines
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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由模型内部完成
# 把3232换成224224
# 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}')