1171 lines
55 KiB
Python
1171 lines
55 KiB
Python
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('<<Drop>>', self.handle_drop)
|
||
self.upload_frame.bind('<Enter>', self.on_drag_enter)
|
||
self.upload_frame.bind('<Leave>', 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("<Button-1>", 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("<Button-1>", 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()
|