已经把重要的配置全部拎出来了,每次训练只需要修改配置文件就可以了。
This commit is contained in:
@@ -25,7 +25,7 @@ pip install -r requirements.txt
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
cd train
|
cd train
|
||||||
python food_classifier.py
|
python train_food_classifier.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',程序会自动检测可用性
|
||||||
@@ -16,6 +16,7 @@ import time
|
|||||||
# sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'net'))
|
# sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'net'))
|
||||||
# from food_net import create_food_cnn
|
# from food_net import create_food_cnn
|
||||||
from net import create_food_cnn
|
from net import create_food_cnn
|
||||||
|
from settings import settings
|
||||||
|
|
||||||
# 设置matplotlib支持中文显示
|
# 设置matplotlib支持中文显示
|
||||||
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans'] # 指定默认字体
|
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__':
|
if __name__ == '__main__':
|
||||||
# 加载数据集
|
# 加载数据集
|
||||||
train_dataset = datasets.ImageFolder('../dataset/train', transform=transform_train)
|
train_dataset = datasets.ImageFolder(settings.TRAIN_DATA_DIR, transform=transform_train)
|
||||||
val_dataset = datasets.ImageFolder('../dataset/val', transform=transform_test)
|
val_dataset = datasets.ImageFolder(settings.VAL_DATA_DIR, transform=transform_test)
|
||||||
test_dataset = datasets.ImageFolder('../dataset/test', transform=transform_test)
|
test_dataset = datasets.ImageFolder(settings.TEST_DATA_DIR, transform=transform_test)
|
||||||
|
|
||||||
# 创建数据加载器
|
# 创建数据加载器
|
||||||
batch_size = 32
|
train_loader = DataLoader(train_dataset, batch_size=settings.BATCH_SIZE, shuffle=True, num_workers=settings.NUM_WORKERS)
|
||||||
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=0)
|
val_loader = DataLoader(val_dataset, batch_size=settings.BATCH_SIZE, shuffle=False, num_workers=settings.NUM_WORKERS)
|
||||||
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
test_loader = DataLoader(test_dataset, batch_size=settings.BATCH_SIZE, shuffle=False, num_workers=settings.NUM_WORKERS)
|
||||||
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=0)
|
|
||||||
|
|
||||||
# 类别名称
|
# 类别名称
|
||||||
class_names = train_dataset.classes
|
class_names = train_dataset.classes
|
||||||
@@ -156,23 +156,22 @@ if __name__ == '__main__':
|
|||||||
model = create_food_cnn().to(device)
|
model = create_food_cnn().to(device)
|
||||||
print(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad)}")
|
print(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad)}")
|
||||||
|
|
||||||
# 定义损失函数和优化器(使用与CIFAR10相同的超参数)
|
# 定义损失函数和优化器
|
||||||
criterion = nn.CrossEntropyLoss()
|
criterion = nn.CrossEntropyLoss()
|
||||||
optimizer = optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4)
|
optimizer = optim.Adam(model.parameters(), lr=settings.LEARNING_RATE, weight_decay=settings.WEIGHT_DECAY)
|
||||||
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
|
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=settings.SCHEDULER_STEP_SIZE, gamma=settings.SCHEDULER_GAMMA)
|
||||||
|
|
||||||
# 训练模型
|
# 训练模型
|
||||||
num_epochs = 100
|
|
||||||
train_losses = []
|
train_losses = []
|
||||||
train_accuracies = []
|
train_accuracies = []
|
||||||
val_losses = []
|
val_losses = []
|
||||||
val_accuracies = []
|
val_accuracies = []
|
||||||
|
|
||||||
best_val_acc = 0.0
|
best_val_acc = 0.0
|
||||||
best_model_path = '../model/02/best_food_model.pth'
|
|
||||||
|
|
||||||
print("开始训练...")
|
print("开始训练...")
|
||||||
start_time = time.time() # 记录训练开始时间
|
start_time = time.time() # 记录训练开始时间
|
||||||
|
num_epochs = settings.NUM_EPOCHS
|
||||||
for epoch in range(num_epochs):
|
for epoch in range(num_epochs):
|
||||||
print(f'\nEpoch {epoch+1}/{num_epochs}')
|
print(f'\nEpoch {epoch+1}/{num_epochs}')
|
||||||
print('-' * 50)
|
print('-' * 50)
|
||||||
@@ -199,7 +198,8 @@ if __name__ == '__main__':
|
|||||||
# 保存最佳模型
|
# 保存最佳模型
|
||||||
if val_acc > best_val_acc:
|
if val_acc > best_val_acc:
|
||||||
best_val_acc = 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}%')
|
print(f'保存最佳模型,验证准确率: {best_val_acc:.2f}%')
|
||||||
|
|
||||||
end_time = time.time() # 记录训练结束时间
|
end_time = time.time() # 记录训练结束时间
|
||||||
Reference in New Issue
Block a user