每次训练完成后,保存训练结果。
This commit is contained in:
@@ -15,16 +15,17 @@ VAL_DATA_DIR = os.path.join(DATASET_DIR, 'val')
|
|||||||
TEST_DATA_DIR = os.path.join(DATASET_DIR, 'test')
|
TEST_DATA_DIR = os.path.join(DATASET_DIR, 'test')
|
||||||
|
|
||||||
# 模型保存路径
|
# 模型保存路径
|
||||||
MODEL_DIR = os.path.join(BASE_DIR, 'model', '04')
|
MODEL_DIR = os.path.join(BASE_DIR, 'model', '05')
|
||||||
BEST_MODEL_PATH = os.path.join(MODEL_DIR, 'best_food_model.pth')
|
BEST_MODEL_PATH = os.path.join(MODEL_DIR, 'best_food_model.pth')
|
||||||
TRAINING_CURVES_PATH = os.path.join(MODEL_DIR, 'training_curves.png')
|
TRAINING_CURVES_PATH = os.path.join(MODEL_DIR, 'training_curves.png')
|
||||||
|
TRAINING_RESULTS_PATH = os.path.join(MODEL_DIR, 'training_results.txt')
|
||||||
|
|
||||||
# INFERENCE_BEST_MODEL_PATH = os.path.join(BASE_DIR, 'model', '03','best_food_model.pth')
|
# INFERENCE_BEST_MODEL_PATH = os.path.join(BASE_DIR, 'model', '03','best_food_model.pth')
|
||||||
INFERENCE_BEST_MODEL_PATH = BEST_MODEL_PATH
|
INFERENCE_BEST_MODEL_PATH = BEST_MODEL_PATH
|
||||||
|
|
||||||
# 训练参数
|
# 训练参数
|
||||||
# NUM_EPOCHS = 100
|
# NUM_EPOCHS = 100
|
||||||
NUM_EPOCHS = 3
|
NUM_EPOCHS = 13
|
||||||
BATCH_SIZE = 32
|
BATCH_SIZE = 32
|
||||||
LEARNING_RATE = 0.001
|
LEARNING_RATE = 0.001
|
||||||
WEIGHT_DECAY = 1e-4
|
WEIGHT_DECAY = 1e-4
|
||||||
|
|||||||
@@ -255,3 +255,77 @@ if __name__ == '__main__':
|
|||||||
print(f'模型已保存为: {best_model_path}')
|
print(f'模型已保存为: {best_model_path}')
|
||||||
print(f'训练曲线已保存为: training_curves.png')
|
print(f'训练曲线已保存为: training_curves.png')
|
||||||
print(f'训练时长: {hours}小时 {minutes}分钟 {seconds}秒')
|
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(3))
|
||||||
|
class_total = list(0. for i in range(3))
|
||||||
|
|
||||||
|
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}')
|
||||||
|
|||||||
Reference in New Issue
Block a user