修改训练脚本,增加网格搜索的可视化
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user