diff --git a/README.md b/README.md index b7da1ed..951c245 100644 --- a/README.md +++ b/README.md @@ -25,7 +25,7 @@ pip install -r requirements.txt ```bash cd train -python food_classifier.py +python train_food_classifier.py ``` 确保您的数据集结构如下: diff --git a/settings/settings.py b/settings/settings.py new file mode 100644 index 0000000..e5427f0 --- /dev/null +++ b/settings/settings.py @@ -0,0 +1,51 @@ +""" +食物分类器配置文件 +包含训练参数、路径配置等 +""" + +import os + +# 基础路径配置 +BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + +# 数据集路径 +DATASET_DIR = os.path.join(BASE_DIR, 'dataset') +TRAIN_DATA_DIR = os.path.join(DATASET_DIR, 'train') +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', '03') +BEST_MODEL_PATH = os.path.join(MODEL_DIR, 'best_food_model.pth') +TRAINING_CURVES_PATH = os.path.join(MODEL_DIR, 'training_curves.png') + +# 训练参数 +# NUM_EPOCHS = 100 +NUM_EPOCHS = 3 +BATCH_SIZE = 32 +LEARNING_RATE = 0.001 +WEIGHT_DECAY = 1e-4 + +# 学习率调度器参数 +SCHEDULER_STEP_SIZE = 30 +SCHEDULER_GAMMA = 0.1 + +# 数据预处理参数 +IMAGE_SIZE = (32, 32) +NORMALIZE_MEAN = (0.485, 0.456, 0.406) +NORMALIZE_STD = (0.229, 0.224, 0.225) + +# 数据增强参数 +RANDOM_HORIZONTAL_FLIP_P = 0.5 +RANDOM_ROTATION_DEGREES = 10 +COLOR_JITTER_BRIGHTNESS = 0.2 +COLOR_JITTER_CONTRAST = 0.2 +COLOR_JITTER_SATURATION = 0.2 +COLOR_JITTER_HUE = 0.1 + +# 模型参数 +NUM_CLASSES = 3 + +# 其他配置 +NUM_WORKERS = 0 # Windows下建议设为0 +DEVICE = 'cuda' # 'cuda' 或 'cpu',程序会自动检测可用性 \ No newline at end of file diff --git a/train/food_classifier.py b/train/train_food_classifier.py similarity index 87% rename from train/food_classifier.py rename to train/train_food_classifier.py index 2bd6997..e736e8e 100644 --- a/train/food_classifier.py +++ b/train/train_food_classifier.py @@ -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() # 记录训练结束时间