372 lines
13 KiB
Python
372 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__)
|
|
|
|
# ==================== 网格搜索空间定义 ====================
|
|
GRID_PARAMS = {
|
|
'dish': {
|
|
's': [56.0, 60.0, 64.0, 68.0], # scale参数
|
|
'm': [0.32, 0.35, 0.38, 0.40], # 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
|
|
)
|