diff --git a/train/train_cosface_embedding.py b/train/train_cosface_embedding.py new file mode 100644 index 0000000..dfe8353 --- /dev/null +++ b/train/train_cosface_embedding.py @@ -0,0 +1,357 @@ +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('--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) diff --git a/train/train_embedding.py b/train/train_triplet_embedding.py similarity index 100% rename from train/train_embedding.py rename to train/train_triplet_embedding.py