From 1f04bb5c018522a437aa6cdc0158e8ac7b91a169 Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Mon, 29 Sep 2025 14:16:50 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0=E4=BA=86=E6=96=B0=E5=A2=9E?= =?UTF-8?q?=E8=8F=9C=E5=93=81=E5=8A=9F=E8=83=BD=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- classifier/embedding_food_classifier_app.py | 524 +++++++++++++++++++- 1 file changed, 513 insertions(+), 11 deletions(-) diff --git a/classifier/embedding_food_classifier_app.py b/classifier/embedding_food_classifier_app.py index 6900ea2..dcd95f5 100644 --- a/classifier/embedding_food_classifier_app.py +++ b/classifier/embedding_food_classifier_app.py @@ -262,22 +262,31 @@ class EmbeddingFoodClassifierApp: ) self.images_display_frame.pack(fill="both", expand=True, padx=15, pady=(0, 15)) - # 右侧框架 - 识别结果区域 + # 右侧框架 - 识别结果和新增Embedding区域 self.right_frame = ctk.CTkFrame(self.main_frame, width=700) self.right_frame.pack(side="right", fill="both", expand=True, padx=(10, 0), pady=0) self.right_frame.pack_propagate(False) - # 右侧标题 - self.right_title = ctk.CTkLabel( - self.right_frame, - text="识别结果 (基于特征相似度)", - font=("Arial", 16, "bold") - ) - self.right_title.pack(pady=(15, 10)) + # 创建右侧的选项卡视图 + 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.create_results_tab() + + # 新增Embedding选项卡 + self.embedding_tab = self.right_tabview.add("新增Embedding") + self.create_embedding_tab() + + # 默认选中识别结果选项卡 + self.right_tabview.set("识别结果") + + def create_results_tab(self): + """创建识别结果选项卡""" # 统计信息框架 - self.stats_frame = ctk.CTkFrame(self.right_frame) - self.stats_frame.pack(fill="x", padx=15, pady=(0, 10)) + self.stats_frame = ctk.CTkFrame(self.results_tab) + self.stats_frame.pack(fill="x", padx=15, pady=(10, 10)) # 统计标签 self.stats_label = ctk.CTkLabel( @@ -289,11 +298,504 @@ class EmbeddingFoodClassifierApp: # 识别结果显示区域 self.results_display_frame = ctk.CTkScrollableFrame( - self.right_frame, + self.results_tab, label_text="识别详情 (相似度排序)" ) self.results_display_frame.pack(fill="both", expand=True, padx=15, pady=(0, 15)) + + def create_embedding_tab(self): + """创建新增Embedding选项卡""" + # 新增Embedding区域的变量初始化 + self.embedding_images = [] # 待添加到向量库的图片 + self.selected_class = ctk.StringVar(value="") # 选中的类别 + self.new_class_name = ctk.StringVar(value="") # 新类别名称 + # 标题 + title_label = ctk.CTkLabel( + self.embedding_tab, + text="新增Embedding向量", + font=("Arial", 16, "bold") + ) + title_label.pack(pady=(10, 15)) + + # 类别选择区域 + class_frame = ctk.CTkFrame(self.embedding_tab) + class_frame.pack(fill="x", padx=15, pady=(0, 10)) + + class_title = ctk.CTkLabel( + class_frame, + text="选择或新增菜品类别:", + font=("Arial", 12, "bold") + ) + class_title.pack(pady=(10, 5)) + + # 现有类别选择 + existing_class_frame = ctk.CTkFrame(class_frame) + existing_class_frame.pack(fill="x", padx=10, pady=5) + + existing_label = ctk.CTkLabel( + existing_class_frame, + text="选择现有类别:", + font=("Arial", 11) + ) + existing_label.pack(side="left", padx=(10, 5), pady=10) + + self.class_dropdown = ctk.CTkComboBox( + existing_class_frame, + values=self.class_names, + variable=self.selected_class, + width=200, + command=self.on_class_selected + ) + self.class_dropdown.pack(side="left", padx=5, pady=10) + + # 新类别输入 + new_class_frame = ctk.CTkFrame(class_frame) + new_class_frame.pack(fill="x", padx=10, pady=5) + + new_label = ctk.CTkLabel( + new_class_frame, + text="或新增类别:", + font=("Arial", 11) + ) + new_label.pack(side="left", padx=(10, 5), pady=10) + + self.new_class_entry = ctk.CTkEntry( + new_class_frame, + textvariable=self.new_class_name, + placeholder_text="输入新的菜品类别名称", + width=200 + ) + self.new_class_entry.pack(side="left", padx=5, pady=10) + self.new_class_entry.bind('', self.on_new_class_input) + + # 图片上传区域 + upload_frame = ctk.CTkFrame(self.embedding_tab) + upload_frame.pack(fill="x", padx=15, pady=(0, 10)) + + upload_title = ctk.CTkLabel( + upload_frame, + text="上传样本图片:", + font=("Arial", 12, "bold") + ) + upload_title.pack(pady=(10, 5)) + + # 拖拽上传区域 + self.embedding_upload_frame = ctk.CTkFrame(upload_frame, fg_color=("gray90", "gray20")) + self.embedding_upload_frame.pack(fill="x", padx=10, pady=(0, 10), ipady=30) + + self.embedding_upload_label = ctk.CTkLabel( + self.embedding_upload_frame, + text="拖拽图片到这里或点击下方按钮选择支持多图片上传", + font=("Arial", 12), + text_color=("gray40", "gray60") + ) + self.embedding_upload_label.pack(expand=True) + + # 绑定拖放事件 + self.embedding_upload_frame.drop_target_register(DND_FILES) + self.embedding_upload_frame.dnd_bind('<>', self.handle_embedding_drop) + + # 按钮区域 + embedding_button_frame = ctk.CTkFrame(upload_frame) + embedding_button_frame.pack(fill="x", padx=10, pady=(0, 10)) + + # 选择图片按钮 + self.select_embedding_button = ctk.CTkButton( + embedding_button_frame, + text="选择图片", + command=self.select_embedding_images, + width=100, + height=30 + ) + self.select_embedding_button.pack(side="left", padx=(10, 5), pady=5) + + # 清空图片按钮 + self.clear_embedding_button = ctk.CTkButton( + embedding_button_frame, + text="清空图片", + command=self.clear_embedding_images, + width=100, + height=30, + fg_color="gray", + hover_color="darkgray" + ) + self.clear_embedding_button.pack(side="left", padx=5, pady=5) + + # 添加到向量库按钮 + self.add_embedding_button = ctk.CTkButton( + embedding_button_frame, + text="添加到向量库", + command=self.add_to_vector_db, + width=120, + height=30, + fg_color="green", + hover_color="darkgreen" + ) + self.add_embedding_button.pack(side="right", padx=(5, 10), pady=5) + self.add_embedding_button.configure(state="disabled") + + # 已选择图片显示区域 + self.embedding_images_frame = ctk.CTkScrollableFrame( + self.embedding_tab, + label_text="待添加的图片" + ) + self.embedding_images_frame.pack(fill="both", expand=True, padx=15, pady=(0, 15)) + + # 状态显示区域 + self.status_frame = ctk.CTkFrame(self.embedding_tab) + self.status_frame.pack(fill="x", padx=15, pady=(0, 15)) + + self.status_label = ctk.CTkLabel( + self.status_frame, + text="请选择类别并上传图片", + font=("Arial", 11), + text_color="gray" + ) + self.status_label.pack(pady=10) + + def on_class_selected(self, value): + """当选择现有类别时""" + if value: + self.new_class_name.set("") # 清空新类别输入 + self.update_add_button_state() + + def on_new_class_input(self, event): + """当输入新类别时""" + if self.new_class_name.get().strip(): + self.selected_class.set("") # 清空现有类别选择 + self.class_dropdown.set("") + self.update_add_button_state() + + def update_add_button_state(self): + """更新添加按钮状态""" + has_class = bool(self.selected_class.get() or self.new_class_name.get().strip()) + has_images = bool(self.embedding_images) + + if has_class and has_images: + self.add_embedding_button.configure(state="normal") + if self.new_class_name.get().strip(): + self.status_label.configure( + text=f"准备为新类别 '{self.new_class_name.get().strip()}' 添加 {len(self.embedding_images)} 张图片", + text_color="green" + ) + else: + self.status_label.configure( + text=f"准备为类别 '{self.selected_class.get()}' 添加 {len(self.embedding_images)} 张图片", + text_color="green" + ) + else: + self.add_embedding_button.configure(state="disabled") + if not has_class: + self.status_label.configure(text="请选择或输入类别", text_color="orange") + elif not has_images: + self.status_label.configure(text="请上传图片", text_color="orange") + + def select_embedding_images(self): + """选择要添加到向量库的图片""" + file_paths = filedialog.askopenfilenames( + title="选择要添加到向量库的图片", + filetypes=[ + ("图像文件", "*.jpg *.jpeg *.png *.bmp *.gif"), + ("JPEG文件", "*.jpg *.jpeg"), + ("PNG文件", "*.png"), + ("所有文件", "*.*") + ] + ) + + if file_paths: + for file_path in file_paths: + self.add_embedding_image(file_path) + + def handle_embedding_drop(self, event): + """处理拖拽到embedding区域的文件""" + files = event.data.split() + for file_path in files: + # 清理文件路径 + file_path = file_path.strip('{}').strip('"') + file_path = os.path.normpath(file_path) + + # 检查是否为图片文件 + valid_extensions = ('.jpg', '.jpeg', '.png', '.bmp', '.gif') + if file_path.lower().endswith(valid_extensions): + self.add_embedding_image(file_path) + + def add_embedding_image(self, file_path): + """添加图片到embedding列表""" + try: + # 检查文件是否存在 + if not os.path.exists(file_path): + messagebox.showerror("错误", f"文件不存在: {file_path}") + return + + # 检查是否已经添加过 + if file_path in [img['path'] for img in self.embedding_images]: + messagebox.showinfo("提示", "该图片已经添加过了") + return + + # 加载图片 + image = self.load_image_with_chinese_path(file_path) + if image is None: + messagebox.showerror("错误", f"无法读取图片: {file_path}") + return + + # 添加到列表 + image_info = { + 'path': file_path, + 'name': os.path.basename(file_path), + 'image': image + } + self.embedding_images.append(image_info) + + # 更新显示 + self.update_embedding_images_display() + self.update_add_button_state() + + except Exception as e: + messagebox.showerror("错误", f"添加图片时出错: {str(e)}") + + def clear_embedding_images(self): + """清空embedding图片列表""" + if self.embedding_images: + result = messagebox.askyesno("确认", "确定要清空所有待添加的图片吗?") + if result: + self.embedding_images.clear() + self.update_embedding_images_display() + self.update_add_button_state() + + def update_embedding_images_display(self): + """更新embedding图片显示""" + # 清空当前显示 + for widget in self.embedding_images_frame.winfo_children(): + widget.destroy() + + if not self.embedding_images: + no_image_label = ctk.CTkLabel( + self.embedding_images_frame, + text="暂无待添加的图片", + font=("Arial", 12), + text_color="gray" + ) + no_image_label.pack(pady=20) + return + + # 显示每张图片 + for i, img_info in enumerate(self.embedding_images): + # 创建图片框架 + img_frame = ctk.CTkFrame(self.embedding_images_frame) + img_frame.pack(fill="x", padx=5, pady=5) + + # 缩放图片用于显示 + display_image = self.resize_image_for_display(img_info['image'], 80, 80) + display_image = cv2.cvtColor(display_image, cv2.COLOR_BGR2RGB) + pil_image = Image.fromarray(display_image) + ctk_image = ctk.CTkImage(light_image=pil_image, dark_image=pil_image, size=(80, 80)) + + # 图片标签 + img_label = ctk.CTkLabel(img_frame, image=ctk_image, text="") + img_label.image = ctk_image + img_label.pack(side="left", padx=10, pady=10) + + # 信息框架 + info_frame = ctk.CTkFrame(img_frame) + info_frame.pack(side="left", fill="both", expand=True, padx=10, pady=10) + + # 文件名 + name_label = ctk.CTkLabel( + info_frame, + text=f"文件名: {img_info['name']}", + anchor="w", + font=("Arial", 11) + ) + name_label.pack(fill="x", padx=5, pady=2) + + # 删除按钮 + delete_button = ctk.CTkButton( + img_frame, + text="删除", + command=lambda idx=i: self.remove_embedding_image(idx), + width=50, + height=25, + fg_color="red", + hover_color="darkred" + ) + delete_button.pack(side="right", padx=10, pady=10) + + def remove_embedding_image(self, index): + """删除embedding图片""" + if index < len(self.embedding_images): + self.embedding_images.pop(index) + self.update_embedding_images_display() + self.update_add_button_state() + + def add_to_vector_db(self): + """添加图片到向量数据库""" + if not self.embedding_images: + messagebox.showwarning("警告", "请先上传图片") + return + + # 确定目标类别 + target_class = self.new_class_name.get().strip() or self.selected_class.get() + if not target_class: + messagebox.showwarning("警告", "请选择或输入类别") + return + + # 确认操作 + if self.new_class_name.get().strip(): + message = f"确定要创建新类别 '{target_class}' 并添加 {len(self.embedding_images)} 张图片到向量库吗?" + else: + message = f"确定要为类别 '{target_class}' 添加 {len(self.embedding_images)} 张图片到向量库吗?" + + result = messagebox.askyesno("确认", message) + if not result: + return + + # 在新线程中执行添加操作 + self.add_embedding_button.configure(state="disabled", text="添加中...") + self.status_label.configure(text="正在处理图片并添加到向量库...", text_color="blue") + + threading.Thread( + target=self.process_embedding_addition, + args=(target_class,), + daemon=True + ).start() + + def process_embedding_addition(self, target_class): + """处理embedding添加的后台任务""" + try: + if self.model is None or self.faiss_index is None: + # 如果模型或索引未加载,显示模拟结果 + self.root.after(0, lambda: self.show_embedding_result( + target_class, len(self.embedding_images), True, "模拟模式:图片已添加到向量库" + )) + return + + # 处理每张图片 + new_embeddings = [] + new_paths = [] + new_labels = [] + + # 确定类别索引 + if target_class in self.class_names: + class_idx = self.class_names.index(target_class) + else: + # 新类别,添加到类别列表 + class_idx = len(self.class_names) + self.class_names.append(target_class) + self.class_to_idx[target_class] = class_idx + self.idx_to_class[str(class_idx)] = target_class + + for img_info in self.embedding_images: + try: + # 转换图片格式 + image_rgb = cv2.cvtColor(img_info['image'], cv2.COLOR_BGR2RGB) + pil_image = Image.fromarray(image_rgb) + + # 提取特征向量 + embedding = self.model.extract_embedding(pil_image, normalize=True) + embedding = embedding.reshape(1, -1).astype(np.float32) + + new_embeddings.append(embedding) + new_paths.append(img_info['path']) + new_labels.append(class_idx) + + except Exception as e: + print(f"处理图片 {img_info['name']} 时出错: {e}") + continue + + if new_embeddings: + # 合并所有新的embedding + all_new_embeddings = np.vstack(new_embeddings) + + # 添加到FAISS索引 + self.faiss_index.add(all_new_embeddings) + + # 更新路径和标签列表 + self.image_paths.extend(new_paths) + self.labels.extend(new_labels) + + # 保存更新后的索引和元数据 + self.save_updated_index() + + success_message = f"成功添加 {len(new_embeddings)} 张图片到向量库" + self.root.after(0, lambda: self.show_embedding_result( + target_class, len(new_embeddings), True, success_message + )) + else: + self.root.after(0, lambda: self.show_embedding_result( + target_class, 0, False, "没有成功处理任何图片" + )) + + except Exception as e: + error_message = f"添加过程中出错: {str(e)}" + self.root.after(0, lambda: self.show_embedding_result( + target_class, 0, False, error_message + )) + + def save_updated_index(self): + """保存更新后的索引和元数据""" + try: + index_dir = "../faiss_vector_db/faiss_index" + + # 保存FAISS索引 + index_path = os.path.join(index_dir, 'faiss_index.bin') + faiss.write_index(self.faiss_index, index_path) + + # 保存图片路径 + paths_path = os.path.join(index_dir, 'image_paths.pkl') + with open(paths_path, 'wb') as f: + pickle.dump(self.image_paths, f) + + # 保存标签 + labels_path = os.path.join(index_dir, 'labels.pkl') + with open(labels_path, 'wb') as f: + pickle.dump(self.labels, f) + + # 更新类别信息 + self.class_info = { + 'class_names': self.class_names, + 'class_to_idx': self.class_to_idx, + 'idx_to_class': self.idx_to_class + } + + # 保存类别信息 + class_info_path = os.path.join(index_dir, 'class_info.json') + with open(class_info_path, 'w', encoding='utf-8') as f: + json.dump(self.class_info, f, ensure_ascii=False, indent=2) + + print("索引和元数据已成功保存") + + except Exception as e: + print(f"保存索引时出错: {e}") + + def show_embedding_result(self, target_class, count, success, message): + """显示embedding添加结果""" + self.add_embedding_button.configure(state="normal", text="添加到向量库") + + if success: + self.status_label.configure(text=message, text_color="green") + messagebox.showinfo("成功", message) + + # 清空已添加的图片 + self.embedding_images.clear() + self.update_embedding_images_display() + + # 重置选择 + self.selected_class.set("") + self.new_class_name.set("") + self.class_dropdown.set("") + + # 更新类别下拉框(如果有新类别) + if target_class not in self.class_dropdown.cget("values"): + current_values = list(self.class_dropdown.cget("values")) + current_values.append(target_class) + self.class_dropdown.configure(values=current_values) + + self.update_add_button_state() + else: + self.status_label.configure(text=message, text_color="red") + messagebox.showerror("错误", message) + def select_images(self): """选择图片文件""" file_paths = filedialog.askopenfilenames(