diff --git a/train/grid_search_cosface.py b/train/grid_search_cosface.py index 9d93c32..e72e4cf 100644 --- a/train/grid_search_cosface.py +++ b/train/grid_search_cosface.py @@ -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 diff --git a/train/train_cosface_embedding.py b/train/train_cosface_embedding.py index 827c517..a213433 100644 --- a/train/train_cosface_embedding.py +++ b/train/train_cosface_embedding.py @@ -411,53 +411,6 @@ def main(task_key: str = 'dish', s: Optional[float] = None, m: Optional[float] = logger.info(f' 最佳验证损失: {best_val_loss:.4f}') logger.info(f' 最佳验证准确率: {best_val_acc:.2f}%') logger.info('=' * 60) - - # ===== 生成特征向量可视化 ===== - 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 build_faiss_index import extract_embeddings_only - from visualize_embeddings import visualize_embeddings_from_files - - # 创建可视化输出目录 - vis_output_dir = os.path.join(save_dir, 'embeddings_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=cfg.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 else: logger.info(f'✓ 完成全部 {num_epochs} 轮训练(未触发早停)')