From 7af923f1cedf55b8fcf3ff0ea02231c11597ffef Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Mon, 2 Mar 2026 17:26:00 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0=E5=90=84=E4=B8=AA=E8=84=9A?= =?UTF-8?q?=E6=9C=AC=E8=AF=B4=E6=98=8E=EF=BC=8C=E6=95=B4=E7=90=86=E9=A1=B9?= =?UTF-8?q?=E7=9B=AE=E7=9B=AE=E5=BD=95=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- CLAUDE.md | 132 +++++++++++ GRID_SEARCH_README.md | 149 ------------ build_faiss_index.py | 519 ------------------------------------------ 3 files changed, 132 insertions(+), 668 deletions(-) create mode 100644 CLAUDE.md delete mode 100644 GRID_SEARCH_README.md delete mode 100644 build_faiss_index.py diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 0000000..5b18cd6 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,132 @@ +# CLAUDE.md + +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. + +## 项目概述 + +本项目采用**度量学习(Metric Learning)**的方式进行菜品识别和食材识别。核心思路是用 ResNet50 提取图像嵌入向量,配合 FAISS 向量库进行相似度检索,而非传统的端到端分类,从而支持无需重新训练即可快速添加新类别。 + +**三个并行识别任务:** +- `DishClassification` — 菜品识别 +- `WholeIngredientRecognition` — 整体食材识别(未处理) +- `ProcessedIngredientRecognition` — 处理后食材识别(切块等) + +## 常用命令 + +**训练嵌入模型(CosFace 损失,主要方案):** +```bash +python train/train_cosface_embedding.py +``` + +**CosFace 超参数网格搜索(scale × margin 组合):** +```bash +python train/grid_search_cosface.py +``` + +**训练嵌入模型(Triplet + Center Loss,备用方案,目前没使用这个了,主要是未采集的图片和现有分类的相似度太高了):** +```bash +python train/train_triplet_embedding.py +``` + +**基于训练好的模型构建 FAISS 索引:** +```bash +python faiss_vector_db/build_faiss_index.py +``` + +**启动主识别界面(嵌入向量方案):** +```bash +python classifier/embedding_food_classifier_app.py +``` + +**启动数据管理与下载界面:** +```bash +python data_management/app.py +``` + +**训练 SegFormer 食物区域分割模型:** +```bash +python SegFormer/training/train_minimal.py +``` + +**测试 SegFormer 推理效果:** +```bash +python SegFormer/inference/test_model.py +``` + +## 架构说明 + +### 度量学习流水线 + +``` +输入图像 → ResNet50 主干网络(ImageNet1K_V2 预训练) + → 512 维嵌入向量 → L2 归一化 + → FAISS IndexFlatIP(余弦相似度检索) + → 返回 Top-K 相似食物及置信度分数 +``` + +**训练损失函数:** +- **CosFace**(`train_cosface_embedding.py`):大间距余弦损失;scale ≈ 64,margin ≈ 0.35–0.40(因任务而异)。泛化能力更强,为首选方案。 +- **Triplet + Center Loss**(`train_triplet_embedding.py`):在线难样本挖掘,margin = 0.5,Center Loss 权重因任务而异(0.5–20)。 + +### 关键模块 + +| 路径 | 功能说明 | +|------|---------| +| `net/resnet_embedding.py` | ResNet50 嵌入主干网络(512 维,L2 归一化) | +| `net/food_net.py` | 遗留的自定义 CNN,用于端到端分类(4 类) | +| `faiss_vector_db/faiss_manager.py` | FAISS 索引的增删查存核心类 | +| `faiss_vector_db/build_faiss_index.py` | 遍历数据集、提取嵌入、构建索引 | +| `classifier/embedding_food_classifier_app.py` | 主 GUI(CustomTkinter,支持拖拽上传) | +| `exp_multimodal/vlm_classifier.py` | 实验性:VLM(Ollama/Kimi)+ FAISS 开放集识别 | +| `data_management/app.py` | 数据采集 GUI,从 MySQL 后端下载图片 | +| `settings/settings.py` | 全局配置(NUM_CLASSES、MODEL_DIR、DEVICE 等) | + +### FAISS 索引存储结构 + +每个任务在 `faiss_vector_db/<任务名>/faiss_index/` 下独立存储: +- `food_embeddings.index` — FAISS 二进制索引文件 +- `labels.json` / `food_labels.json` — 标签到食物名称的映射 +- `embeddings.npy` — 原始向量的 numpy 数组 + +### 数据集目录结构 + +``` +dataset/<任务名>/ +├── train/<类别名>/ (图片) +├── val/<类别名>/ +└── test/<类别名>/ +``` + +### VLM 多模态集成(实验性) + +`exp_multimodal/` 包含一套绕过封闭类别列表的开放集识别流水线: +- `vlm_providers/ollama_provider.py` — 本地 Ollama 服务(Qwen2.5-VL、LLaVA),地址:`http://192.168.1.250:11434` +- `vlm_providers/kimi_provider.py` — Kimi/Moonshot API(`moonshot-v1-128k-vision-preview`) +- `build_dish_name_index.py` — 使用 `bge-large-zh-v1.5` 为菜品名称构建文本向量索引 +- VLM 配置保存在 `vlm_config.json`(提供商选择、API Key、候选菜品列表) + +### SegFormer 食物分割 + +`SegFormer/` 包含用于提取食物区域的微调 SegFormer: +- 两阶段训练:先冻结编码器训练,再全量微调 +- 损失函数:CrossEntropy + 0.5 × Dice Loss +- 数据以 COCO 格式存放于 `SegFormer/data/` + +## 模型文件 + +训练好的模型保存在 `model/` 下的版本子目录中: +- `best_cosface_model.pth` — CosFace 训练的嵌入模型 +- `best_embedding_model.pth` — Triplet/Center Loss 训练的嵌入模型 +- `grid_search_<时间戳>/best_model_s_m.pth` — 网格搜索最优结果 + +当前使用的模型路径通过 `settings/settings.py` 中的 `MODEL_DIR` 指定。 + +## 依赖安装 + +```bash +pip install -r requirements.txt +``` + +核心依赖:`torch==2.7.0`、`torchvision==0.22.0`、`faiss-cpu>=1.7.0`、`customtkinter==5.2.2`、`tkinterdnd2==0.4.3`、`opencv-python==4.11.0.86`。 + +数据管理界面连接远端 MySQL 数据库(`192.168.1.250:3308`)获取图片元数据,同时使用本地 SQLite 记录下载历史。 diff --git a/GRID_SEARCH_README.md b/GRID_SEARCH_README.md deleted file mode 100644 index 3f00253..0000000 --- a/GRID_SEARCH_README.md +++ /dev/null @@ -1,149 +0,0 @@ -# CosFace超参数网格搜索使用指南 - -## 概述 - -已完成CosFace超参数网格搜索功能的实施,用于优化`scale (s)`和`margin (m)`两个核心超参数。 - -## 修改文件清单 - -### 1. `requirements.txt` -- ✅ 添加依赖:`pandas>=2.0.0` 和 `seaborn>=0.12.0` - -### 2. `train/train_cosface_embedding.py` -- ✅ `TaskConfig`添加`test_dir`字段 -- ✅ `TASKS`字典中所有任务添加测试集路径 - -### 3. `train/grid_search_cosface.py` (新建) -- ✅ 完整的网格搜索脚本 (~350行) - -## 搜索空间配置 - -当前为dish任务配置的搜索空间: -```python -GRID_PARAMS = { - 'dish': { - 's': [56.0, 60.0, 64.0, 68.0], # 4个scale值 - 'm': [0.32, 0.35, 0.38, 0.40], # 4个margin值 - } -} -``` - -**总配置数**: 4 × 4 = **16组实验** - -## 使用方法 - -### 1. 安装依赖 -```bash -pip install pandas>=2.0.0 seaborn>=0.12.0 -``` - -### 2. 运行网格搜索 -```bash -# 完整搜索(16组配置) -python train/grid_search_cosface.py --task dish - -# 测试运行(限制配置数) -python train/grid_search_cosface.py --task dish --max_configs 2 - -# 自定义参数 -python train/grid_search_cosface.py --task dish --epochs 100 --patience 10 --min_epochs 25 -``` - -### 3. 参数说明 -- `--task`: 任务名称 (`dish` / `whole_ingredient` / `processed_ingredient`) -- `--max_configs`: 限制最大配置数(用于测试,可选) -- `--epochs`: 每个配置的最大训练轮数(默认100) -- `--patience`: 早停容忍轮数(默认10) -- `--min_epochs`: 最小训练轮数(默认25) - -## 输出文件 - -运行后会在`model/DishClassification/grid_search_YYYYMMDD_HHMMSS/`目录下生成: - -1. **`grid_search_results.csv`** - 所有配置的详细结果表格 -2. **`heatmap.png`** - 参数热力图(测试集准确率) -3. **`summary.txt`** - 搜索总结报告(含最佳配置) -4. **`model_s{s}_m{m}.pth`** - 每个配置的模型权重 -5. **`grid_search.log`** - 完整训练日志 - -## 评估指标 - -- **主要指标**: 测试集准确率 (`test_acc`) - 用于选择最佳配置 -- **辅助指标**: 验证集准确率、训练轮数、训练时间 - -## 应用最佳配置 - -网格搜索完成后: - -1. 查看`summary.txt`找到最佳配置 -2. 手动更新`train/train_cosface_embedding.py`中的`TASKS`字典: -```python -'dish': TaskConfig( - # ... 其他配置保持不变 ... - cosface_s=64.0, # 更新为最佳s值 - cosface_m=0.38, # 更新为最佳m值 -) -``` -3. 后续训练将自动使用最佳配置 - -## 预计耗时 - -基于以下假设: -- 每个配置平均训练25-40个epoch(早停机制) -- 每个epoch约1-2分钟 -- **单个配置**: ~30-60分钟 -- **16组配置总耗时**: ~8-16小时 - -**建议**: 使用GPU运行,可在夜间或周末执行完整搜索。 - -## 注意事项 - -1. **数据集要求**: 确保`dataset/DishClassification/test/`目录存在且有数据 -2. **GPU推荐**: 网格搜索计算量大,强烈建议使用GPU -3. **磁盘空间**: 每个配置约占用500MB,16组需8GB空间 -4. **中断恢复**: 当前版本不支持断点续训,建议一次性完成 - -## 高级用法 - -### 修改搜索空间 -编辑`train/grid_search_cosface.py`中的`GRID_PARAMS`字典: -```python -GRID_PARAMS = { - 'dish': { - 's': [60.0, 64.0, 68.0, 72.0], # 自定义scale范围 - 'm': [0.30, 0.35, 0.40, 0.45], # 自定义margin范围 - } -} -``` - -### 调整早停策略 -通过命令行参数调整: -```bash -python train/grid_search_cosface.py --task dish --patience 15 --min_epochs 30 -``` - -## 故障排除 - -**问题1**: `ModuleNotFoundError: No module named 'pandas'` -- 解决: `pip install pandas seaborn` - -**问题2**: 测试集目录不存在 -- 解决: 确认`dataset/DishClassification/test/`路径正确且有数据 - -**问题3**: CUDA out of memory -- 解决: 减小`batch_size`或在CPU上运行(速度较慢) - -## 示例结果解读 - -`summary.txt`示例: -``` -最佳配置: - s (scale) = 64.0 - m (margin) = 0.38 - 测试集准确率 = 98.50% - 验证集准确率 = 100.00% - 训练轮数 = 32 - 训练时间 = 45.3 分钟 -``` - -这表示s=64.0, m=0.38是最优组合,在测试集上达到98.50%准确率。 diff --git a/build_faiss_index.py b/build_faiss_index.py deleted file mode 100644 index d01fa90..0000000 --- a/build_faiss_index.py +++ /dev/null @@ -1,519 +0,0 @@ -""" -食物分类器 - FAISS索引构建器 -将训练集图片转换为特征向量并构建FAISS索引用于快速相似度检索 - -功能: -1. 加载训练好的ResNet50 Embedding模型 -2. 遍历训练集所有图片,提取特征向量 -3. 构建FAISS索引文件 -4. 创建标签映射文件 -5. 实现快速相似度检索功能 - -作者: zhangpu -日期: 2025/9/22 -""" - -import os -import json -import pickle -import numpy as np -import torch -from PIL import Image -from tqdm import tqdm -from datetime import datetime -from typing import List, Tuple, Dict, Optional -import faiss - -# 导入项目模块 -from net.resnet_embedding import create_resnet50_embedding -from settings import settings - - -class FAISSIndexBuilder: - """FAISS索引构建器""" - - def __init__(self, model_path: str, embedding_dim: int = 512): - """ - 初始化FAISS索引构建器 - - Args: - model_path: 训练好的模型路径 - embedding_dim: 特征向量维度 - """ - self.model_path = model_path - self.embedding_dim = embedding_dim - self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') - - # 初始化模型 - self.model = None - self.load_model() - - # 数据存储 - self.embeddings = [] - self.image_paths = [] - self.labels = [] - self.class_names = [] - self.class_to_idx = {} - self.idx_to_class = {} - - print(f"使用设备: {self.device}") - print(f"特征向量维度: {self.embedding_dim}") - - def load_model(self): - """加载训练好的ResNet50 Embedding模型""" - try: - print(f"正在加载模型: {self.model_path}") - - # 创建模型 - self.model = create_resnet50_embedding( - embedding_dim=self.embedding_dim, - pretrained=True, - use_internal_preprocess=False - ) - - # 加载权重 - checkpoint = torch.load(self.model_path, map_location=self.device) - - # 处理不同的保存格式 - if isinstance(checkpoint, dict): - if 'model_state_dict' in checkpoint: - self.model.load_state_dict(checkpoint['model_state_dict']) - elif 'state_dict' in checkpoint: - self.model.load_state_dict(checkpoint['state_dict']) - else: - self.model.load_state_dict(checkpoint) - else: - self.model.load_state_dict(checkpoint) - - # 移动到设备并设置为评估模式 - self.model.to(self.device) - self.model.eval() - - print("模型加载成功!") - - except Exception as e: - print(f"模型加载失败: {e}") - raise - - def scan_training_data(self, train_dir: str) -> Tuple[List[str], List[str], List[int]]: - """ - 扫描训练数据目录,获取所有图片路径和标签 - - Args: - train_dir: 训练数据目录路径 - - Returns: - Tuple[List[str], List[str], List[int]]: (图片路径列表, 类别名称列表, 标签索引列表) - """ - print(f"正在扫描训练数据目录: {train_dir}") - - image_paths = [] - class_names = [] - labels = [] - - # 获取所有类别目录 - class_dirs = [d for d in os.listdir(train_dir) - if os.path.isdir(os.path.join(train_dir, d))] - class_dirs.sort() # 确保顺序一致 - - print(f"发现类别: {class_dirs}") - - # 创建类别映射 - self.class_names = class_dirs - self.class_to_idx = {class_name: idx for idx, class_name in enumerate(class_dirs)} - self.idx_to_class = {idx: class_name for idx, class_name in enumerate(class_dirs)} - - # 扫描每个类别目录 - for class_name in class_dirs: - class_dir = os.path.join(train_dir, class_name) - class_idx = self.class_to_idx[class_name] - - # 获取该类别下的所有图片 - image_files = [f for f in os.listdir(class_dir) - if f.lower().endswith(('.jpg', '.jpeg', '.png', '.bmp'))] - - print(f"类别 '{class_name}': {len(image_files)} 张图片") - - for image_file in image_files: - image_path = os.path.join(class_dir, image_file) - image_paths.append(image_path) - class_names.append(class_name) - labels.append(class_idx) - - print(f"总计扫描到 {len(image_paths)} 张图片") - return image_paths, class_names, labels - - def extract_features_batch(self, image_paths: List[str], batch_size: int = 32) -> np.ndarray: - """ - 批量提取图片特征向量 - - Args: - image_paths: 图片路径列表 - batch_size: 批处理大小 - - Returns: - np.ndarray: 特征向量数组,形状为 (N, embedding_dim) - """ - print(f"开始提取特征向量,批大小: {batch_size}") - - all_embeddings = [] - - with torch.no_grad(): - for i in tqdm(range(0, len(image_paths), batch_size), desc="提取特征"): - batch_paths = image_paths[i:i + batch_size] - batch_images = [] - valid_indices = [] - - # 加载并预处理批次图片 - for j, img_path in enumerate(batch_paths): - try: - # 加载图片 - image = Image.open(img_path).convert('RGB') - - # 预处理 - tensor = self.model.preprocess_image(image) - batch_images.append(tensor.squeeze(0)) # 移除batch维度 - valid_indices.append(i + j) - - except Exception as e: - print(f"警告: 无法处理图片 {img_path}: {e}") - continue - - if not batch_images: - continue - - # 堆叠成批次并移动到设备 - batch_tensor = torch.stack(batch_images).to(self.device) - - # 提取特征 - batch_embeddings = self.model(batch_tensor, normalize=True) - - # 转换为numpy并添加到结果中 - all_embeddings.append(batch_embeddings.cpu().numpy()) - - # 合并所有批次的结果 - if all_embeddings: - embeddings = np.vstack(all_embeddings) - print(f"成功提取 {embeddings.shape[0]} 个特征向量") - return embeddings - else: - raise ValueError("没有成功提取任何特征向量") - - def build_faiss_index(self, embeddings: np.ndarray, index_type: str = 'flat') -> faiss.Index: - """ - 构建FAISS索引 - - Args: - embeddings: 特征向量数组 - index_type: 索引类型 ('flat', 'ivf', 'hnsw') - - Returns: - faiss.Index: 构建好的FAISS索引 - """ - print(f"构建FAISS索引,类型: {index_type}") - - dimension = embeddings.shape[1] - n_vectors = embeddings.shape[0] - - if index_type == 'flat': - # 精确搜索索引 - index = faiss.IndexFlatIP(dimension) # 内积索引(适合归一化向量) - - elif index_type == 'ivf': - # IVF索引(适合大规模数据) - nlist = min(100, max(1, n_vectors // 10)) # 聚类中心数量 - quantizer = faiss.IndexFlatIP(dimension) - index = faiss.IndexIVFFlat(quantizer, dimension, nlist) - - # 训练索引 - print("训练IVF索引...") - index.train(embeddings.astype(np.float32)) - - elif index_type == 'hnsw': - # HNSW索引(快速近似搜索) - M = 16 # 连接数 - index = faiss.IndexHNSWFlat(dimension, M) - index.hnsw.efConstruction = 200 - - else: - raise ValueError(f"不支持的索引类型: {index_type}") - - # 添加向量到索引 - print("添加向量到索引...") - index.add(embeddings.astype(np.float32)) - - print(f"索引构建完成,包含 {index.ntotal} 个向量") - return index - - def save_index_and_metadata(self, index: faiss.Index, output_dir: str): - """ - 保存FAISS索引和相关元数据 - - Args: - index: FAISS索引 - output_dir: 输出目录 - """ - # 创建输出目录 - os.makedirs(output_dir, exist_ok=True) - - # 保存FAISS索引 - index_path = os.path.join(output_dir, 'faiss_index.bin') - faiss.write_index(index, index_path) - print(f"FAISS索引已保存到: {index_path}") - - # 保存图片路径映射 - paths_path = os.path.join(output_dir, 'image_paths.pkl') - with open(paths_path, 'wb') as f: - pickle.dump(self.image_paths, f) - print(f"图片路径映射已保存到: {paths_path}") - - # 保存标签映射 - labels_path = os.path.join(output_dir, 'labels.pkl') - with open(labels_path, 'wb') as f: - pickle.dump(self.labels, f) - print(f"标签映射已保存到: {labels_path}") - - # 保存类别信息 - class_info = { - 'class_names': self.class_names, - 'class_to_idx': self.class_to_idx, - 'idx_to_class': self.idx_to_class, - 'embedding_dim': self.embedding_dim, - 'total_images': len(self.image_paths), - 'build_time': datetime.now().isoformat() - } - - class_info_path = os.path.join(output_dir, 'class_info.json') - with open(class_info_path, 'w', encoding='utf-8') as f: - json.dump(class_info, f, ensure_ascii=False, indent=2) - print(f"类别信息已保存到: {class_info_path}") - - # 保存特征向量(可选,用于调试) - embeddings_path = os.path.join(output_dir, 'embeddings.npy') - embeddings_array = np.array(self.embeddings) - np.save(embeddings_path, embeddings_array) - print(f"特征向量已保存到: {embeddings_path}") - - def build_complete_index(self, train_dir: str, output_dir: str, - batch_size: int = 32, index_type: str = 'flat'): - """ - 完整的索引构建流程 - - Args: - train_dir: 训练数据目录 - output_dir: 输出目录 - batch_size: 批处理大小 - index_type: 索引类型 - """ - print("=" * 60) - print("开始构建FAISS索引") - print("=" * 60) - - # 1. 扫描训练数据 - image_paths, class_names, labels = self.scan_training_data(train_dir) - self.image_paths = image_paths - self.labels = labels - - # 2. 提取特征向量 - embeddings = self.extract_features_batch(image_paths, batch_size) - self.embeddings = embeddings - - # 3. 构建FAISS索引 - index = self.build_faiss_index(embeddings, index_type) - - # 4. 保存索引和元数据 - self.save_index_and_metadata(index, output_dir) - - print("=" * 60) - print("FAISS索引构建完成!") - print(f"输出目录: {output_dir}") - print(f"总图片数: {len(image_paths)}") - print(f"类别数: {len(self.class_names)}") - print(f"特征维度: {self.embedding_dim}") - print("=" * 60) - - return index - - -class FAISSSearcher: - """FAISS相似度检索器""" - - def __init__(self, index_dir: str, model_path: str): - """ - 初始化FAISS检索器 - - Args: - index_dir: 索引文件目录 - model_path: 模型路径 - """ - self.index_dir = index_dir - self.model_path = model_path - self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') - - # 加载索引和元数据 - self.load_index_and_metadata() - - # 加载模型 - self.load_model() - - def load_index_and_metadata(self): - """加载FAISS索引和元数据""" - # 加载FAISS索引 - index_path = os.path.join(self.index_dir, 'faiss_index.bin') - self.index = faiss.read_index(index_path) - print(f"FAISS索引已加载: {self.index.ntotal} 个向量") - - # 加载图片路径 - paths_path = os.path.join(self.index_dir, 'image_paths.pkl') - with open(paths_path, 'rb') as f: - self.image_paths = pickle.load(f) - - # 加载标签 - labels_path = os.path.join(self.index_dir, 'labels.pkl') - with open(labels_path, 'rb') as f: - self.labels = pickle.load(f) - - # 加载类别信息 - class_info_path = os.path.join(self.index_dir, 'class_info.json') - with open(class_info_path, 'r', encoding='utf-8') as f: - self.class_info = json.load(f) - - self.class_names = self.class_info['class_names'] - self.class_to_idx = self.class_info['class_to_idx'] - self.idx_to_class = self.class_info['idx_to_class'] - self.embedding_dim = self.class_info['embedding_dim'] - - print(f"元数据已加载: {len(self.image_paths)} 张图片, {len(self.class_names)} 个类别") - - def load_model(self): - """加载特征提取模型""" - self.model = create_resnet50_embedding( - embedding_dim=self.embedding_dim, - pretrained=True, - use_internal_preprocess=False - ) - - checkpoint = torch.load(self.model_path, map_location=self.device) - if isinstance(checkpoint, dict): - if 'model_state_dict' in checkpoint: - self.model.load_state_dict(checkpoint['model_state_dict']) - elif 'state_dict' in checkpoint: - self.model.load_state_dict(checkpoint['state_dict']) - else: - self.model.load_state_dict(checkpoint) - else: - self.model.load_state_dict(checkpoint) - - self.model.to(self.device) - self.model.eval() - print("特征提取模型已加载") - - def search_similar(self, query_image_path: str, k: int = 5) -> List[Tuple[str, str, float]]: - """ - 搜索相似图片 - - Args: - query_image_path: 查询图片路径 - k: 返回最相似的k张图片 - - Returns: - List[Tuple[str, str, float]]: [(图片路径, 类别名称, 相似度分数), ...] - """ - # 提取查询图片的特征 - query_image = Image.open(query_image_path).convert('RGB') - query_embedding = self.model.extract_embedding(query_image, normalize=True) - query_embedding = query_embedding.reshape(1, -1).astype(np.float32) - - # 在FAISS索引中搜索 - scores, indices = self.index.search(query_embedding, k) - - # 整理结果 - results = [] - for i in range(k): - if i < len(indices[0]): - idx = indices[0][i] - score = scores[0][i] - image_path = self.image_paths[idx] - class_name = self.class_names[self.labels[idx]] - results.append((image_path, class_name, float(score))) - - return results - - def search_by_class(self, class_name: str, k: int = 10) -> List[Tuple[str, float]]: - """ - 按类别搜索图片 - - Args: - class_name: 类别名称 - k: 返回的图片数量 - - Returns: - List[Tuple[str, float]]: [(图片路径, 索引), ...] - """ - if class_name not in self.class_to_idx: - raise ValueError(f"未知类别: {class_name}") - - class_idx = self.class_to_idx[class_name] - class_images = [] - - for i, label in enumerate(self.labels): - if label == class_idx: - class_images.append((self.image_paths[i], i)) - if len(class_images) >= k: - break - - return class_images - - -def main(): - """主函数""" - # 配置参数 - MODEL_PATH = "model/embedding_20250930_102826/best_embedding_model.pth" - TRAIN_DIR = "dataset/train" - OUTPUT_DIR = "faiss_vector_db/DishClassification/faiss_index" - BATCH_SIZE = 16 - INDEX_TYPE = 'flat' # 'flat', 'ivf', 'hnsw' - EMBEDDING_DIM = 512 - - try: - # 构建索引 - builder = FAISSIndexBuilder(MODEL_PATH, EMBEDDING_DIM) - index = builder.build_complete_index( - train_dir=TRAIN_DIR, - output_dir=OUTPUT_DIR, - batch_size=BATCH_SIZE, - index_type=INDEX_TYPE - ) - - print("\n" + "=" * 60) - print("测试相似度检索功能") - print("=" * 60) - - # 测试检索功能 - searcher = FAISSSearcher(OUTPUT_DIR, MODEL_PATH) - - # 随机选择一张图片进行测试 - test_image_path = builder.image_paths[0] - print(f"测试图片: {test_image_path}") - - # 搜索相似图片 - similar_images = searcher.search_similar(test_image_path, k=5) - - print("\n最相似的5张图片:") - for i, (img_path, class_name, score) in enumerate(similar_images, 1): - print(f"{i}. {os.path.basename(img_path)} (类别: {class_name}, 相似度: {score:.4f})") - - # 按类别搜索 - print(f"\n'{builder.class_names[0]}'类别的前5张图片:") - class_images = searcher.search_by_class(builder.class_names[0], k=5) - for i, (img_path, idx) in enumerate(class_images, 1): - print(f"{i}. {os.path.basename(img_path)} (索引: {idx})") - - except Exception as e: - print(f"错误: {e}") - import traceback - traceback.print_exc() - - -if __name__ == "__main__": - main() \ No newline at end of file