diff --git a/train/food_classifier.py b/train/food_classifier.py index 371a181..d43f51d 100644 --- a/train/food_classifier.py +++ b/train/food_classifier.py @@ -9,6 +9,7 @@ import matplotlib import numpy as np from tqdm import tqdm import os +import time # 设置matplotlib支持中文显示 plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans'] # 指定默认字体 @@ -219,6 +220,7 @@ if __name__ == '__main__': best_model_path = '../model/02/best_food_model.pth' print("开始训练...") + start_time = time.time() # 记录训练开始时间 for epoch in range(num_epochs): print(f'\nEpoch {epoch+1}/{num_epochs}') print('-' * 50) @@ -248,6 +250,14 @@ if __name__ == '__main__': torch.save(model.state_dict(), 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}%') # 加载最佳模型进行测试 @@ -277,11 +287,12 @@ if __name__ == '__main__': plt.grid(True) plt.tight_layout() - plt.savefig('training_curves.png', dpi=300, bbox_inches='tight') + 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') \ No newline at end of file + print(f'训练曲线已保存为: training_curves.png') + print(f'训练时长: {hours}小时 {minutes}分钟 {seconds}秒')