增加网格搜索超参数功能。

This commit is contained in:
2025-11-11 18:03:50 +08:00
parent f0468e729b
commit b4b2b18ccd
5 changed files with 526 additions and 0 deletions
+149
View File
@@ -0,0 +1,149 @@
# CosFace超参数网格搜索使用指南
## 概述
已完成CosFace超参数网格搜索功能的实施,用于优化`scale (s)``margin (m)`两个核心超参数。
## 修改文件清单
### 1. `requirements.txt`
- ✅ 添加依赖:`pandas>=2.0.0``seaborn>=0.12.0`
### 2. `train/train_cosface_embedding.py`
-`TaskConfig`添加`test_dir`字段
-`TASKS`字典中所有任务添加测试集路径
### 3. `train/grid_search_cosface.py` (新建)
- ✅ 完整的网格搜索脚本 (~350行)
## 搜索空间配置
当前为dish任务配置的搜索空间:
```python
GRID_PARAMS = {
'dish': {
's': [56.0, 60.0, 64.0, 68.0], # 4个scale值
'm': [0.32, 0.35, 0.38, 0.40], # 4个margin值
}
}
```
**总配置数**: 4 × 4 = **16组实验**
## 使用方法
### 1. 安装依赖
```bash
pip install pandas>=2.0.0 seaborn>=0.12.0
```
### 2. 运行网格搜索
```bash
# 完整搜索(16组配置)
python train/grid_search_cosface.py --task dish
# 测试运行(限制配置数)
python train/grid_search_cosface.py --task dish --max_configs 2
# 自定义参数
python train/grid_search_cosface.py --task dish --epochs 100 --patience 10 --min_epochs 25
```
### 3. 参数说明
- `--task`: 任务名称 (`dish` / `whole_ingredient` / `processed_ingredient`)
- `--max_configs`: 限制最大配置数(用于测试,可选)
- `--epochs`: 每个配置的最大训练轮数(默认100)
- `--patience`: 早停容忍轮数(默认10
- `--min_epochs`: 最小训练轮数(默认25
## 输出文件
运行后会在`model/DishClassification/grid_search_YYYYMMDD_HHMMSS/`目录下生成:
1. **`grid_search_results.csv`** - 所有配置的详细结果表格
2. **`heatmap.png`** - 参数热力图(测试集准确率)
3. **`summary.txt`** - 搜索总结报告(含最佳配置)
4. **`model_s{s}_m{m}.pth`** - 每个配置的模型权重
5. **`grid_search.log`** - 完整训练日志
## 评估指标
- **主要指标**: 测试集准确率 (`test_acc`) - 用于选择最佳配置
- **辅助指标**: 验证集准确率、训练轮数、训练时间
## 应用最佳配置
网格搜索完成后:
1. 查看`summary.txt`找到最佳配置
2. 手动更新`train/train_cosface_embedding.py`中的`TASKS`字典:
```python
'dish': TaskConfig(
# ... 其他配置保持不变 ...
cosface_s=64.0, # 更新为最佳s值
cosface_m=0.38, # 更新为最佳m值
)
```
3. 后续训练将自动使用最佳配置
## 预计耗时
基于以下假设:
- 每个配置平均训练25-40个epoch(早停机制)
- 每个epoch约1-2分钟
- **单个配置**: ~30-60分钟
- **16组配置总耗时**: ~8-16小时
**建议**: 使用GPU运行,可在夜间或周末执行完整搜索。
## 注意事项
1. **数据集要求**: 确保`dataset/DishClassification/test/`目录存在且有数据
2. **GPU推荐**: 网格搜索计算量大,强烈建议使用GPU
3. **磁盘空间**: 每个配置约占用500MB,16组需8GB空间
4. **中断恢复**: 当前版本不支持断点续训,建议一次性完成
## 高级用法
### 修改搜索空间
编辑`train/grid_search_cosface.py`中的`GRID_PARAMS`字典:
```python
GRID_PARAMS = {
'dish': {
's': [60.0, 64.0, 68.0, 72.0], # 自定义scale范围
'm': [0.30, 0.35, 0.40, 0.45], # 自定义margin范围
}
}
```
### 调整早停策略
通过命令行参数调整:
```bash
python train/grid_search_cosface.py --task dish --patience 15 --min_epochs 30
```
## 故障排除
**问题1**: `ModuleNotFoundError: No module named 'pandas'`
- 解决: `pip install pandas seaborn`
**问题2**: 测试集目录不存在
- 解决: 确认`dataset/DishClassification/test/`路径正确且有数据
**问题3**: CUDA out of memory
- 解决: 减小`batch_size`或在CPU上运行(速度较慢)
## 示例结果解读
`summary.txt`示例:
```
最佳配置:
s (scale) = 64.0
m (margin) = 0.38
测试集准确率 = 98.50%
验证集准确率 = 100.00%
训练轮数 = 32
训练时间 = 45.3 分钟
```
这表示s=64.0, m=0.38是最优组合,在测试集上达到98.50%准确率。
+2
View File
@@ -28,3 +28,5 @@ typing_extensions==4.15.0
faiss-cpu>=1.7.0 faiss-cpu>=1.7.0
requests>=2.28.0 requests>=2.28.0
openai>=1.0.0 openai>=1.0.0
pandas>=2.0.0
seaborn>=0.12.0
View File
+371
View File
@@ -0,0 +1,371 @@
"""
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
)
+4
View File
@@ -35,6 +35,7 @@ class TaskConfig:
name: str name: str
train_dir: str train_dir: str
val_dir: str val_dir: str
test_dir: str # 测试集目录(用于网格搜索评估)
embedding_dim: int embedding_dim: int
batch_size: int batch_size: int
lr: float lr: float
@@ -48,6 +49,7 @@ TASKS = {
name='DishClassification', name='DishClassification',
train_dir=os.path.join(settings.BASE_DIR, 'dataset', 'DishClassification', 'train'), train_dir=os.path.join(settings.BASE_DIR, 'dataset', 'DishClassification', 'train'),
val_dir=os.path.join(settings.BASE_DIR, 'dataset', 'DishClassification', 'val'), val_dir=os.path.join(settings.BASE_DIR, 'dataset', 'DishClassification', 'val'),
test_dir=os.path.join(settings.BASE_DIR, 'dataset', 'DishClassification', 'test'),
embedding_dim=512, embedding_dim=512,
batch_size=32, batch_size=32,
lr=5e-4, lr=5e-4,
@@ -61,6 +63,7 @@ TASKS = {
name='WholeIngredientRecognition', name='WholeIngredientRecognition',
train_dir=os.path.join(settings.BASE_DIR, 'dataset', 'WholeIngredientRecognition', 'train'), train_dir=os.path.join(settings.BASE_DIR, 'dataset', 'WholeIngredientRecognition', 'train'),
val_dir=os.path.join(settings.BASE_DIR, 'dataset', 'WholeIngredientRecognition', 'val'), val_dir=os.path.join(settings.BASE_DIR, 'dataset', 'WholeIngredientRecognition', 'val'),
test_dir=os.path.join(settings.BASE_DIR, 'dataset', 'WholeIngredientRecognition', 'test'),
embedding_dim=512, embedding_dim=512,
batch_size=64, batch_size=64,
lr=8e-4, lr=8e-4,
@@ -72,6 +75,7 @@ TASKS = {
name='ProcessedIngredientRecognition', name='ProcessedIngredientRecognition',
train_dir=os.path.join(settings.BASE_DIR, 'dataset', 'ProcessedIngredientRecognition', 'train'), train_dir=os.path.join(settings.BASE_DIR, 'dataset', 'ProcessedIngredientRecognition', 'train'),
val_dir=os.path.join(settings.BASE_DIR, 'dataset', 'ProcessedIngredientRecognition', 'val'), val_dir=os.path.join(settings.BASE_DIR, 'dataset', 'ProcessedIngredientRecognition', 'val'),
test_dir=os.path.join(settings.BASE_DIR, 'dataset', 'ProcessedIngredientRecognition', 'test'),
embedding_dim=512, embedding_dim=512,
batch_size=32, batch_size=32,
lr=5e-4, lr=5e-4,