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 embedding_dim: int batch_size: int lr: 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=32, lr=5e-4, 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=64, lr=8e-4, 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=32, lr=5e-4, 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]), ]) 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 running_acc = 0.0 n_batches = 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) acc = accuracy_top1(logits, labels) running_loss += loss.item() running_acc += acc n_batches += 1 pbar.set_postfix({ 'Loss': f'{running_loss / n_batches:.4f}', 'Acc': f'{running_acc / n_batches:.2f}%' }) return running_loss / max(n_batches, 1), running_acc / max(n_batches, 1) 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: float = 64.0, m: float = 0.35, num_epochs: int = 60, unknown_dir: Optional[str] = None, far: float = 0.05): cfg = TASKS[task_key] 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}') 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_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}%)') 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('开放集评估数据不足,跳过阈值估计。') # 结果保存 results = { 'training_time': f'{elapsed/3600:.2f} 小时', 'total_epochs': len(train_losses), 'best_val_accuracy': best_val_acc, '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, '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=64.0) parser.add_argument('--m', type=float, default=0.35) parser.add_argument('--epochs', type=int, default=60) 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)