814 lines
28 KiB
Python
814 lines
28 KiB
Python
import torch
|
|
import torch.nn as nn
|
|
import torch.optim as optim
|
|
import torch.nn.functional as F
|
|
from torch.utils.data import DataLoader, Dataset
|
|
from torchvision import datasets, transforms
|
|
import matplotlib.pyplot as plt
|
|
import matplotlib
|
|
import numpy as np
|
|
from tqdm import tqdm
|
|
import os
|
|
import sys
|
|
import time
|
|
import json
|
|
import random
|
|
from collections import defaultdict
|
|
from typing import List, Tuple, Dict, Optional
|
|
import logging
|
|
from datetime import datetime
|
|
|
|
# 添加项目根目录到路径
|
|
sys.path.append(os.path.join(os.path.dirname(__file__), '..'))
|
|
# 引入模型
|
|
from net.resnet_embedding import create_resnet50_embedding
|
|
from settings import settings
|
|
|
|
# 任务配置
|
|
from dataclasses import dataclass
|
|
|
|
@dataclass
|
|
class TaskConfig:
|
|
name: str
|
|
train_dir: str
|
|
val_dir: str
|
|
embedding_dim: int
|
|
batch_size: int
|
|
lr: float
|
|
triplet_margin: float
|
|
center_loss_weight: float
|
|
aug_strength: str # "strong" | "medium" | "shape"
|
|
|
|
# 使用你已创建的目录结构
|
|
TASKS = {
|
|
"dish": TaskConfig(
|
|
name="DishClassification",
|
|
train_dir=os.path.join(settings.BASE_DIR, "dataset", "DishClassification", "train"),
|
|
val_dir=os.path.join(settings.BASE_DIR, "dataset", "DishClassification", "val"),
|
|
embedding_dim=512,
|
|
batch_size=16,
|
|
lr=1e-3,
|
|
triplet_margin=0.3,
|
|
center_loss_weight=0.1,
|
|
aug_strength="medium",
|
|
),
|
|
"whole_ingredient": TaskConfig(
|
|
name="WholeIngredientRecognition",
|
|
train_dir=os.path.join(settings.BASE_DIR, "dataset", "WholeIngredientRecognition", "train"),
|
|
val_dir=os.path.join(settings.BASE_DIR, "dataset", "WholeIngredientRecognition", "val"),
|
|
embedding_dim=512,
|
|
batch_size=32,
|
|
lr=8e-4,
|
|
triplet_margin=0.35,
|
|
center_loss_weight=0.1,
|
|
aug_strength="medium",
|
|
),
|
|
"processed_ingredient": TaskConfig(
|
|
name="ProcessedIngredientRecognition",
|
|
train_dir=os.path.join(settings.BASE_DIR, "dataset", "ProcessedIngredientRecognition", "train"),
|
|
val_dir=os.path.join(settings.BASE_DIR, "dataset", "ProcessedIngredientRecognition", "val"),
|
|
embedding_dim=512,
|
|
batch_size=16,
|
|
lr=1e-3,
|
|
triplet_margin=0.25,
|
|
center_loss_weight=0.1,
|
|
aug_strength="shape",
|
|
),
|
|
}
|
|
|
|
def build_transforms(aug_strength: str):
|
|
if aug_strength == "strong":
|
|
return transforms.Compose([
|
|
transforms.Resize((224, 224)),
|
|
transforms.RandomHorizontalFlip(p=0.5),
|
|
transforms.RandomRotation(20),
|
|
transforms.ColorJitter(0.3,0.3,0.3,0.1),
|
|
transforms.RandomAffine(degrees=0, translate=(0.12,0.12)),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),
|
|
]), transforms.Compose([
|
|
transforms.Resize((224, 224)),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),
|
|
])
|
|
if aug_strength == "medium":
|
|
return transforms.Compose([
|
|
transforms.Resize((224, 224)),
|
|
transforms.RandomHorizontalFlip(p=0.5),
|
|
transforms.RandomRotation(15),
|
|
transforms.ColorJitter(0.2,0.2,0.2,0.1),
|
|
transforms.RandomAffine(degrees=0,translate=(0.1,0.1)),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),
|
|
]), transforms.Compose([
|
|
transforms.Resize((224, 224)),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),
|
|
])
|
|
# shape: 强调几何与尺度,弱化强色抖动
|
|
return transforms.Compose([
|
|
transforms.Resize((224, 224)),
|
|
transforms.RandomHorizontalFlip(p=0.5),
|
|
transforms.RandomAffine(degrees=15, translate=(0.1,0.1), scale=(0.9,1.1)),
|
|
transforms.RandomPerspective(distortion_scale=0.3, p=0.3),
|
|
transforms.GaussianBlur(kernel_size=3, sigma=(0.1,1.0)),
|
|
transforms.ColorJitter(0.1,0.1,0.1,0.03),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),
|
|
]), transforms.Compose([
|
|
transforms.Resize((224, 224)),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),
|
|
])
|
|
|
|
# 设置matplotlib支持中文显示
|
|
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans']
|
|
plt.rcParams['axes.unicode_minus'] = False
|
|
|
|
# 设置日志
|
|
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# 设置设备
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
print(f"使用设备: {device}")
|
|
|
|
|
|
class TripletDataset(Dataset):
|
|
"""
|
|
三元组数据集,用于三元组损失训练
|
|
每个样本包含:锚点(anchor)、正样本(positive)、负样本(negative)
|
|
"""
|
|
|
|
def __init__(self, dataset_path: str, transform=None, samples_per_class: int = 100):
|
|
"""
|
|
初始化三元组数据集
|
|
|
|
Args:
|
|
dataset_path: 数据集路径
|
|
transform: 数据变换
|
|
samples_per_class: 每个类别最多使用的样本数
|
|
"""
|
|
self.dataset_path = dataset_path
|
|
self.transform = transform
|
|
self.samples_per_class = samples_per_class
|
|
|
|
# 加载数据集
|
|
self.class_to_idx = {}
|
|
self.idx_to_class = {}
|
|
self.samples_by_class = defaultdict(list)
|
|
self.all_samples = []
|
|
|
|
self._load_dataset()
|
|
|
|
def _load_dataset(self):
|
|
"""加载数据集并按类别组织"""
|
|
class_names = sorted(os.listdir(self.dataset_path))
|
|
|
|
for idx, class_name in enumerate(class_names):
|
|
if not os.path.isdir(os.path.join(self.dataset_path, class_name)):
|
|
continue
|
|
|
|
self.class_to_idx[class_name] = idx
|
|
self.idx_to_class[idx] = class_name
|
|
|
|
class_path = os.path.join(self.dataset_path, class_name)
|
|
image_files = [f for f in os.listdir(class_path)
|
|
if f.lower().endswith(('.png', '.jpg', '.jpeg'))]
|
|
|
|
# 限制每个类别的样本数,没有必要限制
|
|
# if len(image_files) > self.samples_per_class:
|
|
# image_files = random.sample(image_files, self.samples_per_class)
|
|
|
|
for img_file in image_files:
|
|
img_path = os.path.join(class_path, img_file)
|
|
self.samples_by_class[idx].append(img_path)
|
|
self.all_samples.append((img_path, idx))
|
|
|
|
logger.info(f"加载数据集完成:")
|
|
for class_name, class_idx in self.class_to_idx.items():
|
|
logger.info(f" {class_name}: {len(self.samples_by_class[class_idx])} 张图片")
|
|
|
|
def __len__(self):
|
|
return len(self.all_samples)
|
|
|
|
def __getitem__(self, idx):
|
|
"""
|
|
获取三元组样本
|
|
|
|
Returns:
|
|
tuple: (anchor, positive, negative, anchor_label)
|
|
"""
|
|
# 获取锚点样本
|
|
anchor_path, anchor_label = self.all_samples[idx]
|
|
anchor_img = self._load_image(anchor_path)
|
|
|
|
# 获取正样本(同类别的不同图片)
|
|
# 需要从自己所在类别中先把自己给排除掉
|
|
positive_candidates = [path for path in self.samples_by_class[anchor_label]
|
|
if path != anchor_path]
|
|
if positive_candidates:
|
|
# 从这里可以看出来是随机选的
|
|
positive_path = random.choice(positive_candidates)
|
|
else:
|
|
positive_path = anchor_path # 如果只有一张图片,使用自己作为正样本
|
|
positive_img = self._load_image(positive_path)
|
|
|
|
# 获取负样本(不同类别的图片)
|
|
# 获得其它类别
|
|
negative_classes = [cls for cls in self.samples_by_class.keys() if cls != anchor_label]
|
|
# 随便选一个类别
|
|
negative_class = random.choice(negative_classes)
|
|
# 随便选一个路径
|
|
negative_path = random.choice(self.samples_by_class[negative_class])
|
|
negative_img = self._load_image(negative_path)
|
|
|
|
return anchor_img, positive_img, negative_img, anchor_label
|
|
|
|
def _load_image(self, image_path: str):
|
|
"""加载并预处理图片"""
|
|
from PIL import Image
|
|
|
|
try:
|
|
image = Image.open(image_path).convert('RGB')
|
|
if self.transform:
|
|
image = self.transform(image)
|
|
return image
|
|
except Exception as e:
|
|
logger.error(f"加载图片失败: {image_path}, 错误: {e}")
|
|
# 返回一个黑色图片作为备用
|
|
if self.transform:
|
|
return self.transform(Image.new('RGB', (224, 224), (0, 0, 0)))
|
|
else:
|
|
return Image.new('RGB', (224, 224), (0, 0, 0))
|
|
|
|
def get_class_names(self):
|
|
"""获取类别名称列表"""
|
|
return [self.idx_to_class[i] for i in sorted(self.idx_to_class.keys())]
|
|
|
|
|
|
class TripletLoss(nn.Module):
|
|
"""
|
|
三元组损失函数
|
|
margin是一个超参数,目的是同类聚,异类散
|
|
"""
|
|
|
|
def __init__(self, margin: float = 0.3):
|
|
"""
|
|
初始化三元组损失
|
|
|
|
Args:
|
|
margin: 边界值,正负样本距离的最小差值
|
|
"""
|
|
super(TripletLoss, self).__init__()
|
|
self.margin = margin
|
|
|
|
def forward(self, anchor, positive, negative):
|
|
"""
|
|
计算三元组损失
|
|
|
|
Args:
|
|
anchor: 锚点特征向量 [batch_size, embedding_dim]
|
|
positive: 正样本特征向量 [batch_size, embedding_dim]
|
|
negative: 负样本特征向量 [batch_size, embedding_dim]
|
|
|
|
Returns:
|
|
torch.Tensor: 三元组损失值
|
|
"""
|
|
# 计算距离(p = 2,代表欧式距离 p = 1,代表曼哈顿距离)
|
|
pos_dist = F.pairwise_distance(anchor, positive, p=2)
|
|
neg_dist = F.pairwise_distance(anchor, negative, p=2)
|
|
|
|
# 计算三元组损失
|
|
loss = F.relu(pos_dist - neg_dist + self.margin)
|
|
|
|
return loss.mean()
|
|
|
|
|
|
class CenterLoss(nn.Module):
|
|
"""
|
|
中心损失函数,用于增强类内聚合
|
|
"""
|
|
|
|
def __init__(self, num_classes: int, embedding_dim: int, alpha: float = 0.5):
|
|
"""
|
|
初始化中心损失
|
|
|
|
Args:
|
|
num_classes: 类别数量
|
|
embedding_dim: 特征向量维度
|
|
alpha: 中心更新的学习率
|
|
"""
|
|
super(CenterLoss, self).__init__()
|
|
self.num_classes = num_classes
|
|
self.embedding_dim = embedding_dim
|
|
self.alpha = alpha
|
|
|
|
# 初始化类别中心,维度是类别数量乘以特征向量维度
|
|
self.centers = nn.Parameter(torch.randn(num_classes, embedding_dim))
|
|
|
|
def forward(self, embeddings, labels):
|
|
"""
|
|
计算中心损失
|
|
|
|
Args:
|
|
embeddings: 特征向量 [batch_size, embedding_dim]
|
|
labels: 标签 [batch_size]
|
|
|
|
Returns:
|
|
torch.Tensor: 中心损失值
|
|
"""
|
|
batch_size = embeddings.size(0)
|
|
|
|
# 获取对应的中心向量
|
|
centers_batch = self.centers[labels] # [batch_size, embedding_dim]
|
|
|
|
# 计算特征向量与中心的距离
|
|
loss = F.mse_loss(embeddings, centers_batch)
|
|
|
|
return loss
|
|
|
|
def update_centers(self, embeddings, labels):
|
|
"""
|
|
更新类别中心
|
|
|
|
Args:
|
|
embeddings: 特征向量 [batch_size, embedding_dim]
|
|
labels: 标签 [batch_size]
|
|
"""
|
|
with torch.no_grad():
|
|
for label in torch.unique(labels):
|
|
mask = labels == label
|
|
if mask.sum() > 0:
|
|
# 计算该类别的平均特征向量
|
|
class_embeddings = embeddings[mask]
|
|
class_center = class_embeddings.mean(dim=0)
|
|
|
|
# 更新中心
|
|
self.centers[label] = (1 - self.alpha) * self.centers[label] + self.alpha * class_center
|
|
|
|
|
|
class EarlyStopping:
|
|
"""早停机制"""
|
|
|
|
def __init__(self, patience: int = 10, min_delta: float = 0.001):
|
|
"""
|
|
初始化早停机制
|
|
|
|
Args:
|
|
patience: 容忍的epoch数
|
|
min_delta: 最小改善幅度
|
|
"""
|
|
self.patience = patience
|
|
self.min_delta = min_delta
|
|
self.counter = 0
|
|
self.best_loss = float('inf')
|
|
|
|
|
|
def __call__(self, val_loss: float) -> bool:
|
|
"""
|
|
检查是否应该早停
|
|
|
|
Args:
|
|
val_loss: 验证损失
|
|
|
|
Returns:
|
|
bool: 是否应该早停
|
|
"""
|
|
if val_loss < self.best_loss - self.min_delta:
|
|
self.best_loss = val_loss
|
|
self.counter = 0
|
|
else:
|
|
self.counter += 1
|
|
|
|
return self.counter >= self.patience
|
|
|
|
|
|
def train_epoch(model, train_loader, triplet_criterion, center_criterion,
|
|
optimizer, center_optimizer, device, epoch):
|
|
"""
|
|
训练一个epoch
|
|
|
|
Args:
|
|
model: 模型
|
|
train_loader: 训练数据加载器
|
|
triplet_criterion: 三元组损失函数
|
|
center_criterion: 中心损失函数
|
|
optimizer: 模型优化器
|
|
center_optimizer: 中心损失优化器
|
|
device: 设备
|
|
epoch: 当前epoch
|
|
|
|
Returns:
|
|
tuple: (平均损失, 三元组损失, 中心损失)
|
|
"""
|
|
model.train()
|
|
total_loss = 0.0
|
|
total_triplet_loss = 0.0
|
|
total_center_loss = 0.0
|
|
num_batches = 0
|
|
|
|
|
|
|
|
train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1} 训练中')
|
|
|
|
for batch_idx, (anchor, positive, negative, labels) in enumerate(train_bar):
|
|
anchor = anchor.to(device)
|
|
positive = positive.to(device)
|
|
negative = negative.to(device)
|
|
labels = labels.to(device)
|
|
|
|
# 前向传播
|
|
anchor_emb = model(anchor)
|
|
positive_emb = model(positive)
|
|
negative_emb = model(negative)
|
|
|
|
# 计算三元组损失
|
|
triplet_loss = triplet_criterion(anchor_emb, positive_emb, negative_emb)
|
|
|
|
# 计算中心损失
|
|
center_loss = center_criterion(anchor_emb, labels)
|
|
|
|
# 总损失
|
|
loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重
|
|
|
|
# 反向传播
|
|
optimizer.zero_grad()
|
|
center_optimizer.zero_grad()
|
|
|
|
loss.backward()
|
|
|
|
optimizer.step()
|
|
center_optimizer.step()
|
|
|
|
# 更新中心
|
|
center_criterion.update_centers(anchor_emb.detach(), labels)
|
|
|
|
# 统计
|
|
total_loss += loss.item()
|
|
total_triplet_loss += triplet_loss.item()
|
|
total_center_loss += center_loss.item()
|
|
num_batches += 1
|
|
|
|
# 更新进度条
|
|
train_bar.set_postfix({
|
|
'Loss': f'{total_loss/num_batches:.4f}',
|
|
'Triplet': f'{total_triplet_loss/num_batches:.4f}',
|
|
'Center': f'{total_center_loss/num_batches:.4f}'
|
|
})
|
|
|
|
return (total_loss/num_batches,
|
|
total_triplet_loss/num_batches,
|
|
total_center_loss/num_batches)
|
|
|
|
|
|
def validate_epoch(model, val_loader, triplet_criterion, center_criterion, device):
|
|
"""
|
|
验证一个epoch
|
|
|
|
Args:
|
|
model: 模型
|
|
val_loader: 验证数据加载器
|
|
triplet_criterion: 三元组损失函数
|
|
center_criterion: 中心损失函数
|
|
device: 设备
|
|
|
|
Returns:
|
|
tuple: (平均损失, 三元组损失, 中心损失, 准确率)
|
|
"""
|
|
model.eval()
|
|
total_loss = 0.0
|
|
total_triplet_loss = 0.0
|
|
total_center_loss = 0.0
|
|
correct_predictions = 0
|
|
total_predictions = 0
|
|
num_batches = 0
|
|
|
|
with torch.no_grad():
|
|
val_bar = tqdm(val_loader, desc='验证中')
|
|
|
|
for anchor, positive, negative, labels in val_bar:
|
|
anchor = anchor.to(device)
|
|
positive = positive.to(device)
|
|
negative = negative.to(device)
|
|
labels = labels.to(device)
|
|
|
|
# 前向传播
|
|
anchor_emb = model(anchor)
|
|
positive_emb = model(positive)
|
|
negative_emb = model(negative)
|
|
|
|
# 计算损失
|
|
triplet_loss = triplet_criterion(anchor_emb, positive_emb, negative_emb)
|
|
center_loss = center_criterion(anchor_emb, labels)
|
|
loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss
|
|
|
|
# 计算准确率(基于最近邻分类)
|
|
# 这里简化为检查正样本距离是否小于负样本距离
|
|
pos_dist = F.pairwise_distance(anchor_emb, positive_emb, p=2)
|
|
neg_dist = F.pairwise_distance(anchor_emb, negative_emb, p=2)
|
|
correct = (pos_dist < neg_dist).sum().item()
|
|
|
|
# 统计
|
|
total_loss += loss.item()
|
|
total_triplet_loss += triplet_loss.item()
|
|
total_center_loss += center_loss.item()
|
|
correct_predictions += correct
|
|
total_predictions += anchor.size(0)
|
|
num_batches += 1
|
|
|
|
# 更新进度条
|
|
val_bar.set_postfix({
|
|
'Loss': f'{total_loss/num_batches:.4f}',
|
|
'Acc': f'{100.*correct_predictions/total_predictions:.2f}%'
|
|
})
|
|
|
|
accuracy = 100. * correct_predictions / total_predictions
|
|
|
|
return (total_loss/num_batches,
|
|
total_triplet_loss/num_batches,
|
|
total_center_loss/num_batches,
|
|
accuracy)
|
|
|
|
|
|
def plot_training_curves(train_losses, val_losses, train_accuracies, val_accuracies, save_path):
|
|
"""
|
|
绘制训练曲线
|
|
|
|
Args:
|
|
train_losses: 训练损失列表
|
|
val_losses: 验证损失列表
|
|
train_accuracies: 训练准确率列表
|
|
val_accuracies: 验证准确率列表
|
|
save_path: 保存路径
|
|
"""
|
|
fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(15, 10))
|
|
|
|
epochs = range(1, len(train_losses) + 1)
|
|
|
|
# 总损失
|
|
ax1.plot(epochs, train_losses, 'b-', label='训练损失')
|
|
ax1.plot(epochs, val_losses, 'r-', label='验证损失')
|
|
ax1.set_title('总损失变化')
|
|
ax1.set_xlabel('Epoch')
|
|
ax1.set_ylabel('Loss')
|
|
ax1.legend()
|
|
ax1.grid(True)
|
|
|
|
# 准确率
|
|
ax2.plot(epochs, train_accuracies, 'b-', label='训练准确率')
|
|
ax2.plot(epochs, val_accuracies, 'r-', label='验证准确率')
|
|
ax2.set_title('准确率变化')
|
|
ax2.set_xlabel('Epoch')
|
|
ax2.set_ylabel('Accuracy (%)')
|
|
ax2.legend()
|
|
ax2.grid(True)
|
|
|
|
# 三元组损失(如果有记录的话)
|
|
ax3.set_title('三元组损失变化')
|
|
ax3.set_xlabel('Epoch')
|
|
ax3.set_ylabel('Triplet Loss')
|
|
ax3.grid(True)
|
|
|
|
# 中心损失(如果有记录的话)
|
|
ax4.set_title('中心损失变化')
|
|
ax4.set_xlabel('Epoch')
|
|
ax4.set_ylabel('Center Loss')
|
|
ax4.grid(True)
|
|
|
|
plt.tight_layout()
|
|
plt.savefig(save_path, dpi=300, bbox_inches='tight')
|
|
plt.close()
|
|
|
|
|
|
def save_training_results(results, save_path):
|
|
"""
|
|
保存训练结果
|
|
|
|
Args:
|
|
results: 训练结果字典
|
|
save_path: 保存路径
|
|
"""
|
|
with open(save_path, 'w', encoding='utf-8') as f:
|
|
f.write("=== ResNet50 Embedding 训练结果 ===\n\n")
|
|
f.write(f"训练时间: {results['training_time']}\n")
|
|
f.write(f"总epoch数: {results['total_epochs']}\n")
|
|
f.write(f"最佳验证损失: {results['best_val_loss']:.4f}\n")
|
|
f.write(f"最佳验证准确率: {results['best_val_accuracy']:.2f}%\n")
|
|
f.write(f"最终训练损失: {results['final_train_loss']:.4f}\n")
|
|
f.write(f"最终验证损失: {results['final_val_loss']:.4f}\n")
|
|
f.write(f"最终训练准确率: {results['final_train_accuracy']:.2f}%\n")
|
|
f.write(f"最终验证准确率: {results['final_val_accuracy']:.2f}%\n")
|
|
f.write(f"模型参数数量: {results['model_parameters']:,}\n")
|
|
f.write(f"Embedding维度: {results['embedding_dim']}\n")
|
|
|
|
f.write("\n=== 训练配置 ===\n")
|
|
f.write(f"学习率: {results['learning_rate']}\n")
|
|
f.write(f"批次大小: {results['batch_size']}\n")
|
|
f.write(f"三元组损失边界: {results['triplet_margin']}\n")
|
|
f.write(f"中心损失权重: {results['center_loss_weight']}\n")
|
|
f.write(f"设备: {results['device']}\n")
|
|
|
|
|
|
def main(task_key: str = "dish"):
|
|
"""主训练函数"""
|
|
|
|
# 创建保存目录
|
|
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
|
cfg = TASKS[task_key]
|
|
save_dir = os.path.join(settings.BASE_DIR, 'model', cfg.name, f'embedding_{timestamp}')
|
|
os.makedirs(save_dir, exist_ok=True)
|
|
|
|
logger.info(f"[{cfg.name}] 模型保存目录: {save_dir}")
|
|
|
|
# 训练参数
|
|
EMBEDDING_DIM = cfg.embedding_dim
|
|
BATCH_SIZE = cfg.batch_size # 三元组训练通常使用较小的batch size
|
|
LEARNING_RATE = cfg.lr
|
|
NUM_EPOCHS = 50
|
|
TRIPLET_MARGIN = cfg.triplet_margin
|
|
CENTER_LOSS_WEIGHT = cfg.center_loss_weight
|
|
PATIENCE = 10
|
|
|
|
# 数据预处理
|
|
transform_train, transform_val = build_transforms(cfg.aug_strength)
|
|
|
|
# 创建数据集
|
|
train_dataset = TripletDataset(
|
|
dataset_path=cfg.train_dir,
|
|
transform=transform_train,
|
|
samples_per_class=550
|
|
)
|
|
|
|
val_dataset = TripletDataset(
|
|
dataset_path=cfg.val_dir,
|
|
transform=transform_val,
|
|
samples_per_class=50
|
|
)
|
|
|
|
# 创建数据加载器
|
|
train_loader = DataLoader(
|
|
train_dataset,
|
|
batch_size=BATCH_SIZE,
|
|
shuffle=True,
|
|
num_workers=0,
|
|
drop_last=True
|
|
)
|
|
|
|
val_loader = DataLoader(
|
|
val_dataset,
|
|
batch_size=BATCH_SIZE,
|
|
shuffle=False,
|
|
num_workers=0,
|
|
drop_last=False
|
|
)
|
|
|
|
logger.info(f"训练集大小: {len(train_dataset)}")
|
|
logger.info(f"验证集大小: {len(val_dataset)}")
|
|
logger.info(f"类别数量: {len(train_dataset.get_class_names())}")
|
|
logger.info(f"类别名称: {train_dataset.get_class_names()}")
|
|
|
|
# 创建模型
|
|
model = create_resnet50_embedding(
|
|
embedding_dim=EMBEDDING_DIM,
|
|
pretrained=True,
|
|
use_internal_preprocess=False
|
|
)
|
|
model = model.to(device)
|
|
|
|
# 创建损失函数
|
|
triplet_criterion = TripletLoss(margin=TRIPLET_MARGIN)
|
|
center_criterion = CenterLoss(
|
|
num_classes=len(train_dataset.get_class_names()),
|
|
embedding_dim=EMBEDDING_DIM,
|
|
alpha=0.5
|
|
).to(device)
|
|
|
|
# 创建优化器
|
|
optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE, weight_decay=1e-4)
|
|
center_optimizer = optim.Adam(center_criterion.parameters(), lr=LEARNING_RATE)
|
|
|
|
# 学习率调度器
|
|
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=15, gamma=0.5)
|
|
|
|
# 早停机制
|
|
early_stopping = EarlyStopping(patience=PATIENCE, min_delta=0.001)
|
|
|
|
# 训练记录
|
|
train_losses = []
|
|
val_losses = []
|
|
train_accuracies = []
|
|
val_accuracies = []
|
|
best_val_loss = float('inf')
|
|
best_val_accuracy = 0.0
|
|
|
|
logger.info("开始训练...")
|
|
start_time = time.time()
|
|
|
|
for epoch in range(NUM_EPOCHS):
|
|
# 训练
|
|
train_loss, train_triplet_loss, train_center_loss = train_epoch(
|
|
model, train_loader, triplet_criterion, center_criterion,
|
|
optimizer, center_optimizer, device, epoch
|
|
)
|
|
|
|
# 验证
|
|
val_loss, val_triplet_loss, val_center_loss, val_accuracy = validate_epoch(
|
|
model, val_loader, triplet_criterion, center_criterion, device
|
|
)
|
|
|
|
# 学习率调度
|
|
scheduler.step()
|
|
|
|
# 记录结果
|
|
train_losses.append(train_loss)
|
|
val_losses.append(val_loss)
|
|
train_accuracies.append(0) # 训练准确率暂时设为0,因为三元组训练中不容易计算
|
|
val_accuracies.append(val_accuracy)
|
|
|
|
# 保存最佳模型
|
|
if val_loss < best_val_loss:
|
|
best_val_loss = val_loss
|
|
best_val_accuracy = val_accuracy
|
|
|
|
# 保存模型
|
|
model_save_path = os.path.join(save_dir, 'best_embedding_model.pth')
|
|
torch.save({
|
|
'epoch': epoch,
|
|
'model_state_dict': model.state_dict(),
|
|
'center_state_dict': center_criterion.state_dict(),
|
|
'optimizer_state_dict': optimizer.state_dict(),
|
|
'center_optimizer_state_dict': center_optimizer.state_dict(),
|
|
'val_loss': val_loss,
|
|
'val_accuracy': val_accuracy,
|
|
'class_names': train_dataset.get_class_names(),
|
|
'embedding_dim': EMBEDDING_DIM,
|
|
}, model_save_path)
|
|
|
|
logger.info(f"保存最佳模型: Epoch {epoch+1}, Val Loss: {val_loss:.4f}, Val Acc: {val_accuracy:.2f}%")
|
|
|
|
# 打印epoch结果
|
|
logger.info(f"Epoch {epoch+1}/{NUM_EPOCHS}:")
|
|
logger.info(f" Train Loss: {train_loss:.4f} (Triplet: {train_triplet_loss:.4f}, Center: {train_center_loss:.4f})")
|
|
logger.info(f" Val Loss: {val_loss:.4f} (Triplet: {val_triplet_loss:.4f}, Center: {val_center_loss:.4f})")
|
|
logger.info(f" Val Accuracy: {val_accuracy:.2f}%")
|
|
logger.info(f" Learning Rate: {scheduler.get_last_lr()[0]:.6f}")
|
|
|
|
# 早停检查
|
|
if early_stopping(val_loss):
|
|
logger.info(f"早停触发,在第 {epoch+1} epoch停止训练")
|
|
break
|
|
|
|
# 训练完成
|
|
training_time = time.time() - start_time
|
|
logger.info(f"训练完成,总用时: {training_time/3600:.2f} 小时")
|
|
|
|
# 绘制训练曲线
|
|
curves_save_path = os.path.join(save_dir, 'training_curves.png')
|
|
plot_training_curves(train_losses, val_losses, train_accuracies, val_accuracies, curves_save_path)
|
|
|
|
# 保存训练结果
|
|
results = {
|
|
'training_time': f"{training_time/3600:.2f} 小时",
|
|
'total_epochs': len(train_losses),
|
|
'best_val_loss': best_val_loss,
|
|
'best_val_accuracy': best_val_accuracy,
|
|
'final_train_loss': train_losses[-1],
|
|
'final_val_loss': val_losses[-1],
|
|
'final_train_accuracy': train_accuracies[-1],
|
|
'final_val_accuracy': val_accuracies[-1],
|
|
'model_parameters': sum(p.numel() for p in model.parameters() if p.requires_grad),
|
|
'embedding_dim': EMBEDDING_DIM,
|
|
'learning_rate': LEARNING_RATE,
|
|
'batch_size': BATCH_SIZE,
|
|
'triplet_margin': TRIPLET_MARGIN,
|
|
'center_loss_weight': CENTER_LOSS_WEIGHT,
|
|
'device': str(device),
|
|
}
|
|
|
|
results_save_path = os.path.join(save_dir, 'training_results.txt')
|
|
save_training_results(results, results_save_path)
|
|
|
|
# 保存类别信息
|
|
class_info = {
|
|
'class_names': train_dataset.get_class_names(),
|
|
'class_to_idx': train_dataset.class_to_idx,
|
|
'idx_to_class': train_dataset.idx_to_class,
|
|
'embedding_dim': EMBEDDING_DIM,
|
|
}
|
|
|
|
class_info_path = os.path.join(save_dir, 'class_info.json')
|
|
with open(class_info_path, 'w', encoding='utf-8') as f:
|
|
json.dump(class_info, f, ensure_ascii=False, indent=2)
|
|
|
|
logger.info(f"训练结果已保存到: {save_dir}")
|
|
logger.info(f"最佳验证损失: {best_val_loss:.4f}")
|
|
logger.info(f"最佳验证准确率: {best_val_accuracy:.2f}%")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
import argparse
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish")
|
|
args = parser.parse_args()
|
|
main("whole_ingredient") |