import os import sys import time import json import math import logging from dataclasses import dataclass from datetime import datetime from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms from tqdm import tqdm import matplotlib.pyplot as plt # 将项目根目录加入路径,便于导入 sys.path.append(os.path.join(os.path.dirname(__file__), '..')) from net.resnet_embedding import create_resnet50_embedding from settings import settings # 日志 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}') @dataclass class TaskConfig: name: str train_dir: str val_dir: str test_dir: str # 测试集目录(用于网格搜索评估) embedding_dim: int batch_size: int lr: float aug_strength: str # "strong" | "medium" | "shape" cosface_s: float = 64.0 # CosFace scale factor cosface_m: float = 0.35 # CosFace margin 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'), test_dir=os.path.join(settings.BASE_DIR, 'dataset', 'DishClassification', 'test'), embedding_dim=512, batch_size=32, lr=5e-4, aug_strength='medium', # cosface_s=60.0, # 菜品分类:类内差异大,使用中等scale cosface_s=64.0, # 菜品分类:类内差异大,使用中等scale # cosface_m=0.32, # 较小margin,适应类内多样性(不同做法、角度) cosface_m=0.40, # 较小margin,适应类内多样性(不同做法、角度) ), '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'), test_dir=os.path.join(settings.BASE_DIR, 'dataset', 'WholeIngredientRecognition', 'test'), embedding_dim=512, batch_size=64, lr=8e-4, aug_strength='medium', cosface_s=68.0, # 完整食材:类间区分度高,使用标准scale cosface_m=0.35, # 较大margin,强化类间分离(番茄vs土豆差异明显) ), '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'), test_dir=os.path.join(settings.BASE_DIR, 'dataset', 'ProcessedIngredientRecognition', 'test'), embedding_dim=512, batch_size=32, lr=5e-4, aug_strength='medium', cosface_s=60.0, # 加工食材:中等难度任务 cosface_m=0.33, # 中等margin,平衡类内多样性和类间区分 ), } 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]), ]) class CosFaceHead(nn.Module): def __init__(self, in_features: int, num_classes: int, s: float = 64.0, m: float = 0.35): super().__init__() self.s = s self.m = m self.weight = nn.Parameter(torch.randn(num_classes, in_features)) nn.init.xavier_normal_(self.weight) def forward(self, features: torch.Tensor, labels: Optional[torch.Tensor] = None): # 假设输入features已L2归一化(backbone已做),此处仍做一次以保证数值稳健 x = F.normalize(features, dim=1) W = F.normalize(self.weight, dim=1) cos_theta = torch.mm(x, W.t()) # [B, C] if labels is None: return self.s * cos_theta one_hot = F.one_hot(labels, num_classes=W.size(0)).float() cos_theta_m = cos_theta - one_hot * self.m logits = self.s * cos_theta_m return logits, cos_theta def accuracy_top1(logits: torch.Tensor, labels: torch.Tensor) -> float: pred = logits.argmax(dim=1) return (pred == labels).float().mean().item() * 100.0 def evaluate_open_set(cos_scores_known: torch.Tensor, cos_scores_unknown: torch.Tensor, far: float = 0.05) -> Tuple[float, float, float]: """ 基于max_cos分数的简单阈值估计:给定未知集FAR,返回阈值和两侧TPR/FPR。 返回: (threshold, known_accept_rate, unknown_reject_rate) """ # 阈值取未知集分布的(1 - FAR)分位数(使未知中约 FAR 被错误接受) threshold = torch.quantile(cos_scores_unknown, 1 - far).item() if len(cos_scores_unknown) > 0 else 0.5 known_accept = (cos_scores_known >= threshold).float().mean().item() if len(cos_scores_known) > 0 else 0.0 unknown_reject = (cos_scores_unknown < threshold).float().mean().item() if len(cos_scores_unknown) > 0 else 0.0 return threshold, known_accept, unknown_reject def plot_training_curves(train_losses, val_losses, train_accuracies, val_accuracies, save_path: str): epochs = range(1, len(train_losses) + 1) plt.figure(figsize=(12, 5)) plt.subplot(1, 2, 1) plt.plot(epochs, train_losses, label='Train Loss') plt.plot(epochs, val_losses, label='Val Loss') plt.xlabel('Epoch'); plt.ylabel('Loss'); plt.title('Loss Curve'); plt.grid(True); plt.legend() plt.subplot(1, 2, 2) plt.plot(epochs, train_accuracies, label='Train Acc') plt.plot(epochs, val_accuracies, label='Val Acc') plt.xlabel('Epoch'); plt.ylabel('Accuracy (%)'); plt.title('Accuracy Curve'); plt.grid(True); plt.legend() plt.tight_layout(); plt.savefig(save_path, dpi=300, bbox_inches='tight'); plt.close() def run_epoch(model, head, loader, criterion, optimizer=None, train: bool = True): if train: model.train(); head.train() else: model.eval(); head.eval() running_loss = 0.0 n_batches = 0 # 改为统计总样本数和正确样本数(全局准确率) total_samples = 0 correct_samples = 0 with torch.set_grad_enabled(train): pbar = tqdm(loader, desc='训练中' if train else '验证中') for images, labels in pbar: images = images.to(device); labels = labels.to(device) feats = model(images) # 已L2归一化 if train: logits, _ = head(feats, labels) loss = criterion(logits, labels) optimizer.zero_grad(); loss.backward(); optimizer.step() else: logits = head(feats) # 不加边距 loss = criterion(logits, labels) # 统计样本级别的正确数 pred = logits.argmax(dim=1) correct = (pred == labels).sum().item() batch_size = labels.size(0) total_samples += batch_size correct_samples += correct running_loss += loss.item() n_batches += 1 # 计算当前的全局准确率 current_acc = 100.0 * correct_samples / total_samples pbar.set_postfix({ 'Loss': f'{running_loss / n_batches:.4f}', 'Acc': f'{current_acc:.2f}%' }) avg_loss = running_loss / max(n_batches, 1) global_acc = 100.0 * correct_samples / max(total_samples, 1) return avg_loss, global_acc def collect_max_cos_scores(model, head, loader) -> torch.Tensor: model.eval(); head.eval() scores = [] with torch.no_grad(): for images, _ in loader: images = images.to(device) feats = model(images) _, cos = head(feats, labels=None), None # head(feats)返回logits= s*cos # 还需要原始cos分数:logits/s logits = head(feats) # s*cos cos_scores = logits / getattr(head, 's', 64.0) max_cos = cos_scores.max(dim=1).values scores.append(max_cos.cpu()) return torch.cat(scores, dim=0) if scores else torch.tensor([]) def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] = None, num_epochs: int = 60, unknown_dir: Optional[str] = None, far: float = 0.05, patience: int = 10, min_delta: float = 0.001, min_epochs: int = 25): """ 参数: patience: 早停容忍轮数(验证损失不改善的最大轮数) min_delta: 最小改善阈值(验证损失改善小于此值不算改善) min_epochs: 最小训练轮数(早停不会在此之前触发) """ cfg = TASKS[task_key] # 优先级:命令行参数 > 配置文件默认值 s = s if s is not None else cfg.cosface_s m = m if m is not None else cfg.cosface_m timestamp = datetime.now().strftime('%Y%m%d_%H%M%S') save_dir = os.path.join(settings.BASE_DIR, 'model', cfg.name, f'cosface_{timestamp}') os.makedirs(save_dir, exist_ok=True) logger.info(f'[{cfg.name}] 模型保存目录: {save_dir}') logger.info(f'[{cfg.name}] CosFace超参数: s={s}, m={m}') logger.info(f'[{cfg.name}] 早停参数: patience={patience}, min_delta={min_delta}, min_epochs={min_epochs}') transform_train, transform_val = build_transforms(cfg.aug_strength) # 数据集 train_ds = datasets.ImageFolder(cfg.train_dir, transform=transform_train) val_ds = datasets.ImageFolder(cfg.val_dir, transform=transform_val) train_loader = DataLoader(train_ds, batch_size=cfg.batch_size, shuffle=True, num_workers=0, drop_last=True) val_loader = DataLoader(val_ds, batch_size=cfg.batch_size, shuffle=False, num_workers=0) class_names = train_ds.classes num_classes = len(class_names) logger.info(f'类别数: {num_classes}; 类别: {class_names}') # 模型与头 model = create_resnet50_embedding(embedding_dim=cfg.embedding_dim, pretrained=True, use_internal_preprocess=False).to(device) head = CosFaceHead(in_features=cfg.embedding_dim, num_classes=num_classes, s=s, m=m).to(device) # 损失与优化器 criterion = nn.CrossEntropyLoss() optimizer = optim.AdamW(list(model.parameters()) + list(head.parameters()), lr=cfg.lr, weight_decay=1e-4) scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs) # 早停相关变量 best_val_acc = 0.0 best_val_loss = float('inf') early_stop_counter = 0 best_model_path = os.path.join(save_dir, 'best_cosface_model.pth') train_losses, val_losses, train_accs, val_accs = [], [], [], [] start = time.time() logger.info('开始训练...') for epoch in range(num_epochs): logger.info(f'Epoch {epoch+1}/{num_epochs}') tr_loss, tr_acc = run_epoch(model, head, train_loader, criterion, optimizer, train=True) va_loss, va_acc = run_epoch(model, head, val_loader, criterion, optimizer=None, train=False) scheduler.step() train_losses.append(tr_loss); train_accs.append(tr_acc) val_losses.append(va_loss); val_accs.append(va_acc) logger.info(f' Train Loss: {tr_loss:.4f}, Acc: {tr_acc:.2f}%') logger.info(f' Val Loss: {va_loss:.4f}, Acc: {va_acc:.2f}%') logger.info(f' LR: {optimizer.param_groups[0]["lr"]:.6f}') # 保存最佳准确率模型 if va_acc > best_val_acc: best_val_acc = va_acc torch.save({ 'epoch': epoch, 'backbone_state_dict': model.state_dict(), 'head_state_dict': head.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'val_acc': va_acc, 'val_loss': va_loss, 'class_names': class_names, 'embedding_dim': cfg.embedding_dim, 's': s, 'm': m, }, best_model_path) logger.info(f'✓ 保存最佳模型: {best_model_path} (Val Acc: {va_acc:.2f}%)') # 早停逻辑:基于验证损失 if va_loss < best_val_loss - min_delta: best_val_loss = va_loss early_stop_counter = 0 logger.info(f'✓ 验证损失改善: {va_loss:.4f} (改善 {best_val_loss - va_loss:.4f})') else: early_stop_counter += 1 logger.info(f'⚠ 验证损失未改善,早停计数: {early_stop_counter}/{patience}') # 早停判断(需满足最小epoch要求) if epoch >= min_epochs and early_stop_counter >= patience: logger.info('=' * 60) logger.info(f'🛑 早停触发:验证损失连续 {patience} 轮未改善超过 {min_delta}') logger.info(f' 触发轮次: Epoch {epoch+1}/{num_epochs}') logger.info(f' 最佳验证损失: {best_val_loss:.4f}') logger.info(f' 最佳验证准确率: {best_val_acc:.2f}%') logger.info('=' * 60) break else: logger.info(f'✓ 完成全部 {num_epochs} 轮训练(未触发早停)') elapsed = time.time() - start logger.info(f'训练完成,总耗时: {elapsed/3600:.2f} 小时') # 曲线 curves_path = os.path.join(save_dir, 'training_curves.png') plot_training_curves(train_losses, val_losses, train_accs, val_accs, curves_path) # 开放集评估(可选) threshold = None known_accept = None unknown_reject = None if unknown_dir is not None and os.path.isdir(unknown_dir): unknown_ds = datasets.ImageFolder(unknown_dir, transform=transform_val) if os.path.isdir(os.path.join(unknown_dir, os.listdir(unknown_dir)[0])) else datasets.ImageFolder(unknown_dir, transform=transform_val) unknown_loader = DataLoader(unknown_ds, batch_size=cfg.batch_size, shuffle=False, num_workers=0) # 已知集用验证集 known_scores = collect_max_cos_scores(model, head, val_loader) unknown_scores = collect_max_cos_scores(model, head, unknown_loader) if len(unknown_scores) > 0 and len(known_scores) > 0: threshold, known_accept, unknown_reject = evaluate_open_set(known_scores, unknown_scores, far=far) logger.info(f'开放集阈值@FAR={far*100:.1f}%: T={threshold:.4f}, 已知接受率={known_accept*100:.2f}%, 未知拒识率={unknown_reject*100:.2f}%') else: logger.warning('开放集评估数据不足,跳过阈值估计。') # 结果保存 actual_epochs = len(train_losses) early_stopped = actual_epochs < num_epochs results = { 'training_time': f'{elapsed/3600:.2f} 小时', 'total_epochs': actual_epochs, 'early_stopped': early_stopped, 'early_stop_epoch': actual_epochs if early_stopped else None, 'best_val_accuracy': best_val_acc, 'best_val_loss': best_val_loss, 'final_train_loss': train_losses[-1] if train_losses else None, 'final_val_loss': val_losses[-1] if val_losses else None, 'final_train_accuracy': train_accs[-1] if train_accs else None, 'final_val_accuracy': val_accs[-1] if val_accs else None, 'model_parameters': sum(p.numel() for p in list(model.parameters()) + list(head.parameters()) if p.requires_grad), 'embedding_dim': cfg.embedding_dim, 'learning_rate': cfg.lr, 'batch_size': cfg.batch_size, 's': s, 'm': m, 'patience': patience, 'min_delta': min_delta, 'min_epochs': min_epochs, 'device': str(device), 'threshold_max_cos': threshold, 'known_accept_rate_at_threshold': known_accept, 'unknown_reject_rate_at_threshold': unknown_reject, } with open(os.path.join(save_dir, 'training_results.txt'), 'w', encoding='utf-8') as f: f.write('=== ResNet50 + CosFace 训练结果 ===\n\n') for k, v in results.items(): f.write(f'{k}: {v}\n') # 保存类别信息 with open(os.path.join(save_dir, 'class_info.json'), 'w', encoding='utf-8') as f: json.dump({'class_names': class_names, 'num_classes': num_classes, 'embedding_dim': cfg.embedding_dim}, f, ensure_ascii=False, indent=2) logger.info(f'训练结果已保存到: {save_dir}') if __name__ == '__main__': import argparse parser = argparse.ArgumentParser() # parser.add_argument('task', choices=list(TASKS.keys()), nargs='?', default='dish') parser.add_argument('task', choices=list(TASKS.keys()), nargs='?', default='whole_ingredient') parser.add_argument('--s', type=float, default=None, help='CosFace scale factor (默认使用任务配置值)') parser.add_argument('--m', type=float, default=None, help='CosFace margin (默认使用任务配置值)') # parser.add_argument('--epochs', type=int, default=60) parser.add_argument('--epochs', type=int, default=100) parser.add_argument('--patience', type=int, default=10, help='早停容忍轮数') parser.add_argument('--min_delta', type=float, default=0.001, help='早停最小改善阈值') parser.add_argument('--min_epochs', type=int, default=25, help='最小训练轮数(早停保护)') parser.add_argument('--unknown_dir', type=str, default=None, help='开放集评估用未知类目录(可选)') parser.add_argument('--far', type=float, default=0.05, help='未知集允许的FAR,用于阈值估计') args = parser.parse_args() main(task_key=args.task, s=args.s, m=args.m, num_epochs=args.epochs, unknown_dir=args.unknown_dir, far=args.far, patience=args.patience, min_delta=args.min_delta, min_epochs=args.min_epochs)