From d51853997ff0ae14160a77665c304888131694ac Mon Sep 17 00:00:00 2001 From: zhanghuan <1262329256@qq.com> Date: Tue, 23 Sep 2025 15:34:23 +0800 Subject: [PATCH] =?UTF-8?q?=E5=BB=BA=E7=AB=8BFAISS=E7=B4=A2=E5=BC=95?= =?UTF-8?q?=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- build_faiss_index.py | 519 ++++++++++++++++++++ classifier/embedding_food_classifier_app.py | 4 +- net/resnet_embedding.py | 3 +- 3 files changed, 524 insertions(+), 2 deletions(-) create mode 100644 build_faiss_index.py diff --git a/build_faiss_index.py b/build_faiss_index.py new file mode 100644 index 0000000..3409bdd --- /dev/null +++ b/build_faiss_index.py @@ -0,0 +1,519 @@ +""" +食物分类器 - 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_20250917_145342/best_embedding_model.pth" + TRAIN_DIR = "dataset/train" + OUTPUT_DIR = "faiss_vector_db/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 diff --git a/classifier/embedding_food_classifier_app.py b/classifier/embedding_food_classifier_app.py index a0c2177..0a7a20b 100644 --- a/classifier/embedding_food_classifier_app.py +++ b/classifier/embedding_food_classifier_app.py @@ -69,6 +69,7 @@ class EmbeddingFoodClassifierApp: if os.path.exists(model_path) and os.path.exists(index_dir): # 1. 加载embedding模型 print("正在加载embedding模型...") + # 先把网络结构拿过来,传入embedding维度,预训练,除了加载网络结构,还把预训练的网络参数加载过来 self.model = create_resnet50_embedding( embedding_dim=512, pretrained=True, @@ -86,8 +87,9 @@ class EmbeddingFoodClassifierApp: self.model.load_state_dict(checkpoint) else: self.model.load_state_dict(checkpoint) - + # 把模型放到设备上(GPU或CPU) self.model.to(self.device) + # 模型在评估模式下运行时,会禁用某些操作,如BN层和Dropout层 self.model.eval() print("Embedding模型加载成功!") diff --git a/net/resnet_embedding.py b/net/resnet_embedding.py index 9cd2822..46c9752 100644 --- a/net/resnet_embedding.py +++ b/net/resnet_embedding.py @@ -153,7 +153,8 @@ class ResNet50EmbeddingNet(nn.Module): # 生成embedding embeddings = self._forward_embedding(backbone_features) - # L2归一化(用于余弦相似度计算) + # L2归一化(用于余弦相似度计算),p=2表示L2范数(即平方和再开根号),dim=1表示在第二个维度上做归一化 + # 形状是[batch_size, embedding_dim] if normalize: embeddings = F.normalize(embeddings, p=2, dim=1)