修改训练脚本,增加网格搜索的可视化

This commit is contained in:
2025-12-04 13:53:18 +08:00
parent 3e271a61c2
commit d782083822
2 changed files with 65 additions and 49 deletions
+65 -2
View File
@@ -57,7 +57,7 @@ GRID_PARAMS = {
'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],
's': [60.0 ,64.0, 68.0],
'm': [0.32, 0.35, 0.38, 0.40],
},
'processed_ingredient': {
@@ -107,6 +107,7 @@ def train_single_config(
val_loader: DataLoader,
test_loader: DataLoader,
num_classes: int,
save_dir: str,
max_epochs: int = 100,
patience: int = 10,
min_epochs: int = 25,
@@ -142,6 +143,7 @@ def train_single_config(
best_val_loss = float('inf')
early_stop_counter = 0
best_epoch = 0
early_stopped = False # 标记是否触发早停
# 训练历史
train_losses = []
@@ -176,6 +178,66 @@ def train_single_config(
# 早停判断
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
@@ -190,7 +252,7 @@ def train_single_config(
's': s,
'm': m,
'actual_epochs': actual_epochs,
'early_stopped': actual_epochs < max_epochs,
'early_stopped': early_stopped,
'training_time_minutes': elapsed_time / 60,
'best_val_loss': best_val_loss,
'final_train_loss': train_losses[-1],
@@ -302,6 +364,7 @@ def grid_search_main(
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