已经把重要的配置全部拎出来了,每次训练只需要修改配置文件就可以了。

This commit is contained in:
zhanghuan
2025-09-10 15:02:23 +08:00
parent 4c2fa0e533
commit 06fead1ab3
3 changed files with 65 additions and 14 deletions
@@ -16,6 +16,7 @@ import time
# sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'net'))
# from food_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'] # 指定默认字体
@@ -135,15 +136,14 @@ def test(model, test_loader, device, class_names):
if __name__ == '__main__':
# 加载数据集
train_dataset = datasets.ImageFolder('../dataset/train', transform=transform_train)
val_dataset = datasets.ImageFolder('../dataset/val', transform=transform_test)
test_dataset = datasets.ImageFolder('../dataset/test', transform=transform_test)
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)
# 创建数据加载器
batch_size = 32
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=0)
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
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
@@ -156,23 +156,22 @@ if __name__ == '__main__':
model = create_food_cnn().to(device)
print(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad)}")
# 定义损失函数和优化器(使用与CIFAR10相同的超参数)
# 定义损失函数和优化器
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4)
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
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)
# 训练模型
num_epochs = 100
train_losses = []
train_accuracies = []
val_losses = []
val_accuracies = []
best_val_acc = 0.0
best_model_path = '../model/02/best_food_model.pth'
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)
@@ -199,7 +198,8 @@ if __name__ == '__main__':
# 保存最佳模型
if val_acc > best_val_acc:
best_val_acc = val_acc
torch.save(model.state_dict(), best_model_path)
best_model_path = settings.BEST_MODEL_PATH
torch.save(model.state_dict(), settings.BEST_MODEL_PATH)
print(f'保存最佳模型,验证准确率: {best_val_acc:.2f}%')
end_time = time.time() # 记录训练结束时间