372 lines
16 KiB
Python
372 lines
16 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"
|
|
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'),
|
|
embedding_dim=512,
|
|
batch_size=32,
|
|
lr=5e-4,
|
|
aug_strength='medium',
|
|
cosface_s=60.0, # 菜品分类:类内差异大,使用中等scale
|
|
cosface_m=0.32, # 较小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'),
|
|
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'),
|
|
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
|
|
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: Optional[float] = None, m: Optional[float] = None, num_epochs: int = 60,
|
|
unknown_dir: Optional[str] = None, far: float = 0.05):
|
|
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}')
|
|
|
|
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=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('--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)
|