Files
FoodClassifier/train/grid_search_cosface.py
T
2025-11-26 14:50:00 +08:00

378 lines
13 KiB
Python

"""
CosFace超参数网格搜索脚本
用法: python train/grid_search_cosface.py --task dish
"""
import os
import sys
import time
import json
import logging
import itertools
from datetime import datetime
from typing import List, Dict, Tuple
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from tqdm import tqdm
import pandas as pd
import seaborn as sns
import matplotlib.pyplot as plt
# 导入训练脚本中的组件
sys.path.append(os.path.join(os.path.dirname(__file__), '..'))
from train.train_cosface_embedding import (
TASKS, build_transforms, CosFaceHead, run_epoch, device
)
from net.resnet_embedding import create_resnet50_embedding
from settings import settings
# 日志配置
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(levelname)s - %(message)s',
handlers=[
logging.FileHandler('grid_search.log', encoding='utf-8'),
logging.StreamHandler()
]
)
logger = logging.getLogger(__name__)
# ==================== 网格搜索空间定义 ====================
"""
'dish': {
's': [56.0, 60.0, 64.0, 68.0], # scale参数
'm': [0.32, 0.35, 0.38, 0.40], # margin参数
},
"""
GRID_PARAMS = {
'dish': {
's': [56.0, 60.0, 64.0, 68.0], # scale参数
'm': [0.32, 0.35, 0.38, 0.40,0.45,0.50], # margin参数
},
'whole_ingredient': {
's': [56.0, 60.0, 64.0, 68.0],
'm': [0.32, 0.35, 0.38, 0.40],
},
'processed_ingredient': {
's': [56.0, 60.0, 64.0, 68.0],
'm': [0.30, 0.33, 0.36, 0.39],
},
}
def evaluate_on_test(model: nn.Module, head: nn.Module, test_loader: DataLoader) -> Tuple[float, float]:
"""在测试集上评估模型"""
model.eval()
head.eval()
criterion = nn.CrossEntropyLoss()
total_loss = 0.0
total_samples = 0
correct_samples = 0
with torch.no_grad():
for images, labels in tqdm(test_loader, desc='测试集评估', leave=False):
images = images.to(device)
labels = labels.to(device)
feats = model(images)
logits = head(feats) # 测试时不加margin
loss = criterion(logits, labels)
pred = logits.argmax(dim=1)
correct = (pred == labels).sum().item()
total_loss += loss.item() * labels.size(0)
total_samples += labels.size(0)
correct_samples += correct
avg_loss = total_loss / max(total_samples, 1)
avg_acc = 100.0 * correct_samples / max(total_samples, 1)
return avg_acc, avg_loss
def train_single_config(
cfg,
s: float,
m: float,
train_loader: DataLoader,
val_loader: DataLoader,
test_loader: DataLoader,
num_classes: int,
max_epochs: int = 100,
patience: int = 10,
min_epochs: int = 25,
min_delta: float = 0.001
) -> Tuple[nn.Module, nn.Module, Dict]:
"""训练单个超参数配置"""
logger.info(f'开始训练配置: s={s}, m={m}')
# 创建模型
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=max_epochs)
# 早停变量
best_val_loss = float('inf')
early_stop_counter = 0
best_epoch = 0
# 训练历史
train_losses = []
val_losses = []
train_accs = []
val_accs = []
start_time = time.time()
# 训练循环
for epoch in range(max_epochs):
# 训练一个epoch
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)
val_losses.append(va_loss)
train_accs.append(tr_acc)
val_accs.append(va_acc)
logger.info(f'Epoch {epoch+1}/{max_epochs} | Train Loss: {tr_loss:.4f}, Acc: {tr_acc:.2f}% | Val Loss: {va_loss:.4f}, Acc: {va_acc:.2f}%')
# 早停逻辑
if va_loss < best_val_loss - min_delta:
best_val_loss = va_loss
early_stop_counter = 0
best_epoch = epoch
else:
early_stop_counter += 1
# 早停判断
if epoch >= min_epochs and early_stop_counter >= patience:
logger.info(f'早停触发于Epoch {epoch+1}, 最佳Epoch: {best_epoch+1}')
break
elapsed_time = time.time() - start_time
actual_epochs = len(train_losses)
# 测试集评估
test_acc, test_loss = evaluate_on_test(model, head, test_loader)
logger.info(f'✓ 训练完成 | 耗时: {elapsed_time/60:.1f}分钟 | 实际Epoch: {actual_epochs} | Val Acc={val_accs[-1]:.2f}% | Test Acc={test_acc:.2f}%')
train_info = {
's': s,
'm': m,
'actual_epochs': actual_epochs,
'early_stopped': actual_epochs < max_epochs,
'training_time_minutes': elapsed_time / 60,
'best_val_loss': best_val_loss,
'final_train_loss': train_losses[-1],
'final_val_loss': val_losses[-1],
'final_train_acc': train_accs[-1],
'final_val_acc': val_accs[-1],
'test_acc': test_acc,
'test_loss': test_loss,
}
return model, head, train_info
def save_results_csv(results: List[Dict], save_path: str):
"""保存结果到CSV"""
df = pd.DataFrame(results)
df = df.sort_values('test_acc', ascending=False)
df.to_csv(save_path, index=False, encoding='utf-8-sig')
logger.info(f'✓ 结果CSV已保存: {save_path}')
# 打印前5名
logger.info('\n' + '='*60)
logger.info('测试集准确率TOP5配置:')
logger.info('='*60)
top5 = df.head(5)
for idx, row in top5.iterrows():
logger.info(f'Rank {idx+1}: s={row["s"]}, m={row["m"]} | Test Acc={row["test_acc"]:.2f}% | '
f'Val Acc={row["final_val_acc"]:.2f}% | Epochs={int(row["actual_epochs"])}')
logger.info('='*60 + '\n')
def visualize_heatmap(results: List[Dict], save_path: str):
"""生成参数热力图"""
df = pd.DataFrame(results)
# 创建透视表
pivot = df.pivot_table(values='test_acc', index='m', columns='s', aggfunc='mean')
plt.figure(figsize=(10, 8))
sns.heatmap(pivot, annot=True, fmt='.2f', cmap='YlGnBu', cbar_kws={'label': 'Test Accuracy (%)'})
plt.title('CosFace超参数网格搜索热力图 (测试集准确率)', fontsize=14, fontweight='bold')
plt.xlabel('Scale (s)', fontsize=12)
plt.ylabel('Margin (m)', fontsize=12)
plt.tight_layout()
plt.savefig(save_path, dpi=300, bbox_inches='tight')
plt.close()
logger.info(f'✓ 热力图已保存: {save_path}')
def grid_search_main(
task_key: str = 'dish',
max_configs: int = None,
max_epochs: int = 100,
patience: int = 10,
min_epochs: int = 25
):
"""网格搜索主函数"""
logger.info('='*80)
logger.info(f'开始CosFace超参数网格搜索 - 任务: {task_key}')
logger.info('='*80)
cfg = TASKS[task_key]
params = GRID_PARAMS[task_key]
# 生成所有参数组合
param_combinations = list(itertools.product(params['s'], params['m']))
total_configs = len(param_combinations)
if max_configs is not None and max_configs < total_configs:
param_combinations = param_combinations[:max_configs]
logger.info(f'⚠ 限制最大配置数为 {max_configs}')
logger.info(f'任务配置: {cfg.name}')
logger.info(f'搜索空间: s={params["s"]}, m={params["m"]}')
logger.info(f'总配置数: {len(param_combinations)}')
logger.info(f'早停参数: patience={patience}, min_epochs={min_epochs}')
logger.info(f'评估指标: 测试集准确率 (test_acc)')
logger.info('='*80 + '\n')
# 创建保存目录
timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
save_dir = os.path.join(settings.BASE_DIR, 'model', cfg.name, f'grid_search_{timestamp}')
os.makedirs(save_dir, exist_ok=True)
# 准备数据
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)
test_ds = datasets.ImageFolder(cfg.test_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)
test_loader = DataLoader(test_ds, batch_size=cfg.batch_size, shuffle=False, num_workers=0)
num_classes = len(train_ds.classes)
logger.info(f'数据集: 训练={len(train_ds)}, 验证={len(val_ds)}, 测试={len(test_ds)}, 类别数={num_classes}\n')
# 网格搜索
results = []
total_start = time.time()
for i, (s, m) in enumerate(param_combinations, 1):
logger.info(f'\n[配置 {i}/{len(param_combinations)}] s={s}, m={m}')
logger.info('-'*80)
try:
model, head, train_info = train_single_config(
cfg, s, m,
train_loader, val_loader, test_loader,
num_classes,
max_epochs=max_epochs,
patience=patience,
min_epochs=min_epochs
)
results.append(train_info)
# 保存模型
model_path = os.path.join(save_dir, f'model_s{s}_m{m}.pth')
torch.save({
'backbone_state_dict': model.state_dict(),
'head_state_dict': head.state_dict(),
's': s,
'm': m,
'test_acc': train_info['test_acc'],
'class_names': train_ds.classes,
}, model_path)
except Exception as e:
logger.error(f'配置 s={s}, m={m} 训练失败: {str(e)}')
continue
total_elapsed = time.time() - total_start
# 保存结果
csv_path = os.path.join(save_dir, 'grid_search_results.csv')
save_results_csv(results, csv_path)
heatmap_path = os.path.join(save_dir, 'heatmap.png')
visualize_heatmap(results, heatmap_path)
# 找到最佳配置
best_result = max(results, key=lambda x: x['test_acc'])
# 保存总结报告
summary_path = os.path.join(save_dir, 'summary.txt')
with open(summary_path, 'w', encoding='utf-8') as f:
f.write('='*80 + '\n')
f.write(f'CosFace超参数网格搜索总结 - {cfg.name}\n')
f.write('='*80 + '\n\n')
f.write(f'总耗时: {total_elapsed/3600:.2f} 小时\n')
f.write(f'搜索配置数: {len(results)}/{len(param_combinations)}\n')
f.write(f'评估指标: 测试集准确率\n\n')
f.write('最佳配置:\n')
f.write(f' s (scale) = {best_result["s"]}\n')
f.write(f' m (margin) = {best_result["m"]}\n')
f.write(f' 测试集准确率 = {best_result["test_acc"]:.2f}%\n')
f.write(f' 验证集准确率 = {best_result["final_val_acc"]:.2f}%\n')
f.write(f' 训练轮数 = {best_result["actual_epochs"]}\n')
f.write(f' 训练时间 = {best_result["training_time_minutes"]:.1f} 分钟\n\n')
f.write('使用建议:\n')
f.write(f'在 train/train_cosface_embedding.py 的 TASKS["{task_key}"] 中更新:\n')
f.write(f' cosface_s={best_result["s"]}\n')
f.write(f' cosface_m={best_result["m"]}\n')
logger.info(f'\n✓ 网格搜索完成! 总耗时: {total_elapsed/3600:.2f} 小时')
logger.info(f'✓ 最佳配置: s={best_result["s"]}, m={best_result["m"]}, Test Acc={best_result["test_acc"]:.2f}%')
logger.info(f'✓ 结果已保存到: {save_dir}')
if __name__ == '__main__':
import argparse
parser = argparse.ArgumentParser(description='CosFace超参数网格搜索')
# parser.add_argument('--task', choices=list(TASKS.keys()), default='dish', help='任务名称')
parser.add_argument('--task', choices=list(TASKS.keys()), default='whole_ingredient', help='任务名称')
parser.add_argument('--max_configs', type=int, default=None, help='最大配置数(用于测试)')
parser.add_argument('--epochs', type=int, default=100, help='每个配置的最大训练轮数')
parser.add_argument('--patience', type=int, default=10, help='早停容忍轮数')
parser.add_argument('--min_epochs', type=int, default=25, help='最小训练轮数')
args = parser.parse_args()
grid_search_main(
task_key=args.task,
max_configs=args.max_configs,
max_epochs=args.epochs,
patience=args.patience,
min_epochs=args.min_epochs
)