Files
FoodClassifier/train/train_embedding.py
T

825 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,
triplet_margin=0.5,
# center_loss_weight=0.1,
center_loss_weight=0.5,
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,
triplet_margin=0.5,
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,
triplet_margin=0.5,
# center_loss_weight=0.5,
center_loss_weight=20,
aug_strength="medium",
),
}
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,CENTER_LOSS_WEIGHT):
"""
训练一个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 # 可配置的中心损失权重
loss = triplet_loss + CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重
# print('triplet_loss',triplet_loss)
# print('center_loss',CENTER_LOSS_WEIGHT * center_loss)
# print('CENTER_LOSS_WEIGHT',CENTER_LOSS_WEIGHT)
# 反向传播
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,CENTER_LOSS_WEIGHT
)
# 验证
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("processed_ingredient")
# main("whole_ingredient")
# main("dish")