From 5cfb84084a8633218292be9d367416b3e622ff7d Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Thu, 13 Nov 2025 13:39:30 +0800 Subject: [PATCH] =?UTF-8?q?=E5=8F=AF=E8=A7=86=E5=8C=96=E6=93=8D=E4=BD=9C?= =?UTF-8?q?=E7=95=8C=E9=9D=A2=E5=8A=A0=E5=85=A5=E5=BC=80=E6=94=BE=E9=9B=86?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- exp_multimodal/exp_multimodal_gui.py | 341 +++++++++++++++++++++++++-- vlm_config.json | 9 +- 2 files changed, 323 insertions(+), 27 deletions(-) diff --git a/exp_multimodal/exp_multimodal_gui.py b/exp_multimodal/exp_multimodal_gui.py index a6bdfe5..2824c54 100644 --- a/exp_multimodal/exp_multimodal_gui.py +++ b/exp_multimodal/exp_multimodal_gui.py @@ -13,9 +13,12 @@ from tkinter import filedialog, messagebox from tkinterdnd2 import DND_FILES, TkinterDnD from exp_multimodal.labels import build_labels, _normalize, _base_ingredient -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 as DEFAULT_OLLAMA_URL, DEFAULT_MODEL as DEFAULT_VLM_MODEL from exp_multimodal.vlm_providers import VLMProvider, OllamaProvider, KimiProvider +from exp_multimodal.text_embedder import OllamaEmbedder +from exp_multimodal.vector_matcher import DishNameMatcher +from exp_multimodal.build_dish_name_index import build_index as build_dish_index ctk.set_appearance_mode("System") @@ -63,6 +66,17 @@ class MultiModalFoodApp: } self.extra_labels_file_path: Optional[str] = None + # 开放式识别配置 + self.openset_enabled_var = ctk.BooleanVar(value=False) + self.openset_index_path_var = ctk.StringVar(value="exp_multimodal/dish_name_index") + self.openset_top_k_var = ctk.IntVar(value=3) + self.openset_min_score_var = ctk.DoubleVar(value=0.5) + self.openset_embedder_model_var = ctk.StringVar(value="quentinz/bge-large-zh-v1.5") + + # 开放式识别运行时对象(延迟初始化) + self.openset_embedder: Optional[OllamaEmbedder] = None + self.openset_matcher: Optional[DishNameMatcher] = None + # UI self.create_widgets() @@ -309,7 +323,7 @@ class MultiModalFoodApp: info.pack(side="left", fill="both", expand=True, padx=10, pady=10) ctk.CTkLabel(info, text=f"文件: {r['image_name']}", anchor="w", font=("Arial", 12, "bold")).pack(fill="x", padx=5, pady=2) - # 仅展示解析后的标签和置信度 + # 显示预测结果 is_correct = r.get('is_correct') color = "green" if is_correct is True else ("red" if is_correct is False else "orange") ctk.CTkLabel(info, text=r['predicted_label'], anchor="w", font=("Arial", 18, "bold"), text_color=color).pack(fill="x", padx=5, pady=(4, 2)) @@ -319,6 +333,37 @@ class MultiModalFoodApp: true_cls = r.get('true_class') if true_cls: ctk.CTkLabel(info, text=f"真实类别: {true_cls}", anchor="w", font=("Arial", 10)).pack(fill="x", padx=5, pady=(0, 2)) + + # 开放式识别额外信息 + if r.get('openset_mode'): + raw_dish = r.get('raw_dish', '') + if raw_dish: + ctk.CTkLabel(info, text=f"VLM原始输出: {raw_dish}", anchor="w", font=("Arial", 9), text_color="gray").pack(fill="x", padx=5, pady=(0, 2)) + + # 候选列表(可展开) + candidates = r.get('candidates', []) + if candidates: + # 创建可展开的候选框 + candidates_frame = ctk.CTkFrame(info) + candidates_frame.pack(fill="x", padx=5, pady=(4, 2)) + + def toggle_candidates(frame=candidates_frame, cands=candidates): + # 切换显示/隐藏 + if len(frame.winfo_children()) > 1: + # 隐藏 + for w in frame.winfo_children()[1:]: + w.destroy() + else: + # 显示 + for idx, cand in enumerate(cands, 1): + dish_name = cand.get('dish', 'Unknown') + match_score = cand.get('match_score', 0.0) + final_score = cand.get('final_score', 0.0) + cand_text = f" [{idx}] {dish_name} (匹配:{match_score:.3f} 综合:{final_score:.3f})" + ctk.CTkLabel(frame, text=cand_text, anchor="w", font=("Arial", 9), text_color="gray").pack(fill="x", padx=5, pady=1) + + toggle_btn = ctk.CTkButton(candidates_frame, text=f"▶ 查看候选列表 ({len(candidates)}个)", command=toggle_candidates, width=180, height=24, fg_color="gray", hover_color="darkgray") + toggle_btn.pack(anchor="w", padx=5, pady=2) # -------------------- 右侧-配置与运行 -------------------- def create_config_tab(self): @@ -439,6 +484,44 @@ class MultiModalFoodApp: self.refresh_extra_labels_view() self.refresh_fewshot_view() + # ==================== 开放式识别配置 ==================== + openset_frame = ctk.CTkFrame(panel) + openset_frame.pack(fill="x", pady=(10, 10)) + + top = ctk.CTkFrame(openset_frame) + top.pack(fill="x") + ctk.CTkLabel(top, text="开放式识别配置", font=("Arial", 14, "bold")).pack(side="left", padx=0, pady=(0, 6)) + ctk.CTkSwitch(top, text="启用(仅dish模式)", variable=self.openset_enabled_var, command=self.on_openset_toggle).pack(side="left", padx=12) + + # 索引路径 + openset_row1 = ctk.CTkFrame(openset_frame) + openset_row1.pack(fill="x", pady=(4, 4)) + ctk.CTkLabel(openset_row1, text="向量索引路径:").pack(side="left", padx=6, pady=4) + ctk.CTkEntry(openset_row1, textvariable=self.openset_index_path_var, width=320).pack(side="left", padx=4, pady=4) + ctk.CTkButton(openset_row1, text="浏览", command=self.browse_openset_index, width=80).pack(side="left", padx=4) + + # Top-K 和 最低分数 + openset_row2 = ctk.CTkFrame(openset_frame) + openset_row2.pack(fill="x", pady=(0, 4)) + ctk.CTkLabel(openset_row2, text="Top-K候选数:").pack(side="left", padx=6, pady=4) + ctk.CTkEntry(openset_row2, textvariable=self.openset_top_k_var, width=80).pack(side="left", padx=4, pady=4) + ctk.CTkLabel(openset_row2, text="最低匹配分数:").pack(side="left", padx=(20, 6), pady=4) + ctk.CTkEntry(openset_row2, textvariable=self.openset_min_score_var, width=80).pack(side="left", padx=4, pady=4) + + # Embedding 模型 + openset_row3 = ctk.CTkFrame(openset_frame) + openset_row3.pack(fill="x", pady=(0, 4)) + ctk.CTkLabel(openset_row3, text="Embedding模型:").pack(side="left", padx=6, pady=4) + ctk.CTkEntry(openset_row3, textvariable=self.openset_embedder_model_var, width=380).pack(side="left", padx=4, pady=4) + + # 索引管理按钮 + openset_row4 = ctk.CTkFrame(openset_frame) + openset_row4.pack(fill="x", pady=(4, 6)) + ctk.CTkButton(openset_row4, text="构建索引", command=self.build_openset_index, width=110, fg_color="green", hover_color="darkgreen").pack(side="left", padx=6) + ctk.CTkButton(openset_row4, text="测试索引", command=self.test_openset_index, width=110).pack(side="left", padx=6) + self.openset_status_label = ctk.CTkLabel(openset_row4, text="", font=("Arial", 10), text_color="gray") + self.openset_status_label.pack(side="left", padx=12) + # 配置保存/加载 cfg_manage = ctk.CTkFrame(panel) cfg_manage.pack(fill="x", pady=(10, 10)) @@ -602,6 +685,96 @@ class MultiModalFoodApp: except Exception as e: messagebox.showerror("错误", f"保存失败: {e}") + # -------------------- 开放式识别配置管理 -------------------- + def on_openset_toggle(self): + """开放式识别开关切换时的处理""" + if self.openset_enabled_var.get(): + # 提示用户该功能仅适用于 dish 模式 + if self.mode_var.get() != "dish": + messagebox.showwarning("提示", "开放式识别当前仅支持 dish 模式,请先切换模式") + self.openset_enabled_var.set(False) + return + + # 检查索引是否存在 + index_path = self.openset_index_path_var.get() + if not os.path.exists(index_path): + if messagebox.askyesno("索引不存在", f"索引目录 {index_path} 不存在。是否现在构建索引?"): + self.build_openset_index() + else: + self.openset_enabled_var.set(False) + + def browse_openset_index(self): + """浏览选择索引目录""" + path = filedialog.askdirectory(title="选择向量索引目录") + if path: + self.openset_index_path_var.set(path) + + def build_openset_index(self): + """构建向量索引(后台线程)""" + if messagebox.askyesno("确认", "构建索引可能需要几分钟时间,是否继续?"): + self.openset_status_label.configure(text="正在构建索引...", text_color="orange") + threading.Thread(target=self._build_index_thread, daemon=True).start() + + def _build_index_thread(self): + """索引构建后台线程""" + try: + mode = "dish" # 开放式识别目前仅支持 dish + labels = build_labels(mode, alias_map_path=None) + + if not labels: + self.root.after(0, lambda: messagebox.showerror("错误", "未找到可用的菜品名")) + self.root.after(0, lambda: self.openset_status_label.configure(text="构建失败", text_color="red")) + return + + # 初始化 Embedder + embedder_url = self.ollama_url_var.get().strip() or DEFAULT_OLLAMA_URL + embedder_model = self.openset_embedder_model_var.get().strip() + embedder = OllamaEmbedder(base_url=embedder_url, model=embedder_model) + + # 构建索引 + output_dir = self.openset_index_path_var.get() + build_dish_index( + dish_names=labels, + embedder=embedder, + output_dir=output_dir, + batch_size=100 + ) + + self.root.after(0, lambda: messagebox.showinfo("成功", f"索引构建完成!位置: {output_dir}")) + self.root.after(0, lambda: self.openset_status_label.configure(text=f"索引已构建 ({len(labels)}个菜品)", text_color="green")) + + except Exception as e: + self.root.after(0, lambda: messagebox.showerror("错误", f"构建索引失败: {e}")) + self.root.after(0, lambda: self.openset_status_label.configure(text="构建失败", text_color="red")) + + def test_openset_index(self): + """测试索引是否可用""" + try: + index_path = self.openset_index_path_var.get() + + # 检查必要文件是否存在 + required_files = ["dish_names.json", "dish_embeddings.npy", "faiss_index.bin"] + missing = [] + for fname in required_files: + if not os.path.exists(os.path.join(index_path, fname)): + missing.append(fname) + + if missing: + messagebox.showerror("错误", f"索引目录不完整,缺少文件:{', '.join(missing)}") + self.openset_status_label.configure(text="索引无效", text_color="red") + return + + # 尝试加载索引 + matcher = DishNameMatcher(index_dir=index_path) + num_dishes = len(matcher.dish_names) + + messagebox.showinfo("成功", f"索引测试通过!包含 {num_dishes} 个菜品名") + self.openset_status_label.configure(text=f"索引正常 ({num_dishes}个菜品)", text_color="green") + + except Exception as e: + messagebox.showerror("错误", f"索引测试失败: {e}") + self.openset_status_label.configure(text="索引无效", text_color="red") + # -------------------- 识别流程 -------------------- def start_recognition(self): if not self.uploaded_images: @@ -633,50 +806,138 @@ class MultiModalFoodApp: def recognize_images(self): try: mode = self.mode_var.get() - labels = self._build_final_labels(mode) - ingredient_only = (mode in {"whole", "processed"}) - use_hints = self.fewshot_hints if self.fewshot_enabled_var.get() else None - + use_openset = self.openset_enabled_var.get() and mode == "dish" + # 创建 VLM Provider provider = self._create_vlm_provider() if provider is None: self.root.after(0, lambda: messagebox.showerror("错误", "Provider 配置错误,请检查配置")) self.root.after(0, self.recognition_completed) return + + # 开放式识别分支 + if use_openset: + self._recognize_with_openset(provider) + # 封闭式识别分支(原有逻辑) + else: + self._recognize_with_closedset(provider, mode) + + self.root.after(0, self.recognition_completed) + except Exception as e: + self.root.after(0, lambda: messagebox.showerror("错误", f"识别过程中出错: {e}")) + self.root.after(0, self.recognition_completed) + + def _recognize_with_closedset(self, provider: VLMProvider, mode: str): + """封闭式识别(原有逻辑)""" + labels = self._build_final_labels(mode) + ingredient_only = (mode in {"whole", "processed"}) + use_hints = self.fewshot_hints if self.fewshot_enabled_var.get() else None + self.current_results.clear() + for i, img in enumerate(self.uploaded_images): + try: + result = classify_image( + image_path=img['path'], + labels=labels, + provider=provider, + fewshot_hints=use_hints, + ingredient_only=ingredient_only, + ) + pred = result.get("label", "Unknown") + conf = result.get("confidence", 0.0) + # 若非数值,后续展示为 N/A + if not isinstance(conf, (int, float)): + try: + conf = float(conf) + except Exception: + pass + true_cls = self._infer_true_class(img['path'], labels) + is_correct = (pred == true_cls) if true_cls is not None else None + ui_res = { + "image_index": i, + "image_name": img['name'], + "predicted_label": pred, + "confidence": conf, + "true_class": true_cls, + "is_correct": is_correct, + "openset_mode": False, + } + self.current_results.append(ui_res) + img['recognized'] = True + img['result'] = ui_res + except Exception as e: + print(f"[VLM] classify error: {e}") + ui_res = { + "image_index": i, + "image_name": img['name'], + "predicted_label": "Error", + "confidence": "N/A", + "true_class": None, + "is_correct": None, + "openset_mode": False, + } + self.current_results.append(ui_res) + img['recognized'] = True + img['result'] = ui_res + finally: + self.root.after(0, self.update_progress, i + 1, len(self.uploaded_images)) + + def _recognize_with_openset(self, provider: VLMProvider): + """开放式识别(新逻辑)""" + try: + # 初始化 Embedder 和 Matcher(延迟加载) + if self.openset_embedder is None: + embedder_url = self.ollama_url_var.get().strip() or DEFAULT_OLLAMA_URL + embedder_model = self.openset_embedder_model_var.get().strip() + self.openset_embedder = OllamaEmbedder(base_url=embedder_url, model=embedder_model) + + if self.openset_matcher is None: + index_path = self.openset_index_path_var.get() + self.openset_matcher = DishNameMatcher(index_dir=index_path) + + top_k = self.openset_top_k_var.get() + min_score = self.openset_min_score_var.get() + + # 获取所有菜品名(用于准确率计算) + all_dish_names = self.openset_matcher.dish_names + self.current_results.clear() for i, img in enumerate(self.uploaded_images): try: - result = classify_image( + result = classify_image_openset( image_path=img['path'], - labels=labels, provider=provider, - fewshot_hints=use_hints, - ingredient_only=ingredient_only, + embedder=self.openset_embedder, + matcher=self.openset_matcher, + top_k=top_k, + min_match_score=min_score, ) - pred = result.get("label", "Unknown") - conf = result.get("confidence", 0.0) - # 若非数值,后续展示为 N/A - if not isinstance(conf, (int, float)): - try: - conf = float(conf) - except Exception: - pass - true_cls = self._infer_true_class(img['path'], labels) - is_correct = (pred == true_cls) if true_cls is not None else None + + # 解析开放式结果 + best_match = result.get("best_match", "Unknown") + final_conf = result.get("final_confidence", 0.0) + candidates = result.get("candidates", []) + raw_dish = result.get("raw_dish", "") + + true_cls = self._infer_true_class(img['path'], all_dish_names) + is_correct = (best_match == true_cls) if true_cls is not None else None + ui_res = { "image_index": i, "image_name": img['name'], - "predicted_label": pred, - "confidence": conf, + "predicted_label": best_match, + "confidence": final_conf, "true_class": true_cls, "is_correct": is_correct, + "openset_mode": True, + "raw_dish": raw_dish, + "candidates": candidates, } self.current_results.append(ui_res) img['recognized'] = True img['result'] = ui_res except Exception as e: - print(f"[VLM] classify error: {e}") + print(f"[VLM-Openset] classify error: {e}") ui_res = { "image_index": i, "image_name": img['name'], @@ -684,16 +945,18 @@ class MultiModalFoodApp: "confidence": "N/A", "true_class": None, "is_correct": None, + "openset_mode": True, + "raw_dish": "", + "candidates": [], } self.current_results.append(ui_res) img['recognized'] = True img['result'] = ui_res finally: self.root.after(0, self.update_progress, i + 1, len(self.uploaded_images)) - self.root.after(0, self.recognition_completed) + except Exception as e: - self.root.after(0, lambda: messagebox.showerror("错误", f"识别过程中出错: {e}")) - self.root.after(0, self.recognition_completed) + raise Exception(f"开放式识别初始化失败: {e}") def _create_vlm_provider(self) -> Optional[VLMProvider]: """根据配置创建 VLM Provider""" @@ -759,6 +1022,14 @@ class MultiModalFoodApp: "fewshot_file_path": self.fewshot_file_path, "extra_labels": self.extra_labels, "extra_labels_file_path": self.extra_labels_file_path, + # 新增:开放式识别配置 + "openset": { + "enabled": self.openset_enabled_var.get(), + "index_path": self.openset_index_path_var.get(), + "top_k": self.openset_top_k_var.get(), + "min_score": self.openset_min_score_var.get(), + "embedder_model": self.openset_embedder_model_var.get(), + }, } try: @@ -810,6 +1081,15 @@ class MultiModalFoodApp: if "extra_labels_file_path" in config: self.extra_labels_file_path = config["extra_labels_file_path"] + # 新增:开放式识别配置 + if "openset" in config: + openset_cfg = config["openset"] + self.openset_enabled_var.set(openset_cfg.get("enabled", False)) + self.openset_index_path_var.set(openset_cfg.get("index_path", "exp_multimodal/dish_name_index")) + self.openset_top_k_var.set(openset_cfg.get("top_k", 3)) + self.openset_min_score_var.set(openset_cfg.get("min_score", 0.5)) + self.openset_embedder_model_var.set(openset_cfg.get("embedder_model", "quentinz/bge-large-zh-v1.5")) + self.refresh_extra_labels_view() self.refresh_fewshot_view() self.on_provider_change() @@ -862,6 +1142,15 @@ class MultiModalFoodApp: if "extra_labels_file_path" in config: self.extra_labels_file_path = config["extra_labels_file_path"] + # 新增:开放式识别配置 + if "openset" in config: + openset_cfg = config["openset"] + self.openset_enabled_var.set(openset_cfg.get("enabled", False)) + self.openset_index_path_var.set(openset_cfg.get("index_path", "exp_multimodal/dish_name_index")) + self.openset_top_k_var.set(openset_cfg.get("top_k", 3)) + self.openset_min_score_var.set(openset_cfg.get("min_score", 0.5)) + self.openset_embedder_model_var.set(openset_cfg.get("embedder_model", "quentinz/bge-large-zh-v1.5")) + self.refresh_extra_labels_view() self.refresh_fewshot_view() self.on_provider_change() diff --git a/vlm_config.json b/vlm_config.json index e8b6fae..1bbc71e 100644 --- a/vlm_config.json +++ b/vlm_config.json @@ -11023,5 +11023,12 @@ "whole": [], "processed": [] }, - "extra_labels_file_path": "C:/Users/Huawei/Desktop/labels.json" + "extra_labels_file_path": "C:/Users/Huawei/Desktop/labels.json", + "openset": { + "enabled": true, + "index_path": "exp_multimodal/dish_name_index", + "top_k": 3, + "min_score": 0.5, + "embedder_model": "quentinz/bge-large-zh-v1.5" + } } \ No newline at end of file