From f6180ed228426c41c88c16210f30e950d154833a Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Tue, 4 Nov 2025 14:53:31 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0=E9=85=8D=E7=BD=AE=E6=96=87?= =?UTF-8?q?=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- faiss_vector_db/build_faiss_index.py | 13 +++++++------ faiss_vector_db/visualize_embeddings.py | 8 ++++---- train/train_embedding.py | 4 ++-- 3 files changed, 13 insertions(+), 12 deletions(-) diff --git a/faiss_vector_db/build_faiss_index.py b/faiss_vector_db/build_faiss_index.py index 5ed87cf..c105ca0 100644 --- a/faiss_vector_db/build_faiss_index.py +++ b/faiss_vector_db/build_faiss_index.py @@ -480,14 +480,15 @@ def main(): """主函数""" # 配置参数 # MODEL_PATH = "../model/embedding_20251011_133653/best_embedding_model.pth" - # MODEL_PATH = "../model/ProcessedIngredientRecognition/embedding_20251029_170904/best_embedding_model.pth" - MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251030_135024/best_embedding_model.pth" + MODEL_PATH = "../model/ProcessedIngredientRecognition/embedding_20251103_172012/best_embedding_model.pth" + # MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251030_135024/best_embedding_model.pth" # MODEL_PATH = "../model/DishClassification/embedding_20251022_093635/best_embedding_model.pth" - # TRAIN_DIR = "../dataset/ProcessedIngredientRecognition/train" - TRAIN_DIR = "../dataset/WholeIngredientRecognition/train" + TRAIN_DIR = "../dataset/ProcessedIngredientRecognition/train" + # TRAIN_DIR = "../dataset/WholeIngredientRecognition/train" # TRAIN_DIR = "../dataset/DishClassification/train" - # OUTPUT_DIR = "ProcessedIngredientRecognition/faiss_index" - OUTPUT_DIR = "WholeIngredientRecognition/faiss_index" + + OUTPUT_DIR = "ProcessedIngredientRecognition/faiss_index" + # OUTPUT_DIR = "WholeIngredientRecognition/faiss_index" # OUTPUT_DIR = "DishClassification/faiss_index" BATCH_SIZE = 16 INDEX_TYPE = 'flat' # 'flat', 'ivf', 'hnsw' diff --git a/faiss_vector_db/visualize_embeddings.py b/faiss_vector_db/visualize_embeddings.py index ae6eb9a..cab4355 100644 --- a/faiss_vector_db/visualize_embeddings.py +++ b/faiss_vector_db/visualize_embeddings.py @@ -163,11 +163,11 @@ def plot_2d(Z: np.ndarray, y: np.ndarray, title: str, out_path: Optional[str] = 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 路径") - parser.add_argument("--embeddings", type=str, default=os.path.join("WholeIngredientRecognition/faiss_index", "embeddings.json"), help="embeddings.json 路径") - # parser.add_argument("--embeddings", type=str, default=os.path.join("ProcessedIngredientRecognition/faiss_index", "embeddings.json"), help="embeddings.json 路径") + # parser.add_argument("--embeddings", type=str, default=os.path.join("WholeIngredientRecognition/faiss_index", "embeddings.json"), help="embeddings.json 路径") + parser.add_argument("--embeddings", type=str, default=os.path.join("ProcessedIngredientRecognition/faiss_index", "embeddings.json"), help="embeddings.json 路径") # parser.add_argument("--labels", type=str, default=os.path.join("DishClassification/faiss_index", "labels.json"), help="labels.json 路径") - parser.add_argument("--labels", type=str, default=os.path.join("WholeIngredientRecognition/faiss_index", "labels.json"), help="labels.json 路径") - # parser.add_argument("--labels", type=str, default=os.path.join("ProcessedIngredientRecognition/faiss_index", "labels.json"), help="labels.json 路径") + # parser.add_argument("--labels", type=str, default=os.path.join("WholeIngredientRecognition/faiss_index", "labels.json"), help="labels.json 路径") + parser.add_argument("--labels", type=str, default=os.path.join("ProcessedIngredientRecognition/faiss_index", "labels.json"), help="labels.json 路径") parser.add_argument("--method", type=str, default="pca", choices=["pca", "tsne", "umap"], help="降维方法") parser.add_argument("--seed", type=int, default=42) parser.add_argument("--max_points", type=int, default=None, help="抽样上限,避免t-SNE/UMAP过慢;None为全量") diff --git a/train/train_embedding.py b/train/train_embedding.py index 2fa2809..1557b7f 100644 --- a/train/train_embedding.py +++ b/train/train_embedding.py @@ -821,6 +821,6 @@ if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish") args = parser.parse_args() - # main("processed_ingredient") - main("whole_ingredient") + main("processed_ingredient") + # main("whole_ingredient") # main("dish") \ No newline at end of file