import os import json import time import threading from typing import Dict, List, Optional import cv2 import numpy as np from PIL import Image import customtkinter as ctk 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, 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") ctk.set_default_color_theme("blue") class MultiModalFoodApp: def __init__(self, root): self.root = root self.root.title("数字味道-食物识别系统 (多模态版)") self.root.geometry("1400x800") # 数据 self.uploaded_images: List[Dict] = [] self.current_results: List[Dict] = [] # 运行统计 self.recognition_start_time: Optional[float] = None self.recognition_duration: float = 0.0 # 配置 self.mode_var = ctk.StringVar(value="dish") # dish | whole | processed # VLM Provider 配置 self.provider_var = ctk.StringVar(value="ollama") # ollama | kimi self.ollama_url_var = ctk.StringVar(value=DEFAULT_OLLAMA_URL) self.vlm_model_var = ctk.StringVar(value=DEFAULT_VLM_MODEL) self.kimi_api_key_var = ctk.StringVar(value="") self.kimi_base_url_var = ctk.StringVar(value="https://api.moonshot.cn/v1") self.kimi_model_var = ctk.StringVar(value="moonshot-v1-32k-vision-preview") self.alias_map_path: Optional[str] = None # Fewshot 可视化编辑(C 方案): dict[label] = hint self.fewshot_hints: Dict[str, str] = {} self.fewshot_file_path: Optional[str] = None self.fewshot_enabled_var = ctk.BooleanVar(value=False) # 新增可识别类别(可保存到本地 JSON) # 为不同模式维持独立的额外标签列表 self.extra_labels: Dict[str, List[str]] = { "dish": [], "whole": [], "processed": [], } 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() # 加载保存的配置 self.load_config() # -------------------- 图像加载工具 -------------------- def load_image_with_chinese_path(self, file_path: str): try: with open(file_path, 'rb') as f: data = f.read() nparr = np.frombuffer(data, np.uint8) image = cv2.imdecode(nparr, cv2.IMREAD_COLOR) return image except Exception as e: print(f"加载图片失败: {e}") return None def resize_image_for_display(self, image, max_w, max_h): h, w = image.shape[:2] scale = min(max_w / w, max_h / h) if scale < 1: new_w, new_h = int(w * scale), int(h * scale) return cv2.resize(image, (new_w, new_h)) return image # -------------------- 主布局 -------------------- def create_widgets(self): self.main_frame = ctk.CTkFrame(self.root) self.main_frame.pack(fill="both", expand=True, padx=15, pady=15) # 左侧:上传/管理 self.left_frame = ctk.CTkFrame(self.main_frame, width=600) self.left_frame.pack(side="left", fill="both", expand=True, padx=(0, 10)) self.left_frame.pack_propagate(False) self.left_title = ctk.CTkLabel(self.left_frame, text="图片上传区域 (多模态识别)", font=("Arial", 16, "bold")) self.left_title.pack(pady=(15, 10)) self.upload_frame = ctk.CTkFrame(self.left_frame, fg_color=("gray90", "gray20")) self.upload_frame.pack(fill="x", padx=15, pady=(0, 10), ipady=50) self.upload_label = ctk.CTkLabel( self.upload_frame, text="拖拽图片到这里\n或点击下方按钮选择图片\n支持多图片上传", font=("Arial", 14), text_color=("gray40", "gray60"), ) self.upload_label.pack(expand=True) self.upload_frame.drop_target_register(DND_FILES) self.upload_frame.dnd_bind('<>', self.handle_drop) self.upload_frame.bind('', self.on_drag_enter) self.upload_frame.bind('', self.on_drag_leave) self.button_frame = ctk.CTkFrame(self.left_frame) self.button_frame.pack(fill="x", padx=15, pady=(0, 10)) self.select_button = ctk.CTkButton(self.button_frame, text="选择图片", command=self.select_images, width=120, height=35) self.select_button.pack(side="left", padx=(10, 5), pady=10) self.clear_button = ctk.CTkButton(self.button_frame, text="清空图片", command=self.clear_images, width=120, height=35, fg_color="gray", hover_color="darkgray") self.clear_button.pack(side="left", padx=5, pady=10) self.recognize_button = ctk.CTkButton(self.button_frame, text="开始识别", command=self.start_recognition, width=120, height=35, fg_color="green", hover_color="darkgreen") self.recognize_button.pack(side="right", padx=(5, 10), pady=10) self.recognize_button.configure(state="disabled") self.images_display_frame = ctk.CTkScrollableFrame(self.left_frame, label_text="已上传的图片") self.images_display_frame.pack(fill="both", expand=True, padx=15, pady=(0, 15)) # 右侧:结果与配置 self.right_frame = ctk.CTkFrame(self.main_frame, width=700) self.right_frame.pack(side="right", fill="both", expand=True, padx=(10, 0)) self.right_frame.pack_propagate(False) self.right_tabview = ctk.CTkTabview(self.right_frame) self.right_tabview.pack(fill="both", expand=True, padx=15, pady=15) self.results_tab = self.right_tabview.add("识别结果") self.config_tab = self.right_tabview.add("配置与运行") self.right_tabview.set("识别结果") self.create_results_tab() self.create_config_tab() # -------------------- 左侧交互 -------------------- def select_images(self): paths = filedialog.askopenfilenames( title="选择图片文件", filetypes=[("图像文件", "*.jpg *.jpeg *.png *.bmp *.gif"), ("所有文件", "*.*")], ) if paths: for p in paths: self.add_image(p) def handle_drop(self, event): files = event.data.split() for p in files: p = p.strip('{}').strip('"') p = os.path.normpath(p) if p.lower().endswith((".jpg", ".jpeg", ".png", ".bmp", ".gif")): self.add_image(p) def on_drag_enter(self, _): self.upload_frame.configure(fg_color=("gray80", "gray30")) self.upload_label.configure(text="释放鼠标上传图片") def on_drag_leave(self, _): self.upload_frame.configure(fg_color=("gray90", "gray20")) self.upload_label.configure(text="拖拽图片到这里\n或点击下方按钮选择图片\n支持多图片上传") def add_image(self, file_path: str): try: if not os.path.exists(file_path): messagebox.showerror("错误", f"文件不存在: {file_path}") return if any(img['path'] == file_path for img in self.uploaded_images): messagebox.showinfo("提示", "该图片已经添加") return image = self.load_image_with_chinese_path(file_path) if image is None: messagebox.showerror("错误", f"无法读取图片: {file_path}") return info = { "path": file_path, "name": os.path.basename(file_path), "image": image, "recognized": False, "result": None, } self.uploaded_images.append(info) self.update_images_display() self.update_recognize_button() except Exception as e: messagebox.showerror("错误", f"添加图片时出错: {e}") def clear_images(self): if self.uploaded_images: if messagebox.askyesno("确认", "确定清空所有图片吗?"): self.uploaded_images.clear() self.current_results.clear() self.recognition_start_time = None self.recognition_duration = 0 self.update_images_display() self.update_recognize_button() self.update_results_display() self.update_stats() def remove_image(self, index: int): if 0 <= index < len(self.uploaded_images): self.uploaded_images.pop(index) self.update_images_display() self.update_recognize_button() self.update_results_display() def preview_image(self, index: int): if not (0 <= index < len(self.uploaded_images)): return img_info = self.uploaded_images[index] win = ctk.CTkToplevel(self.root) win.title(f"预览 - {img_info['name']}") win.geometry("800x600") win.transient(self.root) win.grab_set() win.lift() win.focus_set() win.update_idletasks() x = (win.winfo_screenwidth() // 2) - (800 // 2) y = (win.winfo_screenheight() // 2) - (600 // 2) win.geometry(f"800x600+{x}+{y}") display = self.resize_image_for_display(img_info['image'], 750, 550) display = cv2.cvtColor(display, cv2.COLOR_BGR2RGB) pil = Image.fromarray(display) w, h = pil.size tkimg = ctk.CTkImage(light_image=pil, dark_image=pil, size=(w, h)) lbl = ctk.CTkLabel(win, image=tkimg, text="") lbl.image = tkimg lbl.pack(expand=True, padx=20, pady=20) def update_images_display(self): for w in self.images_display_frame.winfo_children(): w.destroy() for i, img in enumerate(self.uploaded_images): row = ctk.CTkFrame(self.images_display_frame) row.pack(fill="x", padx=5, pady=5) disp = self.resize_image_for_display(img['image'], 100, 100) disp = cv2.cvtColor(disp, cv2.COLOR_BGR2RGB) pil = Image.fromarray(disp) tkimg = ctk.CTkImage(light_image=pil, dark_image=pil, size=(100, 100)) img_label = ctk.CTkLabel(row, image=tkimg, text="") img_label.image = tkimg img_label.pack(side="left", padx=10, pady=10) img_label.bind("", lambda e, idx=i: self.preview_image(idx)) info = ctk.CTkFrame(row) info.pack(side="left", fill="both", expand=True, padx=10, pady=10) ctk.CTkLabel(info, text=f"文件名: {img['name']}", anchor="w").pack(fill="x", padx=5, pady=2) status = "已识别" if img['recognized'] else "未识别" color = "green" if img['recognized'] else None ctk.CTkLabel(info, text=f"状态: {status}", anchor="w", text_color=color).pack(fill="x", padx=5, pady=2) del_btn = ctk.CTkButton(row, text="删除", width=60, height=30, fg_color="red", hover_color="darkred", command=lambda idx=i: self.remove_image(idx)) del_btn.pack(side="right", padx=10, pady=10) def update_recognize_button(self): self.recognize_button.configure(state=("normal" if self.uploaded_images else "disabled")) # -------------------- 右侧-结果 -------------------- def create_results_tab(self): self.stats_frame = ctk.CTkFrame(self.results_tab) self.stats_frame.pack(fill="x", padx=15, pady=(10, 10)) self.stats_label = ctk.CTkLabel(self.stats_frame, text="总图片: 0 | 已识别: 0 | 平均准确率: 0% | 耗时: 0.00s", font=("Arial", 12)) self.stats_label.pack(pady=10) self.results_display_frame = ctk.CTkScrollableFrame(self.results_tab, label_text="识别详情(多模态)") self.results_display_frame.pack(fill="both", expand=True, padx=15, pady=(0, 15)) def update_stats(self): total = len(self.uploaded_images) done = sum(1 for x in self.uploaded_images if x['recognized']) # 使用可推断的真实类别:从图片父目录名与预测对比(若能匹配到 labels) correct = sum(1 for r in self.current_results if r.get('is_correct') is True) acc = (correct / len(self.current_results) * 100.0) if self.current_results else 0.0 self.stats_label.configure(text=f"总图片: {total} | 已识别: {done} | 平均准确率: {acc:.1f}% | 耗时: {self.recognition_duration:.2f}s") def update_results_display(self): for w in self.results_display_frame.winfo_children(): w.destroy() if not self.current_results: ctk.CTkLabel(self.results_display_frame, text="暂无识别结果", font=("Arial", 14), text_color="gray").pack(pady=20) return for r in self.current_results: row = ctk.CTkFrame(self.results_display_frame) row.pack(fill="x", padx=5, pady=5) img = self.uploaded_images[r['image_index']]['image'] disp = self.resize_image_for_display(img, 120, 120) disp = cv2.cvtColor(disp, cv2.COLOR_BGR2RGB) pil = Image.fromarray(disp) tkimg = ctk.CTkImage(light_image=pil, dark_image=pil, size=(120, 120)) img_label = ctk.CTkLabel(row, image=tkimg, text="") img_label.image = tkimg img_label.pack(side="left", padx=10, pady=10) img_label.bind("", lambda e, idx=r['image_index']: self.preview_image(idx)) info = ctk.CTkFrame(row) 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)) conf = r.get('confidence') conf_str = (f"{conf:.3f}" if isinstance(conf, (int, float)) else "N/A") ctk.CTkLabel(info, text=f"置信度: {conf_str}", anchor="w").pack(fill="x", padx=5, pady=(0, 2)) 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): panel = ctk.CTkScrollableFrame(self.config_tab) panel.pack(fill="both", expand=True, padx=15, pady=15) # 模式 ctk.CTkLabel(panel, text="识别模式", font=("Arial", 14, "bold")).pack(anchor="w", pady=(0, 6)) mode_frame = ctk.CTkFrame(panel) mode_frame.pack(fill="x", pady=(0, 10)) for val, text in [("dish", "菜品(dish)"), ("whole", "整食材(whole)"), ("processed", "处理后食材(processed)")]: rb = ctk.CTkRadioButton(mode_frame, text=text, variable=self.mode_var, value=val) rb.pack(side="left", padx=8, pady=8) # VLM 服务提供商选择 ctk.CTkLabel(panel, text="VLM 服务提供商", font=("Arial", 14, "bold")).pack(anchor="w", pady=(10, 6)) provider_frame = ctk.CTkFrame(panel) provider_frame.pack(fill="x", pady=(0, 10)) for val, text in [("ollama", "Ollama (本地/自建)"), ("kimi", "Kimi 1.5 (厂商API)")]: rb = ctk.CTkRadioButton( provider_frame, text=text, variable=self.provider_var, value=val, command=self.on_provider_change ) rb.pack(side="left", padx=8, pady=8) # Ollama 配置区域 self.ollama_config_frame = ctk.CTkFrame(panel) self.ollama_config_frame.pack(fill="x", pady=(0, 10)) ctk.CTkLabel(self.ollama_config_frame, text="Ollama 配置", font=("Arial", 12, "bold")).pack(anchor="w", pady=(4, 4)) ollama_row1 = ctk.CTkFrame(self.ollama_config_frame) ollama_row1.pack(fill="x", pady=(0, 4)) ctk.CTkLabel(ollama_row1, text="OLLAMA_URL:").pack(side="left", padx=6, pady=4) ctk.CTkEntry(ollama_row1, textvariable=self.ollama_url_var, width=420).pack(side="left", padx=4, pady=4) ollama_row2 = ctk.CTkFrame(self.ollama_config_frame) ollama_row2.pack(fill="x", pady=(0, 4)) ctk.CTkLabel(ollama_row2, text="MODEL:").pack(side="left", padx=6, pady=4) ctk.CTkEntry(ollama_row2, textvariable=self.vlm_model_var, width=420).pack(side="left", padx=4, pady=4) # Kimi 配置区域 self.kimi_config_frame = ctk.CTkFrame(panel) self.kimi_config_frame.pack(fill="x", pady=(0, 10)) ctk.CTkLabel(self.kimi_config_frame, text="Kimi 配置", font=("Arial", 12, "bold")).pack(anchor="w", pady=(4, 4)) kimi_row1 = ctk.CTkFrame(self.kimi_config_frame) kimi_row1.pack(fill="x", pady=(0, 4)) ctk.CTkLabel(kimi_row1, text="API_KEY:").pack(side="left", padx=6, pady=4) ctk.CTkEntry(kimi_row1, textvariable=self.kimi_api_key_var, width=420, show="*").pack(side="left", padx=4, pady=4) kimi_row2 = ctk.CTkFrame(self.kimi_config_frame) kimi_row2.pack(fill="x", pady=(0, 4)) ctk.CTkLabel(kimi_row2, text="BASE_URL:").pack(side="left", padx=6, pady=4) ctk.CTkEntry(kimi_row2, textvariable=self.kimi_base_url_var, width=420).pack(side="left", padx=4, pady=4) kimi_row3 = ctk.CTkFrame(self.kimi_config_frame) kimi_row3.pack(fill="x", pady=(0, 4)) ctk.CTkLabel(kimi_row3, text="MODEL:").pack(side="left", padx=6, pady=4) ctk.CTkEntry(kimi_row3, textvariable=self.kimi_model_var, width=420).pack(side="left", padx=4, pady=4) # 初始化显示状态 self.on_provider_change() # Alias Map alias = ctk.CTkFrame(panel) alias.pack(fill="x", pady=(10, 10)) ctk.CTkLabel(alias, text="Alias 映射(JSON,可选)", font=("Arial", 14, "bold")).pack(anchor="w", pady=(0, 6)) alias_row = ctk.CTkFrame(alias) alias_row.pack(fill="x") self.alias_label_var = ctk.StringVar(value="未选择") ctk.CTkLabel(alias_row, textvariable=self.alias_label_var).pack(side="left", padx=6) ctk.CTkButton(alias_row, text="选择文件", command=self.pick_alias_file, width=100).pack(side="left", padx=8) ctk.CTkButton(alias_row, text="清除", command=self.clear_alias_file, width=80, fg_color="gray", hover_color="darkgray").pack(side="left", padx=4) # 新增可识别类别(当前模式) ext = ctk.CTkFrame(panel) ext.pack(fill="x", pady=(10, 10)) ctk.CTkLabel(ext, text="新增可识别类别(当前模式)", font=("Arial", 14, "bold")).pack(anchor="w", pady=(0, 6)) ext_row = ctk.CTkFrame(ext) ext_row.pack(fill="x", pady=(0, 6)) self.new_label_var = ctk.StringVar(value="") ctk.CTkEntry(ext_row, textvariable=self.new_label_var, placeholder_text="输入新类别名", width=260).pack(side="left", padx=6) ctk.CTkButton(ext_row, text="添加", command=self.add_extra_label, width=80).pack(side="left", padx=6) ctk.CTkButton(ext_row, text="删除选中", command=self.remove_selected_extra_label, width=100, fg_color="red", hover_color="darkred").pack(side="left", padx=6) self.extra_labels_listbox = ctk.CTkTextbox(ext, width=520, height=120) self.extra_labels_listbox.pack(fill="x", padx=6, pady=(4, 6)) extra_row2 = ctk.CTkFrame(ext) extra_row2.pack(fill="x") ctk.CTkButton(extra_row2, text="从JSON加载", command=self.load_extra_labels_json, width=110).pack(side="left", padx=6) ctk.CTkButton(extra_row2, text="保存到JSON", command=self.save_extra_labels_json, width=110).pack(side="left", padx=6) # Fewshot 提示(仅在 dish 有明显意义,但允许各模式使用,开启与否由开关控制) fs = ctk.CTkFrame(panel) fs.pack(fill="x", pady=(10, 10)) top = ctk.CTkFrame(fs) top.pack(fill="x") ctk.CTkLabel(top, text="Fewshot 提示(可视化编辑)", font=("Arial", 14, "bold")).pack(side="left", padx=0, pady=(0, 6)) ctk.CTkSwitch(top, text="启用", variable=self.fewshot_enabled_var).pack(side="left", padx=12) fs_row = ctk.CTkFrame(fs) fs_row.pack(fill="x", pady=(4, 6)) self.fs_label_var = ctk.StringVar(value="") self.fs_hint_var = ctk.StringVar(value="") ctk.CTkEntry(fs_row, textvariable=self.fs_label_var, placeholder_text="类别名", width=160).pack(side="left", padx=6) ctk.CTkEntry(fs_row, textvariable=self.fs_hint_var, placeholder_text="提示文本", width=320).pack(side="left", padx=6) ctk.CTkButton(fs_row, text="添加/更新", command=self.add_or_update_fewshot, width=100).pack(side="left", padx=6) ctk.CTkButton(fs_row, text="删除选中", command=self.remove_selected_fewshot, width=100, fg_color="red", hover_color="darkred").pack(side="left", padx=6) self.fewshot_text = ctk.CTkTextbox(fs, width=520, height=160) self.fewshot_text.pack(fill="x", padx=6, pady=(4, 6)) fs_row2 = ctk.CTkFrame(fs) fs_row2.pack(fill="x") ctk.CTkButton(fs_row2, text="从JSON加载", command=self.load_fewshot_json, width=110).pack(side="left", padx=6) ctk.CTkButton(fs_row2, text="保存到JSON", command=self.save_fewshot_json, width=110).pack(side="left", padx=6) 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)) ctk.CTkLabel(cfg_manage, text="配置管理", font=("Arial", 14, "bold")).pack(anchor="w", pady=(0, 6)) cfg_row = ctk.CTkFrame(cfg_manage) cfg_row.pack(fill="x") ctk.CTkButton(cfg_row, text="保存当前配置", command=self.save_config, width=140, fg_color="blue", hover_color="darkblue").pack(side="left", padx=6) ctk.CTkButton(cfg_row, text="加载配置", command=self.load_config_from_file, width=140).pack(side="left", padx=6) ctk.CTkLabel(cfg_row, text="(自动保存到 config.json)", font=("Arial", 10), text_color="gray").pack(side="left", padx=12) def pick_alias_file(self): path = filedialog.askopenfilename(title="选择 alias_map.json", filetypes=[("JSON 文件", "*.json"), ("所有文件", "*.*")]) if path: self.alias_map_path = path self.alias_label_var.set(os.path.basename(path)) def clear_alias_file(self): self.alias_map_path = None self.alias_label_var.set("未选择") def on_provider_change(self): """切换 Provider 时显示/隐藏对应的配置区域""" provider = self.provider_var.get() if provider == "ollama": self.ollama_config_frame.pack(fill="x", pady=(0, 10), before=self.kimi_config_frame) self.kimi_config_frame.pack_forget() elif provider == "kimi": self.kimi_config_frame.pack(fill="x", pady=(0, 10), before=self.ollama_config_frame) self.ollama_config_frame.pack_forget() # 额外标签编辑 def refresh_extra_labels_view(self): self.extra_labels_listbox.configure(state="normal") self.extra_labels_listbox.delete("1.0", "end") mode = self.mode_var.get() for s in self.extra_labels.get(mode, []): self.extra_labels_listbox.insert("end", s + "\n") self.extra_labels_listbox.configure(state="disabled") def add_extra_label(self): raw = self.new_label_var.get().strip() if not raw: return mode = self.mode_var.get() # 与 labels.py 一致的归一化 n = _normalize(raw) if mode in ("whole", "processed"): n = _base_ingredient(n) if not n: return lst = self.extra_labels.setdefault(mode, []) if n not in lst: lst.append(n) self.new_label_var.set("") self.refresh_extra_labels_view() def remove_selected_extra_label(self): try: # 通过选中文本的行来删除 sel = self.extra_labels_listbox.get("sel.first", "sel.last").strip() except Exception: sel = "" if not sel: return mode = self.mode_var.get() if sel in self.extra_labels.get(mode, []): self.extra_labels[mode].remove(sel) self.refresh_extra_labels_view() def load_extra_labels_json(self): path = filedialog.askopenfilename(title="加载额外类别 JSON", filetypes=[("JSON 文件", "*.json")]) if not path: return try: with open(path, "r", encoding="utf-8") as f: data = json.load(f) # 支持:数组 或 {mode: [..]} 两种格式 mode = self.mode_var.get() if isinstance(data, list): self.extra_labels[mode] = [str(x) for x in data] elif isinstance(data, dict): for k in ("dish", "whole", "processed"): if k in data and isinstance(data[k], list): self.extra_labels[k] = [str(x) for x in data[k]] self.extra_labels_file_path = path self.refresh_extra_labels_view() except Exception as e: messagebox.showerror("错误", f"加载失败: {e}") def save_extra_labels_json(self): # 保存为 {mode: [...]} 方便多模式复用 path = filedialog.asksaveasfilename(title="保存额外类别 JSON", defaultextension=".json", filetypes=[("JSON 文件", "*.json")]) if not path: return try: with open(path, "w", encoding="utf-8") as f: json.dump(self.extra_labels, f, ensure_ascii=False, indent=2) self.extra_labels_file_path = path messagebox.showinfo("成功", "已保存额外类别 JSON") except Exception as e: messagebox.showerror("错误", f"保存失败: {e}") # Fewshot 编辑 def refresh_fewshot_view(self): self.fewshot_text.configure(state="normal") self.fewshot_text.delete("1.0", "end") for k, v in self.fewshot_hints.items(): self.fewshot_text.insert("end", f"{k}:{v}\n") self.fewshot_text.configure(state="disabled") def add_or_update_fewshot(self): k = _normalize(self.fs_label_var.get().strip()) v = self.fs_hint_var.get().strip() if not k or not v: return self.fewshot_hints[k] = v self.fs_label_var.set("") self.fs_hint_var.set("") self.refresh_fewshot_view() def remove_selected_fewshot(self): try: sel = self.fewshot_text.get("sel.first", "sel.last") except Exception: sel = "" if not sel: return # 选中行以全角冒号或中文冒号分割 line = sel.strip().split(":", 1)[0] key = _normalize(line) if key in self.fewshot_hints: del self.fewshot_hints[key] self.refresh_fewshot_view() def load_fewshot_json(self): path = filedialog.askopenfilename(title="加载 Fewshot JSON", filetypes=[("JSON 文件", "*.json")]) if not path: return try: with open(path, "r", encoding="utf-8") as f: data = json.load(f) if isinstance(data, dict): # 仅接收 {label: hint} self.fewshot_hints = {str(k): str(v) for k, v in data.items()} self.fewshot_file_path = path self.refresh_fewshot_view() else: raise ValueError("JSON 格式应为 {label: hint}") except Exception as e: messagebox.showerror("错误", f"加载失败: {e}") def save_fewshot_json(self): path = filedialog.asksaveasfilename(title="保存 Fewshot JSON", defaultextension=".json", filetypes=[("JSON 文件", "*.json")]) if not path: return try: with open(path, "w", encoding="utf-8") as f: json.dump(self.fewshot_hints, f, ensure_ascii=False, indent=2) self.fewshot_file_path = path messagebox.showinfo("成功", "已保存 Fewshot JSON") 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: messagebox.showinfo("提示", "请先上传图片") return self.recognition_start_time = time.time() self.recognize_button.configure(state="disabled", text="识别中...") threading.Thread(target=self.recognize_images, daemon=True).start() def _build_final_labels(self, mode: str) -> List[str]: # 基础 labels 来自数据集 + alias base_labels = build_labels(mode, self.alias_map_path) # 合并额外标签 extra = self.extra_labels.get(mode, []) final = list(dict.fromkeys(list(base_labels) + list(extra))) return final def _infer_true_class(self, img_path: str, labels: List[str]) -> Optional[str]: # 从父目录名中尝试匹配到 labels try: parent = os.path.basename(os.path.dirname(os.path.normpath(img_path))) n = _normalize(parent) if self.mode_var.get() in ("whole", "processed"): n = _base_ingredient(n) return n if n in labels else None except Exception: return None def recognize_images(self): try: mode = self.mode_var.get() 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_openset( image_path=img['path'], provider=provider, embedder=self.openset_embedder, matcher=self.openset_matcher, top_k=top_k, min_match_score=min_score, ) # 解析开放式结果 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": 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-Openset] 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": 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)) except Exception as e: raise Exception(f"开放式识别初始化失败: {e}") def _create_vlm_provider(self) -> Optional[VLMProvider]: """根据配置创建 VLM Provider""" provider_type = self.provider_var.get() try: if provider_type == "ollama": url = self.ollama_url_var.get().strip() or DEFAULT_OLLAMA_URL model = self.vlm_model_var.get().strip() or DEFAULT_VLM_MODEL return OllamaProvider(ollama_url=url, model=model) elif provider_type == "kimi": api_key = self.kimi_api_key_var.get().strip() if not api_key: messagebox.showerror("错误", "Kimi API Key 不能为空") return None base_url = self.kimi_base_url_var.get().strip() or "https://api.moonshot.cn/v1" model = self.kimi_model_var.get().strip() or "moonshot-v1-32k-vision-preview" return KimiProvider(api_key=api_key, base_url=base_url, model=model) else: messagebox.showerror("错误", f"未知的 Provider 类型: {provider_type}") return None except Exception as e: messagebox.showerror("错误", f"创建 Provider 失败: {e}") return None def update_progress(self, current: int, total: int): self.recognize_button.configure(text=f"识别中... ({current}/{total})") self.update_images_display() self.update_results_display() def recognition_completed(self): if self.recognition_start_time is not None: self.recognition_duration = time.time() - self.recognition_start_time self.recognize_button.configure(state="normal", text="开始识别") self.update_stats() messagebox.showinfo("完成", f"所有图片识别完成!耗时: {self.recognition_duration:.2f}秒") # -------------------- 配置管理 -------------------- def get_config_path(self) -> str: """获取配置文件路径(项目根目录下的 config.json)""" return os.path.join(os.path.dirname(os.path.dirname(__file__)), "vlm_config.json") def save_config(self): """保存当前配置到 JSON 文件""" config = { "mode": self.mode_var.get(), "provider": self.provider_var.get(), "ollama": { "url": self.ollama_url_var.get(), "model": self.vlm_model_var.get(), }, "kimi": { "api_key": self.kimi_api_key_var.get(), "base_url": self.kimi_base_url_var.get(), "model": self.kimi_model_var.get(), }, "alias_map_path": self.alias_map_path, "fewshot_enabled": self.fewshot_enabled_var.get(), "fewshot_hints": self.fewshot_hints, "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: config_path = self.get_config_path() with open(config_path, "w", encoding="utf-8") as f: json.dump(config, f, ensure_ascii=False, indent=2) messagebox.showinfo("成功", f"配置已保存到 {os.path.basename(config_path)}") except Exception as e: messagebox.showerror("错误", f"保存配置失败: {e}") def load_config(self): """自动加载配置(启动时调用)""" config_path = self.get_config_path() if not os.path.exists(config_path): return try: with open(config_path, "r", encoding="utf-8") as f: config = json.load(f) # 恢复配置 if "mode" in config: self.mode_var.set(config["mode"]) if "provider" in config: self.provider_var.set(config["provider"]) if "ollama" in config: self.ollama_url_var.set(config["ollama"].get("url", DEFAULT_OLLAMA_URL)) self.vlm_model_var.set(config["ollama"].get("model", DEFAULT_VLM_MODEL)) if "kimi" in config: self.kimi_api_key_var.set(config["kimi"].get("api_key", "")) self.kimi_base_url_var.set(config["kimi"].get("base_url", "https://api.moonshot.cn/v1")) self.kimi_model_var.set(config["kimi"].get("model", "moonshot-v1-32k-vision-preview")) if "alias_map_path" in config and config["alias_map_path"]: self.alias_map_path = config["alias_map_path"] self.alias_label_var.set(os.path.basename(self.alias_map_path)) if "fewshot_enabled" in config: self.fewshot_enabled_var.set(config["fewshot_enabled"]) if "fewshot_hints" in config: self.fewshot_hints = config["fewshot_hints"] if "fewshot_file_path" in config: self.fewshot_file_path = config["fewshot_file_path"] if "extra_labels" in config: self.extra_labels = config["extra_labels"] 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() print(f"[Config] 配置已从 {config_path} 加载") except Exception as e: print(f"[Config] 加载配置失败: {e}") def load_config_from_file(self): """从用户选择的文件加载配置""" path = filedialog.askopenfilename( title="选择配置文件", filetypes=[("JSON 文件", "*.json"), ("所有文件", "*.*")] ) if not path: return try: with open(path, "r", encoding="utf-8") as f: config = json.load(f) # 恢复配置(同上) if "mode" in config: self.mode_var.set(config["mode"]) if "provider" in config: self.provider_var.set(config["provider"]) if "ollama" in config: self.ollama_url_var.set(config["ollama"].get("url", DEFAULT_OLLAMA_URL)) self.vlm_model_var.set(config["ollama"].get("model", DEFAULT_VLM_MODEL)) if "kimi" in config: self.kimi_api_key_var.set(config["kimi"].get("api_key", "")) self.kimi_base_url_var.set(config["kimi"].get("base_url", "https://api.moonshot.cn/v1")) self.kimi_model_var.set(config["kimi"].get("model", "moonshot-v1-32k-vision-preview")) if "alias_map_path" in config and config["alias_map_path"]: self.alias_map_path = config["alias_map_path"] self.alias_label_var.set(os.path.basename(self.alias_map_path)) if "fewshot_enabled" in config: self.fewshot_enabled_var.set(config["fewshot_enabled"]) if "fewshot_hints" in config: self.fewshot_hints = config["fewshot_hints"] if "fewshot_file_path" in config: self.fewshot_file_path = config["fewshot_file_path"] if "extra_labels" in config: self.extra_labels = config["extra_labels"] 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() messagebox.showinfo("成功", f"配置已从 {os.path.basename(path)} 加载") except Exception as e: messagebox.showerror("错误", f"加载配置失败: {e}") def main(): root = TkinterDnD.Tk() # 必须使用 TkinterDnD.Tk 以支持拖拽 app = MultiModalFoodApp(root) root.mainloop() if __name__ == "__main__": main()