From 97401e908a6ee53a66b33db20fb7391f2594291b Mon Sep 17 00:00:00 2001 From: zhanghuan <1262329256@qq.com> Date: Wed, 10 Sep 2025 15:32:21 +0800 Subject: [PATCH] =?UTF-8?q?=E6=AF=8F=E6=AC=A1=E8=AE=AD=E7=BB=83=E5=AE=8C?= =?UTF-8?q?=E6=88=90=E5=90=8E=EF=BC=8C=E4=BF=9D=E5=AD=98=E8=AE=AD=E7=BB=83?= =?UTF-8?q?=E7=BB=93=E6=9E=9C=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- settings/settings.py | 5 ++- train/train_food_classifier.py | 74 ++++++++++++++++++++++++++++++++++ 2 files changed, 77 insertions(+), 2 deletions(-) diff --git a/settings/settings.py b/settings/settings.py index 489c792..f5a49f3 100644 --- a/settings/settings.py +++ b/settings/settings.py @@ -15,16 +15,17 @@ VAL_DATA_DIR = os.path.join(DATASET_DIR, 'val') 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') 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 = BEST_MODEL_PATH # 训练参数 # NUM_EPOCHS = 100 -NUM_EPOCHS = 3 +NUM_EPOCHS = 13 BATCH_SIZE = 32 LEARNING_RATE = 0.001 WEIGHT_DECAY = 1e-4 diff --git a/train/train_food_classifier.py b/train/train_food_classifier.py index 7cdfb6b..fd9e203 100644 --- a/train/train_food_classifier.py +++ b/train/train_food_classifier.py @@ -255,3 +255,77 @@ if __name__ == '__main__': 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(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}')