增加cosFace训练脚本。
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user