From 3e271a61c20c48475fecc0874fa25088ba2dfd9b Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Wed, 3 Dec 2025 18:39:15 +0800 Subject: [PATCH] =?UTF-8?q?=E5=9C=A8=E7=BD=91=E6=A0=BC=E6=90=9C=E7=B4=A2?= =?UTF-8?q?=E4=B8=AD=E5=8A=A0=E5=85=A5=E4=BA=86=E5=8F=AF=E8=A7=86=E5=8C=96?= =?UTF-8?q?=EF=BC=8C=E4=BD=86=E6=98=AF=E8=AE=AD=E7=BB=83=E8=84=9A=E6=9C=AC?= =?UTF-8?q?=E5=A5=BD=E5=83=8F=E7=BB=99=E6=94=B9=E9=94=99=E4=BA=86=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- faiss_vector_db/build_faiss_index.py | 58 ++++++++++++++++++ faiss_vector_db/visualize_embeddings.py | 78 +++++++++++++++++++++++++ train/grid_search_cosface.py | 8 ++- train/train_cosface_embedding.py | 47 +++++++++++++++ 4 files changed, 189 insertions(+), 2 deletions(-) diff --git a/faiss_vector_db/build_faiss_index.py b/faiss_vector_db/build_faiss_index.py index a99393c..827dd94 100644 --- a/faiss_vector_db/build_faiss_index.py +++ b/faiss_vector_db/build_faiss_index.py @@ -358,6 +358,64 @@ class FAISSIndexBuilder: return index +def extract_embeddings_only(model_path: str, + train_dir: str, + output_dir: str, + embedding_dim: int = 512, + batch_size: int = 16) -> Tuple[np.ndarray, List[int]]: + """ + 仅提取特征向量并保存为JSON(用于可视化) + 不构建完整的FAISS索引,节省时间 + + Args: + model_path: 模型路径 + train_dir: 训练数据目录 + output_dir: 输出目录 + embedding_dim: 特征向量维度 + batch_size: 批处理大小 + + Returns: + (embeddings, labels): 特征向量数组和标签列表 + """ + print("=" * 60) + print("开始提取特征向量用于可视化") + print("=" * 60) + + builder = FAISSIndexBuilder(model_path, embedding_dim) + + # 扫描数据 + image_paths, class_names, labels = builder.scan_training_data(train_dir) + builder.image_paths = image_paths + builder.labels = labels + + # 提取特征 + embeddings = builder.extract_features_batch(image_paths, batch_size) + builder.embeddings = embeddings + + # 仅保存 embeddings.json 和 labels.json(可视化需要的) + os.makedirs(output_dir, exist_ok=True) + + embeddings_json_path = os.path.join(output_dir, 'embeddings.json') + labels_json_path = os.path.join(output_dir, 'labels.json') + + print(f"保存特征向量到: {embeddings_json_path}") + with open(embeddings_json_path, "w", encoding="utf-8") as f: + json.dump(embeddings.tolist(), f, ensure_ascii=False) + + print(f"保存标签到: {labels_json_path}") + with open(labels_json_path, "w", encoding="utf-8") as f: + json.dump(labels, f, ensure_ascii=False) + + print("=" * 60) + print(f"✓ 特征向量提取完成") + print(f" 输出目录: {output_dir}") + print(f" 样本数: {len(embeddings)}") + print(f" 特征维度: {embedding_dim}") + print("=" * 60) + + return embeddings, labels + + class FAISSSearcher: """FAISS相似度检索器""" diff --git a/faiss_vector_db/visualize_embeddings.py b/faiss_vector_db/visualize_embeddings.py index ae6eb9a..5ddf248 100644 --- a/faiss_vector_db/visualize_embeddings.py +++ b/faiss_vector_db/visualize_embeddings.py @@ -160,6 +160,84 @@ def plot_2d(Z: np.ndarray, y: np.ndarray, title: str, out_path: Optional[str] = plt.show() +def visualize_embeddings_from_files(embeddings_path: str, + labels_path: str, + output_dir: str, + method: str = "pca", + max_points: Optional[int] = None, + seed: int = 42, + **kwargs) -> str: + """ + 从文件加载并可视化embeddings(用于训练脚本调用) + + Args: + embeddings_path: embeddings.json路径 + labels_path: labels.json路径 + output_dir: 输出目录 + method: 降维方法 (pca/tsne/umap) + max_points: 抽样上限 + seed: 随机种子 + **kwargs: 其他降维参数 + + Returns: + 输出的PNG图片路径 + """ + print("=" * 60) + print(f"开始可视化 embeddings") + print(f"降维方法: {method}") + print("=" * 60) + + # 加载数据 + print(f"加载数据: {embeddings_path}") + X = load_embeddings_json(embeddings_path) + y = load_labels_json(labels_path) + + # 数据对齐 + if X.shape[0] != y.shape[0]: + n = min(X.shape[0], y.shape[0]) + print(f"警告: 数据不一致,截断到 {n}") + X, y = X[:n], y[:n] + + # 抽样 + Xs, ys, _ = subsample(X, y, max_points, seed=seed) + if Xs.shape[0] < X.shape[0]: + print(f"已抽样: {Xs.shape[0]}/{X.shape[0]}") + + # 降维参数 + tsne_perplexity = kwargs.get('tsne_perplexity', 30) + umap_n_neighbors = kwargs.get('umap_n_neighbors', 15) + umap_min_dist = kwargs.get('umap_min_dist', 0.1) + + # 降维 + print(f"执行降维: {method}") + Z = reduce_dim(Xs, method, seed, tsne_perplexity, umap_n_neighbors, umap_min_dist) + + # 诊断 + print("\n==== 诊断信息 ====") + diag_info = diagnostics(Xs, ys, reduced2d=Z, method=method) + print(diag_info) + + # 保存诊断信息 + diag_path = os.path.join(output_dir, f"embedding_{method}_diagnostics.txt") + with open(diag_path, 'w', encoding='utf-8') as f: + f.write(diag_info) + print(f"✓ 诊断信息已保存: {diag_path}") + + # 绘图 + out_png = os.path.join(output_dir, f"embedding_{method}_2d.png") + title = f"Embedding {method.upper()} 2D (N={Xs.shape[0]})" + + # 为了在训练脚本中调用时不弹出窗口,我们需要关闭交互模式 + plt.ioff() # 关闭交互模式 + plot_2d(Z, ys, title, out_path=out_png) + plt.close('all') # 关闭所有图形 + + print(f"✓ 可视化图已保存: {out_png}") + print("=" * 60) + + return out_png + + def main(): parser = argparse.ArgumentParser(description="可视化高维 embedding 并进行坍塌诊断") # parser.add_argument("--embeddings", type=str, default=os.path.join("DishClassification/faiss_index", "embeddings.json"), help="embeddings.json 路径") diff --git a/train/grid_search_cosface.py b/train/grid_search_cosface.py index 6df8493..9d93c32 100644 --- a/train/grid_search_cosface.py +++ b/train/grid_search_cosface.py @@ -46,6 +46,10 @@ logger = logging.getLogger(__name__) '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, 72.0], + 'm': [0.32, 0.35, 0.38, 0.40, 0.45], + }, """ GRID_PARAMS = { 'dish': { @@ -53,8 +57,8 @@ 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, 72.0], - 'm': [0.32, 0.35, 0.38, 0.40, 0.45], + '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], diff --git a/train/train_cosface_embedding.py b/train/train_cosface_embedding.py index a213433..827c517 100644 --- a/train/train_cosface_embedding.py +++ b/train/train_cosface_embedding.py @@ -411,6 +411,53 @@ 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} 轮训练(未触发早停)')