From 1c8ceef5b007e6654254ddd687c4ad2e10360ba4 Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Thu, 23 Oct 2025 11:08:46 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E6=94=B9=E9=85=8D=E7=BD=AE=E5=92=8C?= =?UTF-8?q?=E8=B7=AF=E5=BE=84=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- classifier/embedding_food_classifier_app.py | 4 ++-- faiss_vector_db/build_faiss_index.py | 12 +++++------ faiss_vector_db/visualize_embeddings.py | 8 ++++---- toAndroid/toAndroidEmbedding.py | 4 ++-- train/train_embedding.py | 22 ++++++++++++++------- 5 files changed, 29 insertions(+), 21 deletions(-) diff --git a/classifier/embedding_food_classifier_app.py b/classifier/embedding_food_classifier_app.py index 5c800f5..4b9a460 100644 --- a/classifier/embedding_food_classifier_app.py +++ b/classifier/embedding_food_classifier_app.py @@ -65,8 +65,8 @@ class EmbeddingFoodClassifierApp: BASE_DIR = os.path.dirname(os.path.abspath(__file__)) # model_path = "../model/embedding_20251011_133653/best_embedding_model.pth" # model_path = os.path.join(BASE_DIR, "../model/embedding_20251011_133653/best_embedding_model.pth") - model_path = os.path.join(BASE_DIR, "../model/WholeIngredientRecognition/embedding_20251017_145836/best_embedding_model.pth") - # model_path = os.path.join(BASE_DIR, "../model/DishClassification/embedding_20251011_133653/best_embedding_model.pth") + model_path = os.path.join(BASE_DIR, "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model.pth") + # model_path = os.path.join(BASE_DIR, "../model/DishClassification/embedding_20251022_093635/best_embedding_model.pth") # FAISS索引目录 # index_dir = "../faiss_vector_db/faiss_index" diff --git a/faiss_vector_db/build_faiss_index.py b/faiss_vector_db/build_faiss_index.py index f134d19..7f01be5 100644 --- a/faiss_vector_db/build_faiss_index.py +++ b/faiss_vector_db/build_faiss_index.py @@ -480,12 +480,12 @@ def main(): """主函数""" # 配置参数 # MODEL_PATH = "../model/embedding_20251011_133653/best_embedding_model.pth" - MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251017_145836/best_embedding_model.pth" - # MODEL_PATH = "../model/DishClassification/embedding_20251011_133653/best_embedding_model.pth" - TRAIN_DIR = "../dataset/WholeIngredientRecognition/train" - # TRAIN_DIR = "../dataset/DishClassification/train" - OUTPUT_DIR = "WholeIngredientRecognition/faiss_index" - # OUTPUT_DIR = "DishClassification/faiss_index" + # MODEL_PATH = "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model.pth" + MODEL_PATH = "../model/DishClassification/embedding_20251022_093635/best_embedding_model.pth" + # TRAIN_DIR = "../dataset/WholeIngredientRecognition/train" + TRAIN_DIR = "../dataset/DishClassification/train" + # OUTPUT_DIR = "WholeIngredientRecognition/faiss_index" + OUTPUT_DIR = "DishClassification/faiss_index" BATCH_SIZE = 16 INDEX_TYPE = 'flat' # 'flat', 'ivf', 'hnsw' EMBEDDING_DIM = 512 diff --git a/faiss_vector_db/visualize_embeddings.py b/faiss_vector_db/visualize_embeddings.py index 1e67dc3..4f0f60d 100644 --- a/faiss_vector_db/visualize_embeddings.py +++ b/faiss_vector_db/visualize_embeddings.py @@ -162,10 +162,10 @@ 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("--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("--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("--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("--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/toAndroid/toAndroidEmbedding.py b/toAndroid/toAndroidEmbedding.py index 7d3c21f..5527013 100644 --- a/toAndroid/toAndroidEmbedding.py +++ b/toAndroid/toAndroidEmbedding.py @@ -18,7 +18,7 @@ def main(): # base_model = create_resnet50_embedding(embedding_dim=512, pretrained=True) base_model = create_mobile_resnet50_embedding(embedding_dim=512, pretrained=True) # model_path = "../model/embedding_20250930_102826/best_embedding_model.pth" - model_path = "../model/embedding_20251011_133653/best_embedding_model.pth" + model_path = "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model.pth" if not os.path.exists(model_path): print(f"错误:模型文件不存在 {model_path}") @@ -86,7 +86,7 @@ def main(): traced_model = torch.jit.trace(mobile_wrapper, single_input) # 保存模型 - output_path = "../model/embedding_20251011_133653/best_embedding_model_mobile.pt" + output_path = "../model/WholeIngredientRecognition/embedding_20251021_085915/best_embedding_model_mobile.pt" traced_model.save(output_path) print(f"✓ TorchScript模型保存成功: {output_path}") diff --git a/train/train_embedding.py b/train/train_embedding.py index 329a8c3..facae97 100644 --- a/train/train_embedding.py +++ b/train/train_embedding.py @@ -48,8 +48,10 @@ TASKS = { embedding_dim=512, batch_size=16, lr=1e-3, - triplet_margin=0.3, - center_loss_weight=0.1, + # triplet_margin=0.3, + triplet_margin=0.5, + # center_loss_weight=0.1, + center_loss_weight=0.5, aug_strength="medium", ), "whole_ingredient": TaskConfig( @@ -59,7 +61,8 @@ TASKS = { embedding_dim=512, batch_size=32, lr=8e-4, - triplet_margin=0.35, + # triplet_margin=0.35, + triplet_margin=0.5, center_loss_weight=0.1, aug_strength="medium", ), @@ -385,7 +388,7 @@ class EarlyStopping: def train_epoch(model, train_loader, triplet_criterion, center_criterion, - optimizer, center_optimizer, device, epoch): + optimizer, center_optimizer, device, epoch,CENTER_LOSS_WEIGHT): """ 训练一个epoch @@ -430,7 +433,11 @@ def train_epoch(model, train_loader, triplet_criterion, center_criterion, center_loss = center_criterion(anchor_emb, labels) # 总损失 - loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重 + # loss = triplet_loss + settings.CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重 + loss = triplet_loss + CENTER_LOSS_WEIGHT * center_loss # 可配置的中心损失权重 + # print('triplet_loss',triplet_loss) + # print('center_loss',CENTER_LOSS_WEIGHT * center_loss) + # print('CENTER_LOSS_WEIGHT',CENTER_LOSS_WEIGHT) # 反向传播 optimizer.zero_grad() @@ -709,7 +716,7 @@ def main(task_key: str = "dish"): # 训练 train_loss, train_triplet_loss, train_center_loss = train_epoch( model, train_loader, triplet_criterion, center_criterion, - optimizer, center_optimizer, device, epoch + optimizer, center_optimizer, device, epoch,CENTER_LOSS_WEIGHT ) # 验证 @@ -811,4 +818,5 @@ if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("task", choices=list(TASKS.keys()), nargs="?", default="dish") args = parser.parse_args() - main("whole_ingredient") \ No newline at end of file + # main("whole_ingredient") + main("dish") \ No newline at end of file