增加了新增菜品功能。
This commit is contained in:
@@ -262,22 +262,31 @@ class EmbeddingFoodClassifierApp:
|
|||||||
)
|
)
|
||||||
self.images_display_frame.pack(fill="both", expand=True, padx=15, pady=(0, 15))
|
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 = 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(side="right", fill="both", expand=True, padx=(10, 0), pady=0)
|
||||||
self.right_frame.pack_propagate(False)
|
self.right_frame.pack_propagate(False)
|
||||||
|
|
||||||
# 右侧标题
|
# 创建右侧的选项卡视图
|
||||||
self.right_title = ctk.CTkLabel(
|
self.right_tabview = ctk.CTkTabview(self.right_frame)
|
||||||
self.right_frame,
|
self.right_tabview.pack(fill="both", expand=True, padx=15, pady=15)
|
||||||
text="识别结果 (基于特征相似度)",
|
|
||||||
font=("Arial", 16, "bold")
|
|
||||||
)
|
|
||||||
self.right_title.pack(pady=(15, 10))
|
|
||||||
|
|
||||||
|
# 识别结果选项卡
|
||||||
|
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 = ctk.CTkFrame(self.results_tab)
|
||||||
self.stats_frame.pack(fill="x", padx=15, pady=(0, 10))
|
self.stats_frame.pack(fill="x", padx=15, pady=(10, 10))
|
||||||
|
|
||||||
# 统计标签
|
# 统计标签
|
||||||
self.stats_label = ctk.CTkLabel(
|
self.stats_label = ctk.CTkLabel(
|
||||||
@@ -289,11 +298,504 @@ class EmbeddingFoodClassifierApp:
|
|||||||
|
|
||||||
# 识别结果显示区域
|
# 识别结果显示区域
|
||||||
self.results_display_frame = ctk.CTkScrollableFrame(
|
self.results_display_frame = ctk.CTkScrollableFrame(
|
||||||
self.right_frame,
|
self.results_tab,
|
||||||
label_text="识别详情 (相似度排序)"
|
label_text="识别详情 (相似度排序)"
|
||||||
)
|
)
|
||||||
self.results_display_frame.pack(fill="both", expand=True, padx=15, pady=(0, 15))
|
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('<KeyRelease>', 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('<<Drop>>', 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):
|
def select_images(self):
|
||||||
"""选择图片文件"""
|
"""选择图片文件"""
|
||||||
file_paths = filedialog.askopenfilenames(
|
file_paths = filedialog.askopenfilenames(
|
||||||
|
|||||||
Reference in New Issue
Block a user