Files
FoodClassifier/train/train_cosface_embedding.py
T

440 lines
19 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
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=64.0, # 完整食材:类间区分度高,使用标准scale
cosface_m=0.38, # 较大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)