From ced0e4f7ef9a09b08f00adffced90282d877dc5d Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Thu, 13 Nov 2025 11:56:51 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0=E5=A4=9A=E6=A8=A1=E6=80=81?= =?UTF-8?q?=E6=A8=A1=E5=9E=8B=E5=BC=80=E6=94=BE=E9=9B=86=E8=AF=86=E5=88=AB?= =?UTF-8?q?=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- exp_multimodal/build_dish_name_index.py | 158 ++++++++++++++++++++++++ exp_multimodal/dish_name_cleaner.py | 102 +++++++++++++++ exp_multimodal/exp_multimodal2.py | 77 +++++++++--- exp_multimodal/exp_openset_vlm.py | 143 +++++++++++++++++++++ exp_multimodal/prompts.py | 28 +++++ exp_multimodal/vlm_classifier.py | 137 +++++++++++++++++++- train/grid_search_cosface.py | 4 +- 7 files changed, 628 insertions(+), 21 deletions(-) create mode 100644 exp_multimodal/build_dish_name_index.py create mode 100644 exp_multimodal/dish_name_cleaner.py create mode 100644 exp_multimodal/exp_openset_vlm.py diff --git a/exp_multimodal/build_dish_name_index.py b/exp_multimodal/build_dish_name_index.py new file mode 100644 index 0000000..03d79e7 --- /dev/null +++ b/exp_multimodal/build_dish_name_index.py @@ -0,0 +1,158 @@ +""" +构建菜品名向量索引 +将labels.py中的菜品名编码为向量并构建FAISS索引(一次性任务) +""" +import argparse +import json +import os +import sys +from typing import List + +import faiss +import numpy as np + +# 兼容:支持直接运行脚本或用 -m 模块方式运行 +if __name__ == "__main__" and __package__ is None: + sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from exp_multimodal.labels import build_labels +from exp_multimodal.text_embedder import OllamaEmbedder + + +def build_index( + dish_names: List[str], + embedder: OllamaEmbedder, + output_dir: str, + batch_size: int = 100, +) -> None: + """ + 构建FAISS索引 + + 参数: + dish_names: 菜品名列表 + embedder: OllamaEmbedder实例 + output_dir: 输出目录 + batch_size: 批量编码大小 + """ + os.makedirs(output_dir, exist_ok=True) + + print(f"[BuildIndex] Total dishes={len(dish_names)} batch_size={batch_size}") + + # 批量编码 + all_embeddings = [] + for i in range(0, len(dish_names), batch_size): + batch = dish_names[i:i+batch_size] + print(f"[BuildIndex] Encoding batch {i//batch_size + 1}/{(len(dish_names)-1)//batch_size + 1} (size={len(batch)})...") + + try: + batch_embs = embedder.encode(batch) + all_embeddings.append(batch_embs) + except Exception as e: + print(f"[BuildIndex] Error encoding batch {i//batch_size + 1}: {e}") + raise + + # 合并所有向量 + embeddings = np.vstack(all_embeddings) + print(f"[BuildIndex] Concatenated embeddings shape={embeddings.shape}") + + # 归一化向量(用于余弦相似度) + norms = np.linalg.norm(embeddings, axis=1, keepdims=True) + embeddings = embeddings / (norms + 1e-8) + print(f"[BuildIndex] Normalized embeddings") + + # 构建FAISS索引(IndexFlatIP = 内积索引,适合归一化后的向量) + dim = embeddings.shape[1] + index = faiss.IndexFlatIP(dim) + index.add(embeddings.astype(np.float32)) + print(f"[BuildIndex] Built FAISS index dim={dim} ntotal={index.ntotal}") + + # 保存文件 + names_path = os.path.join(output_dir, "dish_names.json") + embeddings_path = os.path.join(output_dir, "dish_embeddings.npy") + index_path = os.path.join(output_dir, "faiss_index.bin") + + with open(names_path, "w", encoding="utf-8") as f: + json.dump(dish_names, f, ensure_ascii=False, indent=2) + print(f"[BuildIndex] Saved dish names to {names_path}") + + np.save(embeddings_path, embeddings) + print(f"[BuildIndex] Saved embeddings to {embeddings_path}") + + faiss.write_index(index, index_path) + print(f"[BuildIndex] Saved FAISS index to {index_path}") + + print(f"[BuildIndex] ✅ Index build complete! Output dir: {output_dir}") + + +def main(): + ap = argparse.ArgumentParser(description="构建菜品名向量索引") + ap.add_argument( + "--mode", + choices=["dish", "whole", "processed"], + default="dish", + help="数据集模式(默认: dish)" + ) + ap.add_argument( + "--output", + default=None, + help="输出目录(默认: faiss_vector_db/_names)" + ) + ap.add_argument( + "--embedder_url", + default="http://192.168.1.250:11434", + help="Ollama服务地址" + ) + ap.add_argument( + "--embedder_model", + default="quentinz/bge-large-zh-v1.5", + help="Embedding模型名称" + ) + ap.add_argument( + "--batch_size", + type=int, + default=100, + help="批量编码大小" + ) + + args = ap.parse_args() + + # 确定输出目录 + if args.output: + output_dir = args.output + else: + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + output_dir = os.path.join(project_root, "faiss_vector_db", f"{args.mode}_names") + + print(f"[Main] mode={args.mode} output_dir={output_dir}") + print(f"[Main] embedder_url={args.embedder_url} model={args.embedder_model}") + + # 构建菜品名列表 + print(f"[Main] Building labels from mode={args.mode}...") + dish_names = build_labels(args.mode, alias_map_path=None) + print(f"[Main] Built {len(dish_names)} dish names") + + if not dish_names: + print("[Main] ❌ No dish names found, abort") + return + + # 初始化Embedder + print(f"[Main] Initializing OllamaEmbedder...") + embedder = OllamaEmbedder( + base_url=args.embedder_url, + model=args.embedder_model, + ) + + # 构建索引 + print(f"[Main] Building FAISS index...") + build_index( + dish_names=dish_names, + embedder=embedder, + output_dir=output_dir, + batch_size=args.batch_size, + ) + + print(f"[Main] ✅ All done!") + + +if __name__ == "__main__": + main() diff --git a/exp_multimodal/dish_name_cleaner.py b/exp_multimodal/dish_name_cleaner.py new file mode 100644 index 0000000..2bc48ae --- /dev/null +++ b/exp_multimodal/dish_name_cleaner.py @@ -0,0 +1,102 @@ +""" +菜品名清洗工具 +用于将VLM输出的菜品名进行标准化处理,提升向量匹配准确率 +""" +import re +from typing import Dict + + +# 同义词映射表(可根据实际情况扩展) +SYNONYM_MAP: Dict[str, str] = { + "西红柿": "番茄", + "土豆": "马铃薯", + "洋芋": "马铃薯", + "青椒": "柿子椒", + # 可继续添加... +} + + +def clean_dish_name(name: str) -> str: + """ + 清洗菜品名:去除括号注释、前缀、英文等干扰信息 + + 示例: + "宫保鸡丁(川菜)" -> "宫保鸡丁" + "川菜-麻婆豆腐" -> "麻婆豆腐" + "红烧肉 Braised Pork" -> "红烧肉" + """ + if not isinstance(name, str): + return "" + + # 1. 去除括号及内容(中英文括号) + name = re.sub(r"[((].*?[))]", "", name) + name = re.sub(r"\[.*?\]", "", name) + + # 2. 去除常见前缀(菜系、地域等) + prefixes = ["川菜", "粤菜", "鲁菜", "苏菜", "浙菜", "闽菜", "湘菜", "徽菜", + "东北", "西北", "西南", "华南", "华北"] + for prefix in prefixes: + if name.startswith(prefix): + name = name[len(prefix):] + break + + # 3. 去除分隔符后的前缀(如 "川菜-宫保鸡丁") + name = re.sub(r"^[^-—]*[-—]", "", name) + + # 4. 去除英文部分(保留中文) + name = re.sub(r"[a-zA-Z\s]+", "", name) + + # 5. 去除多余空格和标点 + name = re.sub(r"[,,、·\s]+", "", name) + + # 6. 去除可能的烹饪方式后缀(如果VLM违规输出) + cooking_suffixes = ["炒制", "烹饪", "料理", "做法"] + for suffix in cooking_suffixes: + if name.endswith(suffix): + name = name[:-len(suffix)] + + return name.strip() + + +def normalize_dish_name(name: str) -> str: + """ + 标准化菜品名:应用同义词映射 + + 示例: + "西红柿炒鸡蛋" -> "番茄炒鸡蛋" + """ + cleaned = clean_dish_name(name) + + # 应用同义词替换 + for synonym, standard in SYNONYM_MAP.items(): + cleaned = cleaned.replace(synonym, standard) + + return cleaned + + +def extract_main_dish_name(text: str) -> str: + """ + 从VLM返回文本中提取主要菜品名(兼容多种输出格式) + + 示例: + "这是宫保鸡丁" -> "宫保鸡丁" + "菜品:红烧肉" -> "红烧肉" + """ + if not text: + return "" + + # 尝试匹配常见模式 + patterns = [ + r"菜品[::]\s*([^,,。\n]+)", + r"识别为[::]\s*([^,,。\n]+)", + r"这是\s*([^,,。\n]+)", + r"应该是\s*([^,,。\n]+)", + ] + + for pattern in patterns: + match = re.search(pattern, text) + if match: + return match.group(1).strip() + + # 如果没有匹配到,返回清洗后的整个文本 + return clean_dish_name(text) diff --git a/exp_multimodal/exp_multimodal2.py b/exp_multimodal/exp_multimodal2.py index 0418155..6134f7e 100644 --- a/exp_multimodal/exp_multimodal2.py +++ b/exp_multimodal/exp_multimodal2.py @@ -10,8 +10,11 @@ if __name__ == "__main__" and __package__ is None: # 统一使用绝对导入,避免相对导入在脚本直跑时失败 from exp_multimodal.labels import build_labels -from exp_multimodal.vlm_classifier import classify_image +from exp_multimodal.vlm_classifier import classify_image, classify_image_openset from exp_multimodal.ollama_client import OLLAMA_URL, DEFAULT_MODEL +from exp_multimodal.text_embedder import OllamaEmbedder +from exp_multimodal.vector_matcher import DishNameMatcher +from exp_multimodal.vlm_providers.ollama_provider import OllamaVLMProvider FEWSHOT_HINTS: Dict[str, str] = { @@ -24,9 +27,14 @@ FEWSHOT_HINTS: Dict[str, str] = { def main() -> None: ap = argparse.ArgumentParser() - ap.add_argument("--mode", choices=["dish", "whole", "processed"], default="dish") + ap.add_argument("--mode", choices=["dish", "whole", "processed", "openset_dish"], default="openset_dish") ap.add_argument("--image", default=r"D:\MyProjects\PythonProjects\FoodClassifier\dataset\DishClassification\test\红烧肉\img04.png") ap.add_argument("--alias_map", default=None) + + # 开放式识别专用参数 + ap.add_argument("--openset-top-k", type=int, default=3, help="开放式识别:向量检索Top-K候选数") + ap.add_argument("--openset-min-score", type=float, default=0.5, help="开放式识别:最低匹配分数阈值") + ap.add_argument("--openset-index-path", default="exp_multimodal/dish_name_index", help="开放式识别:向量索引路径") args = ap.parse_args() @@ -37,24 +45,59 @@ def main() -> None: ) print(f"[Main] Image exists={os.path.exists(args.image)} size={os.path.getsize(args.image) if os.path.exists(args.image) else 'N/A'}") - labels = build_labels(args.mode, args.alias_map) - ingredient_only = args.mode in {"whole", "processed"} - fewshot = FEWSHOT_HINTS if args.mode == "dish" else None - print(f"[Main] Built labels count={len(labels)} ingredient_only={ingredient_only} fewshot={bool(fewshot)}") + # 开放式识别分支 + if args.mode == "openset_dish": + print(f"[Main] Openset mode: top_k={args.openset_top_k} min_score={args.openset_min_score} index_path={args.openset_index_path}") + + # 初始化Provider、Embedder和Matcher + provider = OllamaVLMProvider(base_url=OLLAMA_URL, model_name=DEFAULT_MODEL) + embedder = OllamaEmbedder(base_url=OLLAMA_URL, model_name="bge-large-zh-v1.5") + matcher = DishNameMatcher(index_dir=args.openset_index_path) + + print("[Main] Loading openset components...") + matcher.load() + print(f"[Main] Loaded {len(matcher.dish_names)} dish names from index") + + res = classify_image_openset( + image_path=args.image, + provider=provider, + embedder=embedder, + matcher=matcher, + top_k=args.openset_top_k, + min_match_score=args.openset_min_score, + ) + + print( + json.dumps( + {"mode": args.mode, "result": res}, + ensure_ascii=False, + indent=2, + ) + ) + + # 封闭式识别分支(原有逻辑) + else: + labels = build_labels(args.mode, args.alias_map) + ingredient_only = args.mode in {"whole", "processed"} + fewshot = FEWSHOT_HINTS if args.mode == "dish" else None + print(f"[Main] Built labels count={len(labels)} ingredient_only={ingredient_only} fewshot={bool(fewshot)}") - res = classify_image( - args.image, - labels, - fewshot_hints=fewshot, - ingredient_only=ingredient_only, - ) + provider = OllamaVLMProvider(base_url=OLLAMA_URL, model_name=DEFAULT_MODEL) + + res = classify_image( + args.image, + labels, + provider=provider, + fewshot_hints=fewshot, + ingredient_only=ingredient_only, + ) - print( - json.dumps( - {"mode": args.mode, "result": res, "num_labels": len(labels)}, - ensure_ascii=False, + print( + json.dumps( + {"mode": args.mode, "result": res, "num_labels": len(labels)}, + ensure_ascii=False, + ) ) - ) if __name__ == "__main__": diff --git a/exp_multimodal/exp_openset_vlm.py b/exp_multimodal/exp_openset_vlm.py new file mode 100644 index 0000000..6b6fbb6 --- /dev/null +++ b/exp_multimodal/exp_openset_vlm.py @@ -0,0 +1,143 @@ +""" +开放式VLM识别 + 向量匹配 主程序 +用于测试"开放识别+向量检索"方案的性能 +""" +import argparse +import json +import os +import sys +import time + +# 兼容:支持直接运行脚本或用 -m 模块方式运行 +if __name__ == "__main__" and __package__ is None: + sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from exp_multimodal.vlm_classifier import classify_image_openset +from exp_multimodal.text_embedder import OllamaEmbedder +from exp_multimodal.vector_matcher import DishNameMatcher +from exp_multimodal.vlm_providers.ollama_provider import OllamaProvider + + +def main(): + ap = argparse.ArgumentParser(description="开放式菜品识别测试") + ap.add_argument( + "--image", + default=r"D:\MyProjects\PythonProjects\FoodClassifier\dataset\DishClassification\test\红烧肉\img04.png", + help="测试图片路径" + ) + ap.add_argument( + "--index_dir", + default=None, + help="FAISS索引目录(默认: faiss_vector_db/dish_names)" + ) + ap.add_argument( + "--vlm_url", + default="http://192.168.1.250:11434", + help="VLM服务地址(Ollama)" + ) + ap.add_argument( + "--vlm_model", + default="qwen2.5vl:32b", + help="VLM模型名称" + ) + ap.add_argument( + "--embedder_url", + default="http://192.168.1.250:11434", + help="Embedding服务地址(Ollama)" + ) + ap.add_argument( + "--embedder_model", + default="quentinz/bge-large-zh-v1.5", + help="Embedding模型名称" + ) + ap.add_argument( + "--top_k", + type=int, + default=3, + help="向量检索返回的候选数" + ) + ap.add_argument( + "--min_score", + type=float, + default=0.5, + help="最低匹配分数阈值" + ) + + args = ap.parse_args() + + # 确定索引目录 + if args.index_dir: + index_dir = args.index_dir + else: + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + index_dir = os.path.join(project_root, "faiss_vector_db", "dish_names") + + print("=" * 80) + print(f"[Main] 开放式VLM识别 + 向量匹配测试") + print("=" * 80) + print(f"[Main] Image: {args.image}") + print(f"[Main] Image exists: {os.path.exists(args.image)}") + if os.path.exists(args.image): + print(f"[Main] Image size: {os.path.getsize(args.image)} bytes") + print(f"[Main] Index dir: {index_dir}") + print(f"[Main] VLM: {args.vlm_url} / {args.vlm_model}") + print(f"[Main] Embedder: {args.embedder_url} / {args.embedder_model}") + print(f"[Main] Top-K: {args.top_k}, Min score: {args.min_score}") + print("=" * 80) + + # 初始化组件 + print(f"\n[Main] 初始化VLM Provider...") + vlm_provider = OllamaProvider( + base_url=args.vlm_url, + model=args.vlm_model, + ) + + print(f"\n[Main] 初始化Embedder...") + embedder = OllamaEmbedder( + base_url=args.embedder_url, + model=args.embedder_model, + ) + + print(f"\n[Main] 加载FAISS索引...") + matcher = DishNameMatcher(index_dir) + + # 执行识别 + print(f"\n[Main] 开始识别...") + print("=" * 80) + + t_start = time.time() + result = classify_image_openset( + image_path=args.image, + provider=vlm_provider, + embedder=embedder, + matcher=matcher, + top_k=args.top_k, + min_match_score=args.min_score, + ) + t_elapsed = time.time() - t_start + + print("=" * 80) + print(f"\n[Main] ✅ 识别完成!耗时: {t_elapsed:.2f}秒") + print("=" * 80) + print("\n识别结果:") + print(json.dumps(result, ensure_ascii=False, indent=2)) + print("=" * 80) + + # 简要总结 + print(f"\n📊 性能总结:") + print(f" - 总耗时: {t_elapsed:.2f}秒") + print(f" - VLM原始输出: {result['raw_dish']}") + print(f" - 清洗后菜名: {result['cleaned_dish']}") + print(f" - VLM置信度: {result['vlm_confidence']:.4f}") + print(f" - 最佳匹配: {result['best_match']}") + print(f" - 最终置信度: {result['final_confidence']:.4f}") + print(f" - 候选数量: {len(result['candidates'])}") + + if t_elapsed <= 3.0: + print(f"\n✅ 性能达标!({t_elapsed:.2f}s ≤ 3.0s)") + else: + print(f"\n⚠️ 性能未达标 ({t_elapsed:.2f}s > 3.0s)") + + +if __name__ == "__main__": + main() diff --git a/exp_multimodal/prompts.py b/exp_multimodal/prompts.py index a612d46..fb849fc 100644 --- a/exp_multimodal/prompts.py +++ b/exp_multimodal/prompts.py @@ -38,3 +38,31 @@ def build_closedset_prompt( ) +def build_openset_prompt() -> str: + """ + 构建开放式识别提示词(不提供候选列表,让VLM自由识别) + 用于"开放识别+向量匹配"方案 + """ + return ( + "【重要】你必须且只能输出JSON格式,严禁输出任何其他内容(包括英文描述、解释、分析等)。\n\n" + "你是专业的中国菜品识别专家。任务:识别图片中的菜品并输出标准中文名称。\n\n" + "要求:\n" + "1. 只输出菜品的标准中文名称,例如:宫保鸡丁、红烧肉、西红柿炒鸡蛋\n" + "2. 禁止输出地域/菜系分类(如川菜、粤菜)\n" + "3. 禁止输出烹饪方式后缀(如炒制、烹饪),除非该方式是菜名的一部分(如红烧肉)\n" + "4. 禁止输出括号注释(如(川菜)、(辣味))\n" + "5. 禁止输出英文翻译\n" + "6. 给出0-1的置信度分数,表示你对识别结果的确信程度\n" + "7. 严格按以下JSON格式输出,不允许有任何偏差\n\n" + '输出格式(必须严格遵守):\n{"dish":"<菜品标准中文名>","confidence":0.95}\n\n' + "示例输出:\n" + '{"dish":"宫保鸡丁","confidence":0.92}\n' + '{"dish":"红烧肉","confidence":0.88}\n' + '{"dish":"西红柿炒鸡蛋","confidence":0.95}\n\n' + "【再次强调】\n" + "- 只输出上述JSON格式,不要有任何额外文字\n" + "- 菜品名必须是标准中文,不带任何前缀、后缀、括号、英文\n" + "- 如果无法识别,输出confidence接近0的值" + ) + + diff --git a/exp_multimodal/vlm_classifier.py b/exp_multimodal/vlm_classifier.py index 75de732..3a6022a 100644 --- a/exp_multimodal/vlm_classifier.py +++ b/exp_multimodal/vlm_classifier.py @@ -1,9 +1,12 @@ import json import re -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Tuple from .vlm_providers.base import VLMProvider -from .prompts import build_closedset_prompt +from .prompts import build_closedset_prompt, build_openset_prompt +from .dish_name_cleaner import normalize_dish_name +from .text_embedder import OllamaEmbedder +from .vector_matcher import DishNameMatcher def _extract_json_obj(text: str): @@ -93,3 +96,133 @@ def classify_image( return {"label": "Unknown", "confidence": 0.0} +def classify_image_openset( + image_path: str, + provider: VLMProvider, + embedder: OllamaEmbedder, + matcher: DishNameMatcher, + top_k: int = 3, + min_match_score: float = 0.5, +) -> Dict: + """ + 开放式菜品识别 + 向量匹配方案 + + 流程: + 1. VLM开放式识别(不给候选列表) + 2. 清洗菜品名 + 3. 文本向量化 + 4. FAISS检索Top-K候选 + 5. 综合置信度排序 + + 参数: + image_path: 图片路径 + provider: VLM Provider实例 + embedder: Ollama Embedder实例 + matcher: 菜品名匹配器实例 + top_k: 向量检索返回的候选数 + min_match_score: 最低匹配分数阈值 + + 返回: + { + "raw_dish": "VLM原始输出", + "cleaned_dish": "清洗后的菜品名", + "vlm_confidence": 0.95, + "candidates": [ + {"dish": "匹配的标准菜名", "match_score": 0.88, "final_score": 0.83}, + ... + ], + "best_match": "最佳匹配菜名", + "final_confidence": 0.83 + } + """ + print(f"[VLM_Openset] Start classify image={image_path} top_k={top_k}") + + # 步骤1: VLM开放式识别 + prompt = build_openset_prompt() + print(f"[VLM_Openset] Prompt length={len(prompt)}") + + print(f"[VLM_Openset] Calling {provider}.chat_vision()...") + text = provider.chat_vision(prompt, [image_path], temperature=0.1) + _preview = (str(text)[:200]).replace("\n", " ") + print(f"[VLM_Openset] Received text length={len(str(text))} preview={_preview}...") + + # 步骤2: 解析VLM输出 + obj = _extract_json_obj(str(text)) + if not obj or "dish" not in obj: + print("[VLM_Openset] JSON parse failed or missing 'dish' field") + return { + "raw_dish": str(text), + "cleaned_dish": "", + "vlm_confidence": 0.0, + "candidates": [], + "best_match": "Unknown", + "final_confidence": 0.0, + } + + raw_dish = obj.get("dish", "") + vlm_conf = obj.get("confidence", 0.0) + try: + vlm_conf = float(vlm_conf) + except Exception: + vlm_conf = 0.0 + + print(f"[VLM_Openset] Parsed raw_dish='{raw_dish}' vlm_confidence={vlm_conf}") + + # 步骤3: 清洗菜品名 + cleaned_dish = normalize_dish_name(raw_dish) + print(f"[VLM_Openset] Cleaned dish='{cleaned_dish}'") + + if not cleaned_dish: + print("[VLM_Openset] Cleaned dish is empty, return Unknown") + return { + "raw_dish": raw_dish, + "cleaned_dish": cleaned_dish, + "vlm_confidence": vlm_conf, + "candidates": [], + "best_match": "Unknown", + "final_confidence": 0.0, + } + + # 步骤4: 向量匹配 + print(f"[VLM_Openset] Embedding text='{cleaned_dish}'...") + query_emb = embedder.encode_single(cleaned_dish) + + print(f"[VLM_Openset] Matching with FAISS index...") + matches = matcher.match(query_emb, top_k) + + # 步骤5: 构建候选列表(综合置信度 = VLM置信度 × 匹配分数) + candidates = [] + for dish_name, match_score in matches: + if match_score < min_match_score: + print(f"[VLM_Openset] Skip candidate '{dish_name}' (score={match_score:.4f} < threshold={min_match_score})") + continue + + final_score = vlm_conf * match_score + candidates.append({ + "dish": dish_name, + "match_score": match_score, + "final_score": final_score, + }) + + # 按综合分数降序排序 + candidates.sort(key=lambda x: x["final_score"], reverse=True) + + # 最佳匹配 + if candidates: + best = candidates[0] + best_match = best["dish"] + final_conf = best["final_score"] + print(f"[VLM_Openset] Best match='{best_match}' final_confidence={final_conf:.4f}") + else: + best_match = "Unknown" + final_conf = 0.0 + print("[VLM_Openset] No valid candidates, return Unknown") + + return { + "raw_dish": raw_dish, + "cleaned_dish": cleaned_dish, + "vlm_confidence": vlm_conf, + "candidates": candidates, + "best_match": best_match, + "final_confidence": final_conf, + } diff --git a/train/grid_search_cosface.py b/train/grid_search_cosface.py index 89819a6..aa39199 100644 --- a/train/grid_search_cosface.py +++ b/train/grid_search_cosface.py @@ -354,8 +354,8 @@ def grid_search_main( if __name__ == '__main__': import argparse parser = argparse.ArgumentParser(description='CosFace超参数网格搜索') - # parser.add_argument('--task', choices=list(TASKS.keys()), default='dish', help='任务名称') - parser.add_argument('--task', choices=list(TASKS.keys()), default='whole_ingredient', help='任务名称') + parser.add_argument('--task', choices=list(TASKS.keys()), default='dish', help='任务名称') + # parser.add_argument('--task', choices=list(TASKS.keys()), default='whole_ingredient', help='任务名称') parser.add_argument('--max_configs', type=int, default=None, help='最大配置数(用于测试)') parser.add_argument('--epochs', type=int, default=100, help='每个配置的最大训练轮数') parser.add_argument('--patience', type=int, default=10, help='早停容忍轮数')