修改训练脚本,增加网格搜索的可视化
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参数
|
'm': [0.32, 0.35, 0.38, 0.40,0.45,0.50], # margin参数
|
||||||
},
|
},
|
||||||
'whole_ingredient': {
|
'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],
|
'm': [0.32, 0.35, 0.38, 0.40],
|
||||||
},
|
},
|
||||||
'processed_ingredient': {
|
'processed_ingredient': {
|
||||||
@@ -107,6 +107,7 @@ def train_single_config(
|
|||||||
val_loader: DataLoader,
|
val_loader: DataLoader,
|
||||||
test_loader: DataLoader,
|
test_loader: DataLoader,
|
||||||
num_classes: int,
|
num_classes: int,
|
||||||
|
save_dir: str,
|
||||||
max_epochs: int = 100,
|
max_epochs: int = 100,
|
||||||
patience: int = 10,
|
patience: int = 10,
|
||||||
min_epochs: int = 25,
|
min_epochs: int = 25,
|
||||||
@@ -142,6 +143,7 @@ def train_single_config(
|
|||||||
best_val_loss = float('inf')
|
best_val_loss = float('inf')
|
||||||
early_stop_counter = 0
|
early_stop_counter = 0
|
||||||
best_epoch = 0
|
best_epoch = 0
|
||||||
|
early_stopped = False # 标记是否触发早停
|
||||||
|
|
||||||
# 训练历史
|
# 训练历史
|
||||||
train_losses = []
|
train_losses = []
|
||||||
@@ -176,6 +178,66 @@ def train_single_config(
|
|||||||
# 早停判断
|
# 早停判断
|
||||||
if epoch >= min_epochs and early_stop_counter >= patience:
|
if epoch >= min_epochs and early_stop_counter >= patience:
|
||||||
logger.info(f'早停触发于Epoch {epoch+1}, 最佳Epoch: {best_epoch+1}')
|
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
|
break
|
||||||
|
|
||||||
elapsed_time = time.time() - start_time
|
elapsed_time = time.time() - start_time
|
||||||
@@ -190,7 +252,7 @@ def train_single_config(
|
|||||||
's': s,
|
's': s,
|
||||||
'm': m,
|
'm': m,
|
||||||
'actual_epochs': actual_epochs,
|
'actual_epochs': actual_epochs,
|
||||||
'early_stopped': actual_epochs < max_epochs,
|
'early_stopped': early_stopped,
|
||||||
'training_time_minutes': elapsed_time / 60,
|
'training_time_minutes': elapsed_time / 60,
|
||||||
'best_val_loss': best_val_loss,
|
'best_val_loss': best_val_loss,
|
||||||
'final_train_loss': train_losses[-1],
|
'final_train_loss': train_losses[-1],
|
||||||
@@ -302,6 +364,7 @@ def grid_search_main(
|
|||||||
cfg, s, m,
|
cfg, s, m,
|
||||||
train_loader, val_loader, test_loader,
|
train_loader, val_loader, test_loader,
|
||||||
num_classes,
|
num_classes,
|
||||||
|
save_dir=save_dir,
|
||||||
max_epochs=max_epochs,
|
max_epochs=max_epochs,
|
||||||
patience=patience,
|
patience=patience,
|
||||||
min_epochs=min_epochs
|
min_epochs=min_epochs
|
||||||
|
|||||||
@@ -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_loss:.4f}')
|
||||||
logger.info(f' 最佳验证准确率: {best_val_acc:.2f}%')
|
logger.info(f' 最佳验证准确率: {best_val_acc:.2f}%')
|
||||||
logger.info('=' * 60)
|
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
|
break
|
||||||
else:
|
else:
|
||||||
logger.info(f'✓ 完成全部 {num_epochs} 轮训练(未触发早停)')
|
logger.info(f'✓ 完成全部 {num_epochs} 轮训练(未触发早停)')
|
||||||
|
|||||||
Reference in New Issue
Block a user