Files
FoodClassifier/train/grid_search_cosface.py

445 lines
16 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参数
},
'whole_ingredient': {
's': [56.0, 60.0, 64.0, 68.0, 72.0],
'm': [0.32, 0.35, 0.38, 0.40, 0.45],
},
"""
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': [60.0,64.0],
'm': [0.42, 0.45],
},
'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,
save_dir: str,
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
early_stopped = False # 标记是否触发早停
# 训练历史
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}')
early_stopped = True
# ===== 保存最佳模型 =====
best_model_path = os.path.join(save_dir, f'best_model_s{s}_m{m}.pth')
torch.save({
'backbone_state_dict': model.state_dict(),
'head_state_dict': head.state_dict(),
's': s,
'm': m,
'best_val_loss': best_val_loss,
'best_epoch': best_epoch,
}, best_model_path)
logger.info(f'✓ 最佳模型已保存: {best_model_path}')
# ===== 生成特征向量可视化 =====
logger.info('🎨 开始生成特征向量可视化...')
try:
# 动态导入函数
faiss_db_dir = os.path.join(settings.BASE_DIR, 'faiss_vector_db')
if faiss_db_dir not in sys.path:
sys.path.insert(0, faiss_db_dir)
from faiss_vector_db.build_faiss_index import extract_embeddings_only
from faiss_vector_db.visualize_embeddings import visualize_embeddings_from_files
# 为每个配置创建独立的可视化目录
vis_output_dir = os.path.join(save_dir, f'model_s{s}_m{m}_visualization')
# 1. 提取特征向量
logger.info(' 步骤1/2: 提取训练集特征向量...')
extract_embeddings_only(
model_path=best_model_path,
train_dir=cfg.train_dir,
output_dir=vis_output_dir,
embedding_dim=cfg.embedding_dim,
batch_size=16 # 使用较小的batch_size加快速度
)
# 2. 生成可视化(使用PCA方法,快速)
logger.info(' 步骤2/2: 生成可视化图...')
embeddings_json = os.path.join(vis_output_dir, 'embeddings.json')
labels_json = os.path.join(vis_output_dir, 'labels.json')
visualize_embeddings_from_files(
embeddings_path=embeddings_json,
labels_path=labels_json,
output_dir=vis_output_dir,
method='pca', # 使用PCA方法(快速)
max_points=None, # 训练集全量可视化
seed=42
)
logger.info(f'✓ 可视化完成,保存至: {vis_output_dir}')
except Exception as e:
logger.error(f'⚠ 可视化生成失败: {e}')
import traceback
traceback.print_exc()
# ===== 可视化逻辑结束 =====
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': early_stopped,
'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,
save_dir=save_dir,
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
)