Files
FoodClassifier/train/train_cosface_embedding.py
T
2025-11-07 09:40:54 +08:00

359 lines
15 KiB
Python

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)