From be284442e4330fd5f350f8eafc21a8d31db2374f Mon Sep 17 00:00:00 2001 From: zhanghuan <1262329256@qq.com> Date: Mon, 22 Sep 2025 10:42:57 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0FAISS=E5=90=91=E9=87=8F?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E5=BA=93=E6=A3=80=E7=B4=A2=E5=AD=98=E5=82=A8?= =?UTF-8?q?=E7=AD=89=E5=8A=9F=E8=83=BD=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 2 + demo/向量索引.py | 40 ++ demo/看一下向量的数量.py | 20 + faiss_vector_db/README.md | 312 +++++++++++ .../__init__.py | 0 faiss_vector_db/example_usage.py | 223 ++++++++ faiss_vector_db/faiss_manager.py | 516 ++++++++++++++++++ faiss_vector_db/food_labels.json | 0 faiss_vector_db/test_faiss_manager.py | 459 ++++++++++++++++ requirements.txt | 1 + 10 files changed, 1573 insertions(+) create mode 100644 demo/向量索引.py create mode 100644 demo/看一下向量的数量.py create mode 100644 faiss_vector_db/README.md rename faiss_index/food_labels.json => faiss_vector_db/__init__.py (100%) create mode 100644 faiss_vector_db/example_usage.py create mode 100644 faiss_vector_db/faiss_manager.py create mode 100644 faiss_vector_db/food_labels.json create mode 100644 faiss_vector_db/test_faiss_manager.py diff --git a/.gitignore b/.gitignore index ba0d2b7..d164478 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,5 @@ /dataset/ /.idea/ /model/ +/faiss_vector_db/demo_faiss_index/ +/faiss_vector_db/faiss_index/ diff --git a/demo/向量索引.py b/demo/向量索引.py new file mode 100644 index 0000000..2be0a51 --- /dev/null +++ b/demo/向量索引.py @@ -0,0 +1,40 @@ +import faiss +import numpy as np + +# 数据归一化函数 +def normalize_vectors(vectors): + """对向量进行L2归一化""" + norms = np.linalg.norm(vectors, axis=1, keepdims=True) + # 避免除零 + norms = np.where(norms == 0, 1, norms) + return vectors / norms + + +data = np.array([[2, 3], [2, 4], [3, 7]], dtype='float32') + +# 归一化数据 +data_normalized = normalize_vectors(data) +print("原始数据:") +print(data) +print("归一化后数据:") +print(data_normalized) + +# 普通索引 +# base_index = faiss.IndexFlatL2(2) +base_index = faiss.IndexFlatIP(2) + +# 包一层 IDMap +index = faiss.IndexIDMap(base_index) + +# 指定 ID +ids = np.array([101, 102, 103]) # 自定义 ID +# index.add_with_ids(data, ids) +index.add_with_ids(data_normalized, ids) + +# 查询 +query = np.array([[3, 4.5]], dtype='float32') +query_normalized = normalize_vectors(query) +# D, I = index.search(query, k=2) +D, I = index.search(query_normalized, k=2) +print(D) +print(I) # 可能输出 [[101 102]] diff --git a/demo/看一下向量的数量.py b/demo/看一下向量的数量.py new file mode 100644 index 0000000..34ea4f0 --- /dev/null +++ b/demo/看一下向量的数量.py @@ -0,0 +1,20 @@ +import faiss +import numpy as np + +# 建一个 2 维向量的 L2 索引 +index = faiss.IndexFlatL2(2) + +print(index.ntotal) # 初始是 0 + +# 插入 5 个向量 +data = np.random.rand(5, 2).astype("float32") +index.add(data) + +print(index.ntotal) # 现在是 5 + +# 再插入 3 个 +more_data = np.random.rand(3, 2).astype("float32") +index.add(more_data) + +print(index.ntotal) # 现在是 8 + diff --git a/faiss_vector_db/README.md b/faiss_vector_db/README.md new file mode 100644 index 0000000..b37531f --- /dev/null +++ b/faiss_vector_db/README.md @@ -0,0 +1,312 @@ +# FAISS向量管理器 + +这是一个基于FAISS的向量管理器,专为食物分类项目设计,提供了完整的向量索引管理功能。 + +## 功能特性 + +- ✅ **FAISS索引管理**: 支持创建、保存和加载FAISS索引 +- ✅ **向量操作**: 支持向量的增删改查操作 +- ✅ **相似度搜索**: 支持余弦相似度搜索和阈值过滤 +- ✅ **批量操作**: 提供高效的批量向量操作接口 +- ✅ **元数据管理**: 支持向量元数据的存储和检索 +- ✅ **多种索引类型**: 支持不同的FAISS索引类型 +- ✅ **持久化存储**: 支持索引的保存和加载 + +## 安装依赖 + +```bash +pip install faiss-cpu numpy +# 或者如果需要GPU支持 +pip install faiss-gpu numpy +``` + +## 快速开始 + +### 1. 基本使用 + +```python +from faiss.faiss_manager import FAISSManager +import numpy as np + +# 创建管理器 +manager = FAISSManager( + dimension=512, # 向量维度 + index_type="IndexFlatIP", # 索引类型 + normalize_vectors=True, # 启用向量标准化 + metric_type="cosine" # 使用余弦相似度 +) + +# 添加单个向量 +vector = np.random.random(512).astype(np.float32) +manager.add_vector( + vector=vector, + vector_id="food_001", + metadata={"name": "回锅肉", "category": "川菜", "spicy_level": 3} +) + +# 搜索相似向量 +query_vector = np.random.random(512).astype(np.float32) +results = manager.search_similar(query_vector, k=5) + +for vector_id, similarity, metadata in results: + print(f"ID: {vector_id}, 相似度: {similarity:.4f}, 菜名: {metadata['name']}") +``` + +### 2. 批量操作 + +```python +# 批量添加向量 +batch_vectors = np.random.random((100, 512)).astype(np.float32) +batch_ids = [f"food_{i:03d}" for i in range(100)] +batch_metadata = [ + {"name": f"菜品{i}", "category": "川菜", "spicy_level": i % 5} + for i in range(100) +] + +results = manager.add_vectors_batch(batch_vectors, batch_ids, batch_metadata) +print(f"成功添加: {sum(results)}/{len(results)} 个向量") + +# 批量删除向量 +delete_ids = [f"food_{i:03d}" for i in range(10)] +delete_results = manager.delete_vectors_batch(delete_ids) +print(f"成功删除: {sum(delete_results)}/{len(delete_results)} 个向量") +``` + +### 3. 索引保存和加载 + +```python +# 保存索引 +save_path = "./my_faiss_index" +success = manager.save_index(save_path) +if success: + print("索引保存成功") + +# 加载索引 +new_manager = FAISSManager(dimension=512) +success = new_manager.load_index(save_path) +if success: + print("索引加载成功") + + # 验证加载结果 + info = new_manager.get_index_info() + print(f"加载的向量数量: {info['active_vectors']}") +``` + +## API 参考 + +### FAISSManager 类 + +#### 初始化参数 + +- `dimension` (int): 向量维度 +- `index_type` (str): 索引类型,支持: + - `"IndexFlatIP"`: 内积索引(适合余弦相似度) + - `"IndexFlatL2"`: L2距离索引 + - `"IndexIVFFlat"`: IVF索引(适合大规模数据) + - `"IndexHNSW"`: HNSW索引(高性能近似搜索) +- `normalize_vectors` (bool): 是否标准化向量 +- `metric_type` (str): 距离度量类型 ("cosine", "l2", "ip") + +#### 主要方法 + +##### 向量操作 + +```python +# 添加单个向量 +add_vector(vector, vector_id, metadata=None) -> bool + +# 批量添加向量 +add_vectors_batch(vectors, vector_ids, metadata_list=None) -> List[bool] + +# 更新向量 +update_vector(vector, vector_id, metadata=None) -> bool + +# 删除单个向量 +delete_vector(vector_id) -> bool + +# 批量删除向量 +delete_vectors_batch(vector_ids) -> List[bool] +``` + +##### 搜索操作 + +```python +# 相似度搜索 +search_similar(query_vector, k=10, threshold=None) -> List[Tuple[str, float, Dict]] + +# 根据ID获取向量信息 +get_vector_by_id(vector_id) -> Optional[Tuple[np.ndarray, Dict]] +``` + +##### 索引管理 + +```python +# 保存索引 +save_index(save_path) -> bool + +# 加载索引 +load_index(load_path) -> bool + +# 清空索引 +clear_index() -> bool + +# 获取索引信息 +get_index_info() -> Dict[str, Any] +``` + +## 使用场景 + +### 1. 食物图像检索 + +```python +# 为食物分类项目设计的示例 +manager = FAISSManager(dimension=2048, index_type="IndexFlatIP") + +# 添加食物向量(通过CNN模型提取的特征) +food_features = extract_features_from_images(food_images) # 假设的特征提取函数 +food_metadata = [ + {"name": "回锅肉", "category": "川菜", "ingredients": ["猪肉", "青椒", "豆瓣酱"]}, + {"name": "西红柿鸡蛋", "category": "家常菜", "ingredients": ["西红柿", "鸡蛋"]}, + # ... 更多食物数据 +] + +manager.add_vectors_batch(food_features, food_ids, food_metadata) + +# 查询相似食物 +query_image_feature = extract_features_from_image(query_image) +similar_foods = manager.search_similar(query_image_feature, k=5, threshold=0.7) + +for food_id, similarity, metadata in similar_foods: + print(f"相似食物: {metadata['name']} (相似度: {similarity:.3f})") +``` + +### 2. 文本向量检索 + +```python +# 用于文本相似度检索 +manager = FAISSManager(dimension=768, index_type="IndexFlatIP") + +# 添加文本向量(通过BERT等模型编码) +text_embeddings = encode_texts(recipe_texts) # 假设的文本编码函数 +recipe_metadata = [ + {"title": "回锅肉制作方法", "difficulty": "中等", "time": "30分钟"}, + # ... 更多菜谱数据 +] + +manager.add_vectors_batch(text_embeddings, recipe_ids, recipe_metadata) + +# 搜索相关菜谱 +query_embedding = encode_text("如何做川菜") +related_recipes = manager.search_similar(query_embedding, k=3) +``` + +## 性能优化建议 + +### 1. 索引类型选择 + +- **小规模数据 (<10K向量)**: 使用 `IndexFlatIP` 或 `IndexFlatL2` +- **中等规模数据 (10K-1M向量)**: 使用 `IndexIVFFlat` +- **大规模数据 (>1M向量)**: 使用 `IndexHNSW` 或更复杂的索引 + +### 2. 内存优化 + +```python +# 对于大规模数据,考虑使用IVF索引 +manager = FAISSManager( + dimension=512, + index_type="IndexIVFFlat", + normalize_vectors=True +) + +# 训练IVF索引(需要足够的训练数据) +if hasattr(manager.index, 'train') and not manager.index.is_trained: + training_vectors = np.random.random((10000, 512)).astype(np.float32) + manager.index.train(training_vectors) +``` + +### 3. 批量操作 + +```python +# 优先使用批量操作而不是循环调用单个操作 +# 好的做法 +manager.add_vectors_batch(vectors, ids, metadata_list) + +# 避免的做法 +for vector, id, metadata in zip(vectors, ids, metadata_list): + manager.add_vector(vector, id, metadata) +``` + +## 测试 + +运行测试以验证功能: + +```bash +# 运行基本测试 +python faiss/test_faiss_manager.py + +# 运行使用示例 +python faiss/example_usage.py +``` + +## 注意事项 + +1. **向量维度**: 所有向量必须具有相同的维度 +2. **ID唯一性**: 向量ID必须唯一,重复ID会导致更新操作 +3. **内存使用**: FAISS索引会占用内存,大规模数据需要考虑内存限制 +4. **删除操作**: FAISS不支持真正的删除,删除操作只是标记,需要重建索引来释放空间 +5. **向量标准化**: 使用余弦相似度时建议启用向量标准化 + +## 故障排除 + +### 常见问题 + +1. **维度不匹配错误** + ```python + # 确保所有向量维度一致 + assert vector.shape[-1] == manager.dimension + ``` + +2. **索引未训练错误** + ```python + # 对于IVF索引,需要先训练 + if hasattr(manager.index, 'train') and not manager.index.is_trained: + manager.index.train(training_data) + ``` + +3. **内存不足** + ```python + # 使用更节省内存的索引类型 + manager = FAISSManager(dimension=512, index_type="IndexIVFFlat") + ``` + +## 扩展功能 + +### 自定义距离度量 + +```python +class CustomFAISSManager(FAISSManager): + def custom_similarity_function(self, query_vector, k=10): + """自定义相似度计算""" + # 实现自定义逻辑 + pass +``` + +### 多模态检索 + +```python +# 结合图像和文本特征 +image_manager = FAISSManager(dimension=2048, index_type="IndexFlatIP") +text_manager = FAISSManager(dimension=768, index_type="IndexFlatIP") + +# 融合检索结果 +def multimodal_search(image_query, text_query, alpha=0.7): + image_results = image_manager.search_similar(image_query, k=20) + text_results = text_manager.search_similar(text_query, k=20) + + # 融合结果逻辑 + # ... +``` + +## 许可证 + +本项目采用 MIT 许可证。 \ No newline at end of file diff --git a/faiss_index/food_labels.json b/faiss_vector_db/__init__.py similarity index 100% rename from faiss_index/food_labels.json rename to faiss_vector_db/__init__.py diff --git a/faiss_vector_db/example_usage.py b/faiss_vector_db/example_usage.py new file mode 100644 index 0000000..ec3eb72 --- /dev/null +++ b/faiss_vector_db/example_usage.py @@ -0,0 +1,223 @@ +""" +FAISS向量管理器使用示例 +演示如何使用FAISSManager进行向量操作 +""" + +import numpy as np +import sys +import os +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from faiss_vector_db.faiss_manager import FAISSManager + +def demo_basic_operations(): + """演示基本操作""" + print("=== FAISS向量管理器基本操作演示 ===\n") + + # 1. 创建管理器 + print("1. 创建FAISS管理器...") + manager = FAISSManager( + dimension=128, # 向量维度 + index_type="IndexFlatIP", # 使用内积索引(适合余弦相似度) + normalize_vectors=True, # 标准化向量 + metric_type="cosine" # 余弦相似度 + ) + print(f" 管理器创建成功,维度: {manager.dimension}") + + # 2. 添加单个向量 + print("\n2. 添加单个向量...") + vector1 = np.random.random(128).astype(np.float32) + success = manager.add_vector( + vector=vector1, + vector_id="food_001", + metadata={"name": "回锅肉", "category": "川菜", "spicy_level": 3} + ) + print(f" 添加结果: {'成功' if success else '失败'}") + + # 3. 批量添加向量 + print("\n3. 批量添加向量...") + batch_vectors = np.random.random((5, 128)).astype(np.float32) + # 向量ID + batch_ids = ["food_002", "food_003", "food_004", "food_005", "food_006"] + # 向量元数据 + batch_metadata = [ + {"name": "炒细面", "category": "川菜", "spicy_level": 2}, + {"name": "西红柿鸡蛋", "category": "家常菜", "spicy_level": 0}, + {"name": "麻辣小面", "category": "川菜", "spicy_level": 4}, + {"name": "宫保鸡丁", "category": "川菜", "spicy_level": 3}, + {"name": "糖醋里脊", "category": "鲁菜", "spicy_level": 0} + ] + + results = manager.add_vectors_batch(batch_vectors, batch_ids, batch_metadata) + success_count = sum(results) + print(f" 批量添加结果: {success_count}/{len(batch_ids)} 成功") + + # 4. 查看索引信息 + print("\n4. 索引信息:") + info = manager.get_index_info() + for key, value in info.items(): + print(f" {key}: {value}") + + return manager + +def demo_search_operations(manager): + """演示搜索操作""" + print("\n=== 搜索操作演示 ===\n") + + # 1. 相似度搜索 + print("1. 相似度搜索...") + query_vector = np.random.random(128).astype(np.float32) + + # 搜索最相似的3个向量 + results = manager.search_similar(query_vector, k=3) + print(f" 找到 {len(results)} 个相似向量:") + for i, (vector_id, similarity, metadata) in enumerate(results, 1): + print(f" {i}. ID: {vector_id}") + print(f" 相似度: {similarity:.4f}") + print(f" 菜名: {metadata.get('name', 'N/A')}") + print(f" 类别: {metadata.get('category', 'N/A')}") + print(f" 辣度: {metadata.get('spicy_level', 'N/A')}") + print() + + # 2. 带阈值的搜索 + print("2. 带阈值的搜索(相似度 > 0.5)...") + results_with_threshold = manager.search_similar(query_vector, k=10, threshold=0.5) + print(f" 找到 {len(results_with_threshold)} 个高相似度向量") + + # 3. 根据ID获取向量信息 + print("\n3. 根据ID获取向量信息...") + vector_info = manager.get_vector_by_id("food_001") + if vector_info: + vector_data, metadata = vector_info + print(f" ID: food_001") + print(f" 元数据: {metadata}") + else: + print(" 向量不存在") + +def demo_update_delete_operations(manager): + """演示更新和删除操作""" + print("\n=== 更新和删除操作演示 ===\n") + + # 1. 更新向量 + print("1. 更新向量...") + new_vector = np.random.random(128).astype(np.float32) + new_metadata = {"name": "回锅肉(改良版)", "category": "川菜", "spicy_level": 2} + success = manager.update_vector(new_vector, "food_001", new_metadata) + print(f" 更新结果: {'成功' if success else '失败'}") + + # 验证更新 + vector_info = manager.get_vector_by_id("food_001") + if vector_info: + _, metadata = vector_info + print(f" 更新后的元数据: {metadata}") + + # 2. 删除单个向量 + print("\n2. 删除单个向量...") + success = manager.delete_vector("food_006") + print(f" 删除结果: {'成功' if success else '失败'}") + + # 3. 批量删除向量 + print("\n3. 批量删除向量...") + delete_ids = ["food_004", "food_005"] + results = manager.delete_vectors_batch(delete_ids) + success_count = sum(results) + print(f" 批量删除结果: {success_count}/{len(delete_ids)} 成功") + + # 4. 查看删除后的索引信息 + print("\n4. 删除后的索引信息:") + info = manager.get_index_info() + for key, value in info.items(): + print(f" {key}: {value}") + +def demo_save_load_operations(manager): + """演示保存和加载操作""" + print("\n=== 保存和加载操作演示 ===\n") + + # 1. 保存索引 + print("1. 保存索引...") + save_path = "./demo_faiss_index" + success = manager.save_index(save_path) + print(f" 保存结果: {'成功' if success else '失败'}") + + # 2. 创建新的管理器并加载索引 + print("\n2. 加载索引到新管理器...") + new_manager = FAISSManager(dimension=128) + success = new_manager.load_index(save_path) + print(f" 加载结果: {'成功' if success else '失败'}") + + # 3. 验证加载的索引 + print("\n3. 验证加载的索引...") + info = new_manager.get_index_info() + print(f" 加载后的向量数量: {info['active_vectors']}") + + # 4. 测试加载后的搜索功能 + print("\n4. 测试加载后的搜索功能...") + query_vector = np.random.random(128).astype(np.float32) + results = new_manager.search_similar(query_vector, k=2) + print(f" 搜索到 {len(results)} 个结果") + for vector_id, similarity, metadata in results: + print(f" - {vector_id}: {metadata.get('name', 'N/A')} (相似度: {similarity:.4f})") + + return new_manager + +def demo_advanced_features(): + """演示高级功能""" + print("\n=== 高级功能演示 ===\n") + + # 1. 不同索引类型的比较 + print("1. 不同索引类型的性能比较...") + + # 创建测试数据 + test_vectors = np.random.random((1000, 64)).astype(np.float32) + test_ids = [f"test_{i}" for i in range(1000)] + + index_types = ["IndexFlatIP", "IndexFlatL2"] + + for index_type in index_types: + print(f"\n 测试索引类型: {index_type}") + manager = FAISSManager(dimension=64, index_type=index_type) + + # 批量添加 + import time + start_time = time.time() + manager.add_vectors_batch(test_vectors, test_ids) + add_time = time.time() - start_time + + # 搜索测试 + query = np.random.random(64).astype(np.float32) + start_time = time.time() + results = manager.search_similar(query, k=10) + search_time = time.time() - start_time + + print(f" - 添加1000个向量耗时: {add_time:.4f}秒") + print(f" - 搜索耗时: {search_time:.6f}秒") + print(f" - 找到结果数: {len(results)}") + +def main(): + """主函数""" + try: + # 基本操作演示 + manager = demo_basic_operations() + + # 搜索操作演示 + demo_search_operations(manager) + + # 更新删除操作演示 + demo_update_delete_operations(manager) + + # 保存加载操作演示 + loaded_manager = demo_save_load_operations(manager) + + # 高级功能演示 + demo_advanced_features() + + print("\n=== 演示完成 ===") + print("所有功能测试通过!") + + 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/faiss_vector_db/faiss_manager.py b/faiss_vector_db/faiss_manager.py new file mode 100644 index 0000000..3a65ed8 --- /dev/null +++ b/faiss_vector_db/faiss_manager.py @@ -0,0 +1,516 @@ +import faiss +import numpy as np +import pickle +import os +from typing import List, Tuple, Optional, Union, Dict, Any +import logging +from pathlib import Path +import json +from datetime import datetime + +class FAISSManager: + """ + FAISS向量管理器 + 支持向量的增删改查、余弦相似度搜索、批量操作等功能 + """ + + def __init__(self, dimension: int, index_type: str = "IndexFlatIP", + normalize_vectors: bool = True, metric_type: str = "cosine"): + """ + 初始化FAISS管理器 + + Args: + dimension: 向量维度 + index_type: 索引类型 ("IndexFlatIP", "IndexFlatL2", "IndexIVFFlat", "IndexHNSW") + normalize_vectors: 是否标准化向量(用于余弦相似度) + metric_type: 距离度量类型 ("cosine", "l2", "ip") + """ + self.dimension = dimension + self.index_type = index_type + self.normalize_vectors = normalize_vectors + self.metric_type = metric_type + + # 初始化索引 + self.index = self._create_index() + + # 存储向量ID到实际ID的映射 + self.id_mapping = {} # faiss_id -> actual_id + self.reverse_id_mapping = {} # actual_id -> faiss_id + self.next_faiss_id = 0 + + # 存储向量元数据 + self.metadata = {} # actual_id -> metadata + + # 日志设置 + logging.basicConfig(level=logging.INFO) + self.logger = logging.getLogger(__name__) + + def _create_index(self) -> 'faiss.Index': + """创建FAISS索引""" + if self.index_type == "IndexFlatIP": + # 内积索引(适合余弦相似度,需要标准化向量) + # 初始化的时候需要告诉维度 + index = faiss.IndexFlatIP(self.dimension) + elif self.index_type == "IndexFlatL2": + # L2距离索引 + index = faiss.IndexFlatL2(self.dimension) + elif self.index_type == "IndexIVFFlat": + # IVF索引(适合大规模数据) + quantizer = faiss.IndexFlatIP(self.dimension) if self.metric_type == "cosine" else faiss.IndexFlatL2(self.dimension) + nlist = 100 # 聚类中心数量 + index = faiss.IndexIVFFlat(quantizer, self.dimension, nlist) + elif self.index_type == "IndexHNSW": + # HNSW索引(高性能近似搜索) + index = faiss.IndexHNSWFlat(self.dimension, 32) + else: + raise ValueError(f"不支持的索引类型: {self.index_type}") + + return index + + def _normalize_vector(self, vector: np.ndarray) -> np.ndarray: + """标准化向量(用于余弦相似度)""" + if self.normalize_vectors: + norm = np.linalg.norm(vector, axis=-1, keepdims=True) + # 避免除零 + norm = np.where(norm == 0, 1, norm) + return vector / norm + return vector + + def add_vector(self, vector: Union[np.ndarray, List[float]], + vector_id: str, metadata: Optional[Dict[str, Any]] = None) -> bool: + """ + 添加单个向量 + + Args: + vector: 向量数据 + vector_id: 向量唯一标识 + metadata: 向量元数据 + + Returns: + bool: 是否添加成功 + """ + try: + # 检查向量是否已存在 + if vector_id in self.reverse_id_mapping: + self.logger.warning(f"向量ID {vector_id} 已存在,将更新该向量") + return self.update_vector(vector, vector_id, metadata) + + # 转换为numpy数组并标准化 + vector = np.array(vector, dtype=np.float32).reshape(1, -1) + if vector.shape[1] != self.dimension: + raise ValueError(f"向量维度不匹配: 期望 {self.dimension}, 实际 {vector.shape[1]}") + + vector = self._normalize_vector(vector) + + # 添加到索引 + self.index.add(vector) + + # 更新映射关系 + faiss_id = self.next_faiss_id + self.id_mapping[faiss_id] = vector_id + self.reverse_id_mapping[vector_id] = faiss_id + self.next_faiss_id += 1 + + # 存储元数据 + if metadata: + self.metadata[vector_id] = metadata + + self.logger.info(f"成功添加向量: {vector_id}") + return True + + except Exception as e: + self.logger.error(f"添加向量失败: {e}") + return False + + def add_vectors_batch(self, vectors: Union[np.ndarray, List[List[float]]], + vector_ids: List[str], + metadata_list: Optional[List[Dict[str, Any]]] = None) -> List[bool]: + """ + 批量添加向量 + + Args: + vectors: 向量数据矩阵 + vector_ids: 向量ID列表 + metadata_list: 元数据列表 + + Returns: + List[bool]: 每个向量的添加结果 + """ + results = [] + + # 转换为numpy数组 + vectors = np.array(vectors, dtype=np.float32) + if vectors.ndim == 1: + vectors = vectors.reshape(1, -1) + + if vectors.shape[1] != self.dimension: + raise ValueError(f"向量维度不匹配: 期望 {self.dimension}, 实际 {vectors.shape[1]}") + + if len(vector_ids) != vectors.shape[0]: + raise ValueError("向量数量与ID数量不匹配") + + # 标准化向量 + vectors = self._normalize_vector(vectors) + + # 检查重复ID + new_vectors = [] + new_ids = [] + new_metadata = [] + + for i, vector_id in enumerate(vector_ids): + if vector_id not in self.reverse_id_mapping: + new_vectors.append(vectors[i]) + new_ids.append(vector_id) + if metadata_list: + new_metadata.append(metadata_list[i]) + results.append(True) + else: + self.logger.warning(f"向量ID {vector_id} 已存在,跳过") + results.append(False) + + if new_vectors: + # 批量添加到索引 + new_vectors = np.array(new_vectors) + self.index.add(new_vectors) + + # 批量更新映射关系 + for i, vector_id in enumerate(new_ids): + faiss_id = self.next_faiss_id + i + self.id_mapping[faiss_id] = vector_id + self.reverse_id_mapping[vector_id] = faiss_id + + # 存储元数据 + if new_metadata and i < len(new_metadata): + self.metadata[vector_id] = new_metadata[i] + + self.next_faiss_id += len(new_ids) + self.logger.info(f"批量添加了 {len(new_ids)} 个向量") + + return results + + def search_similar(self, query_vector: Union[np.ndarray, List[float]], + k: int = 10, threshold: Optional[float] = None) -> List[Tuple[str, float, Dict[str, Any]]]: + """ + 搜索相似向量 + + Args: + query_vector: 查询向量 + k: 返回最相似的k个结果 + threshold: 相似度阈值 + + Returns: + List[Tuple[str, float, Dict]]: (vector_id, similarity_score, metadata) + """ + if self.index.ntotal == 0: + return [] + + # 转换并标准化查询向量 + query_vector = np.array(query_vector, dtype=np.float32).reshape(1, -1) + if query_vector.shape[1] != self.dimension: + raise ValueError(f"查询向量维度不匹配: 期望 {self.dimension}, 实际 {query_vector.shape[1]}") + + query_vector = self._normalize_vector(query_vector) + + # 搜索 + k = min(k, self.index.ntotal) + distances, indices = self.index.search(query_vector, k) + + results = [] + for i in range(k): + faiss_id = indices[0][i] + distance = distances[0][i] + + if faiss_id == -1: # 无效结果 + continue + + # 转换距离为相似度分数 + if self.metric_type == "cosine" or self.index_type == "IndexFlatIP": + similarity = float(distance) # 内积已经是相似度 + else: + similarity = 1.0 / (1.0 + float(distance)) # L2距离转相似度 + + # 应用阈值过滤 + if threshold is not None and similarity < threshold: + continue + + vector_id = self.id_mapping.get(faiss_id) + if vector_id: + metadata = self.metadata.get(vector_id, {}) + results.append((vector_id, similarity, metadata)) + + return results + + def get_vector_by_id(self, vector_id: str) -> Optional[Tuple[np.ndarray, Dict[str, Any]]]: + """ + 根据ID获取向量 + + Args: + vector_id: 向量ID + + Returns: + Optional[Tuple[np.ndarray, Dict]]: (vector, metadata) 或 None + """ + faiss_id = self.reverse_id_mapping.get(vector_id) + if faiss_id is None: + return None + + # FAISS不直接支持根据ID获取向量,这里返回元数据 + # 如果需要向量数据,建议单独存储 + metadata = self.metadata.get(vector_id, {}) + return None, metadata + + def update_vector(self, vector: Union[np.ndarray, List[float]], + vector_id: str, metadata: Optional[Dict[str, Any]] = None) -> bool: + """ + 更新向量(先删除再添加) + + Args: + vector: 新向量数据 + vector_id: 向量ID + metadata: 新元数据 + + Returns: + bool: 是否更新成功 + """ + # FAISS不支持直接更新,需要重建索引 + if vector_id not in self.reverse_id_mapping: + self.logger.warning(f"向量ID {vector_id} 不存在") + return False + + # 删除旧向量 + self.delete_vector(vector_id) + + # 添加新向量 + return self.add_vector(vector, vector_id, metadata) + + def delete_vector(self, vector_id: str) -> bool: + """ + 删除向量 + + Args: + vector_id: 向量ID + + Returns: + bool: 是否删除成功 + """ + if vector_id not in self.reverse_id_mapping: + self.logger.warning(f"向量ID {vector_id} 不存在") + return False + + # FAISS不支持直接删除,标记为删除 + faiss_id = self.reverse_id_mapping[vector_id] + + # 从映射中移除 + del self.id_mapping[faiss_id] + del self.reverse_id_mapping[vector_id] + + # 删除元数据 + if vector_id in self.metadata: + del self.metadata[vector_id] + + self.logger.info(f"标记删除向量: {vector_id}") + return True + + def delete_vectors_batch(self, vector_ids: List[str]) -> List[bool]: + """ + 批量删除向量 + + Args: + vector_ids: 向量ID列表 + + Returns: + List[bool]: 每个向量的删除结果 + """ + results = [] + for vector_id in vector_ids: + results.append(self.delete_vector(vector_id)) + return results + + def rebuild_index(self) -> bool: + """ + 重建索引(清理已删除的向量) + """ + try: + # 获取所有有效向量的信息 + valid_vectors = [] + valid_ids = [] + valid_metadata = {} + + # 这里需要外部提供向量数据,因为FAISS不能直接获取 + # 建议在实际使用时维护一个向量存储 + self.logger.warning("重建索引需要外部提供向量数据") + return False + + except Exception as e: + self.logger.error(f"重建索引失败: {e}") + return False + + def get_index_info(self) -> Dict[str, Any]: + """ + 获取索引信息 + + Returns: + Dict: 索引统计信息 + """ + return { + "total_vectors": self.index.ntotal, + "dimension": self.dimension, + "index_type": self.index_type, + "metric_type": self.metric_type, + "normalize_vectors": self.normalize_vectors, + "active_vectors": len(self.reverse_id_mapping), + "metadata_count": len(self.metadata) + } + + def save_index(self, save_path: str) -> bool: + """ + 保存索引到文件 + + Args: + save_path: 保存路径 + + Returns: + bool: 是否保存成功 + """ + try: + save_path = Path(save_path) + save_path.mkdir(parents=True, exist_ok=True) + + # 保存FAISS索引 + index_file = save_path / "faiss_index.bin" + faiss.write_index(self.index, str(index_file)) + + # 保存映射关系和元数据 + metadata_file = save_path / "metadata.pkl" + with open(metadata_file, 'wb') as f: + pickle.dump({ + 'id_mapping': self.id_mapping, + 'reverse_id_mapping': self.reverse_id_mapping, + 'next_faiss_id': self.next_faiss_id, + 'metadata': self.metadata, + 'dimension': self.dimension, + 'index_type': self.index_type, + 'normalize_vectors': self.normalize_vectors, + 'metric_type': self.metric_type + }, f) + + # 保存配置信息 + config_file = save_path / "config.json" + with open(config_file, 'w', encoding='utf-8') as f: + json.dump({ + 'dimension': self.dimension, + 'index_type': self.index_type, + 'normalize_vectors': self.normalize_vectors, + 'metric_type': self.metric_type, + 'save_time': datetime.now().isoformat(), + 'total_vectors': self.index.ntotal + }, f, indent=2, ensure_ascii=False) + + self.logger.info(f"索引已保存到: {save_path}") + return True + + except Exception as e: + self.logger.error(f"保存索引失败: {e}") + return False + + def load_index(self, load_path: str) -> bool: + """ + 从文件加载索引 + + Args: + load_path: 加载路径 + + Returns: + bool: 是否加载成功 + """ + try: + load_path = Path(load_path) + + # 检查文件是否存在 + index_file = load_path / "faiss_index.bin" + metadata_file = load_path / "metadata.pkl" + + if not index_file.exists() or not metadata_file.exists(): + self.logger.error(f"索引文件不存在: {load_path}") + return False + + # 加载FAISS索引 + self.index = faiss.read_index(str(index_file)) + + # 加载映射关系和元数据 + with open(metadata_file, 'rb') as f: + data = pickle.load(f) + self.id_mapping = data['id_mapping'] + self.reverse_id_mapping = data['reverse_id_mapping'] + self.next_faiss_id = data['next_faiss_id'] + self.metadata = data['metadata'] + self.dimension = data['dimension'] + self.index_type = data['index_type'] + self.normalize_vectors = data['normalize_vectors'] + self.metric_type = data['metric_type'] + + self.logger.info(f"索引已从 {load_path} 加载") + return True + + except Exception as e: + self.logger.error(f"加载索引失败: {e}") + return False + + def clear_index(self) -> bool: + """ + 清空索引 + + Returns: + bool: 是否清空成功 + """ + try: + # 重新创建索引 + self.index = self._create_index() + + # 清空映射关系 + self.id_mapping.clear() + self.reverse_id_mapping.clear() + self.next_faiss_id = 0 + + # 清空元数据 + self.metadata.clear() + + self.logger.info("索引已清空") + return True + + except Exception as e: + self.logger.error(f"清空索引失败: {e}") + return False + + +# 使用示例 +if __name__ == "__main__": + # 创建FAISS管理器 + manager = FAISSManager(dimension=512, index_type="IndexFlatIP") + + # 添加向量 + vector1 = np.random.random(512).astype(np.float32) + manager.add_vector(vector1, "vec_1", {"label": "食物1", "category": "川菜"}) + + # 批量添加向量 + vectors = np.random.random((10, 512)).astype(np.float32) + ids = [f"vec_{i}" for i in range(2, 12)] + metadata_list = [{"label": f"食物{i}", "category": "川菜"} for i in range(2, 12)] + manager.add_vectors_batch(vectors, ids, metadata_list) + + # 搜索相似向量 + query = np.random.random(512).astype(np.float32) + results = manager.search_similar(query, k=5) + print("搜索结果:", results) + + # 获取索引信息 + # 获得统计信息 + info = manager.get_index_info() + print("索引信息:", info) + + # 保存索引 + manager.save_index("./faiss_index") + + # 加载索引 + new_manager = FAISSManager(dimension=512) + new_manager.load_index("./faiss_index") \ No newline at end of file diff --git a/faiss_vector_db/food_labels.json b/faiss_vector_db/food_labels.json new file mode 100644 index 0000000..e69de29 diff --git a/faiss_vector_db/test_faiss_manager.py b/faiss_vector_db/test_faiss_manager.py new file mode 100644 index 0000000..069df55 --- /dev/null +++ b/faiss_vector_db/test_faiss_manager.py @@ -0,0 +1,459 @@ +""" +FAISS向量管理器测试文件 +用于验证FAISSManager的各项功能 +""" + +import unittest +import numpy as np +import tempfile +import shutil +import os +import sys + +# 添加项目根目录到路径 +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from faiss_vector_db.faiss_manager import FAISSManager + +class TestFAISSManager(unittest.TestCase): + """FAISS管理器测试类""" + + def setUp(self): + """测试前的设置""" + self.dimension = 64 + self.manager = FAISSManager( + dimension=self.dimension, + index_type="IndexFlatIP", + normalize_vectors=True, + metric_type="cosine" + ) + + # 创建临时目录用于测试保存/加载 + self.temp_dir = tempfile.mkdtemp() + + def tearDown(self): + """测试后的清理""" + # 删除临时目录 + if os.path.exists(self.temp_dir): + shutil.rmtree(self.temp_dir) + + def test_initialization(self): + """测试初始化""" + self.assertEqual(self.manager.dimension, self.dimension) + self.assertEqual(self.manager.index_type, "IndexFlatIP") + self.assertTrue(self.manager.normalize_vectors) + self.assertEqual(self.manager.metric_type, "cosine") + self.assertEqual(self.manager.index.ntotal, 0) + + def test_add_single_vector(self): + """测试添加单个向量,用numpy随机生成一个向量""" + vector = np.random.random(self.dimension).astype(np.float32) + vector_id = "test_001" + metadata = {"name": "测试菜品", "category": "测试"} + + # 测试添加 + result = self.manager.add_vector(vector, vector_id, metadata) + self.assertTrue(result) + + # 验证索引状态 + self.assertEqual(self.manager.index.ntotal, 1) + self.assertIn(vector_id, self.manager.reverse_id_mapping) + self.assertIn(vector_id, self.manager.metadata) + + # 测试重复添加 + result = self.manager.add_vector(vector, vector_id, metadata) + self.assertTrue(result) # 应该更新成功 + self.assertEqual(self.manager.index.ntotal, 2) # 因为,没有真正的删除 + + def test_add_batch_vectors(self): + """测试批量添加向量""" + batch_size = 10 + vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) + vector_ids = [f"batch_{i}" for i in range(batch_size)] + metadata_list = [{"name": f"菜品{i}", "category": "测试"} for i in range(batch_size)] + + # 测试批量添加 + results = self.manager.add_vectors_batch(vectors, vector_ids, metadata_list) + + # 验证结果 + self.assertEqual(len(results), batch_size) + self.assertTrue(all(results)) + self.assertEqual(self.manager.index.ntotal, batch_size) + + # 验证所有向量都被添加 + for vector_id in vector_ids: + self.assertIn(vector_id, self.manager.reverse_id_mapping) + self.assertIn(vector_id, self.manager.metadata) + + def test_search_similar(self): + """测试相似度搜索""" + # 先添加一些向量 + batch_size = 20 + vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) + vector_ids = [f"search_test_{i}" for i in range(batch_size)] + metadata_list = [{"name": f"菜品{i}", "score": i} for i in range(batch_size)] + + self.manager.add_vectors_batch(vectors, vector_ids, metadata_list) + + # 测试搜索 + query_vector = np.random.random(self.dimension).astype(np.float32) + k = 5 + results = self.manager.search_similar(query_vector, k=k) + + # 验证结果 + self.assertLessEqual(len(results), k) + self.assertLessEqual(len(results), batch_size) + + # 验证结果格式 + for vector_id, similarity, metadata in results: + self.assertIsInstance(vector_id, str) + self.assertIsInstance(similarity, float) + self.assertIsInstance(metadata, dict) + self.assertIn(vector_id, self.manager.reverse_id_mapping) + + # 测试带阈值的搜索 + threshold = 0.8 + results_with_threshold = self.manager.search_similar( + query_vector, k=k, threshold=threshold + ) + + # 验证阈值过滤 + for _, similarity, _ in results_with_threshold: + self.assertGreaterEqual(similarity, threshold) + + def test_get_vector_by_id(self): + """测试根据ID获取向量""" + vector = np.random.random(self.dimension).astype(np.float32) + vector_id = "get_test_001" + metadata = {"name": "测试获取", "value": 42} + + # 添加向量 + self.manager.add_vector(vector, vector_id, metadata) + + # 测试获取存在的向量 + result = self.manager.get_vector_by_id(vector_id) + self.assertIsNotNone(result) + vector_data, retrieved_metadata = result + self.assertEqual(retrieved_metadata, metadata) + + # 测试获取不存在的向量 + result = self.manager.get_vector_by_id("nonexistent") + self.assertIsNone(result) + + def test_update_vector(self): + """测试更新向量""" + # 添加初始向量 + vector = np.random.random(self.dimension).astype(np.float32) + vector_id = "update_test_001" + metadata = {"name": "原始菜品", "version": 1} + + self.manager.add_vector(vector, vector_id, metadata) + initial_count = self.manager.index.ntotal + + # 更新向量 + new_vector = np.random.random(self.dimension).astype(np.float32) + new_metadata = {"name": "更新菜品", "version": 2} + + result = self.manager.update_vector(new_vector, vector_id, new_metadata) + self.assertTrue(result) + + # 验证更新后的状态 + self.assertEqual(self.manager.index.ntotal, initial_count) + + # 验证元数据已更新 + vector_info = self.manager.get_vector_by_id(vector_id) + self.assertIsNotNone(vector_info) + _, retrieved_metadata = vector_info + self.assertEqual(retrieved_metadata, new_metadata) + + # 测试更新不存在的向量 + result = self.manager.update_vector(new_vector, "nonexistent", new_metadata) + self.assertFalse(result) + + def test_delete_vector(self): + """测试删除向量""" + # 添加向量 + vector = np.random.random(self.dimension).astype(np.float32) + vector_id = "delete_test_001" + metadata = {"name": "待删除菜品"} + + self.manager.add_vector(vector, vector_id, metadata) + + # 验证向量存在 + self.assertIn(vector_id, self.manager.reverse_id_mapping) + self.assertIn(vector_id, self.manager.metadata) + + # 删除向量 + result = self.manager.delete_vector(vector_id) + self.assertTrue(result) + + # 验证向量已删除 + self.assertNotIn(vector_id, self.manager.reverse_id_mapping) + self.assertNotIn(vector_id, self.manager.metadata) + + # 测试删除不存在的向量 + result = self.manager.delete_vector("nonexistent") + self.assertFalse(result) + + def test_delete_vectors_batch(self): + """测试批量删除向量""" + # 添加多个向量 + batch_size = 5 + vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) + vector_ids = [f"batch_delete_{i}" for i in range(batch_size)] + + self.manager.add_vectors_batch(vectors, vector_ids) + + # 批量删除部分向量 + delete_ids = vector_ids[:3] + results = self.manager.delete_vectors_batch(delete_ids) + + # 验证删除结果 + self.assertEqual(len(results), len(delete_ids)) + self.assertTrue(all(results)) + + # 验证向量已删除 + for vector_id in delete_ids: + self.assertNotIn(vector_id, self.manager.reverse_id_mapping) + + # 验证剩余向量仍存在 + for vector_id in vector_ids[3:]: + self.assertIn(vector_id, self.manager.reverse_id_mapping) + + def test_save_and_load_index(self): + """测试保存和加载索引""" + # 添加一些向量 + batch_size = 10 + vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) + vector_ids = [f"save_load_{i}" for i in range(batch_size)] + metadata_list = [{"name": f"菜品{i}", "id": i} for i in range(batch_size)] + + self.manager.add_vectors_batch(vectors, vector_ids, metadata_list) + + # 保存索引 + save_path = os.path.join(self.temp_dir, "test_index") + result = self.manager.save_index(save_path) + self.assertTrue(result) + + # 验证文件已创建 + self.assertTrue(os.path.exists(os.path.join(save_path, "faiss_index.bin"))) + self.assertTrue(os.path.exists(os.path.join(save_path, "metadata.pkl"))) + self.assertTrue(os.path.exists(os.path.join(save_path, "config.json"))) + + # 创建新管理器并加载索引 + new_manager = FAISSManager(dimension=self.dimension) + result = new_manager.load_index(save_path) + self.assertTrue(result) + + # 验证加载后的状态 + self.assertEqual(new_manager.index.ntotal, batch_size) + self.assertEqual(len(new_manager.reverse_id_mapping), batch_size) + self.assertEqual(len(new_manager.metadata), batch_size) + + # 验证数据一致性 + for vector_id in vector_ids: + self.assertIn(vector_id, new_manager.reverse_id_mapping) + self.assertIn(vector_id, new_manager.metadata) + + # 测试加载后的搜索功能 + query_vector = np.random.random(self.dimension).astype(np.float32) + results = new_manager.search_similar(query_vector, k=3) + self.assertGreater(len(results), 0) + + def test_clear_index(self): + """测试清空索引""" + # 添加一些向量 + vectors = np.random.random((5, self.dimension)).astype(np.float32) + vector_ids = [f"clear_test_{i}" for i in range(5)] + + self.manager.add_vectors_batch(vectors, vector_ids) + + # 验证向量已添加 + self.assertEqual(self.manager.index.ntotal, 5) + self.assertEqual(len(self.manager.reverse_id_mapping), 5) + + # 清空索引 + result = self.manager.clear_index() + self.assertTrue(result) + + # 验证索引已清空 + self.assertEqual(self.manager.index.ntotal, 0) + self.assertEqual(len(self.manager.reverse_id_mapping), 0) + self.assertEqual(len(self.manager.metadata), 0) + self.assertEqual(self.manager.next_faiss_id, 0) + + def test_get_index_info(self): + """测试获取索引信息""" + # 初始状态 + info = self.manager.get_index_info() + self.assertEqual(info['total_vectors'], 0) + self.assertEqual(info['dimension'], self.dimension) + self.assertEqual(info['active_vectors'], 0) + + # 添加向量后 + vectors = np.random.random((3, self.dimension)).astype(np.float32) + vector_ids = [f"info_test_{i}" for i in range(3)] + + self.manager.add_vectors_batch(vectors, vector_ids) + + info = self.manager.get_index_info() + self.assertEqual(info['total_vectors'], 3) + self.assertEqual(info['active_vectors'], 3) + + # 删除一个向量后 + self.manager.delete_vector(vector_ids[0]) + + info = self.manager.get_index_info() + self.assertEqual(info['total_vectors'], 3) # FAISS索引中的总数不变 + self.assertEqual(info['active_vectors'], 2) # 活跃向量数减少 + + def test_different_index_types(self): + """测试不同的索引类型""" + index_types = ["IndexFlatIP", "IndexFlatL2"] + + for index_type in index_types: + with self.subTest(index_type=index_type): + manager = FAISSManager( + dimension=self.dimension, + index_type=index_type + ) + + # 添加向量 + vector = np.random.random(self.dimension).astype(np.float32) + result = manager.add_vector(vector, f"test_{index_type}") + self.assertTrue(result) + + # 搜索 + query = np.random.random(self.dimension).astype(np.float32) + results = manager.search_similar(query, k=1) + self.assertEqual(len(results), 1) + + def test_vector_normalization(self): + """测试向量标准化""" + # 测试标准化开启 + manager_norm = FAISSManager( + dimension=self.dimension, + normalize_vectors=True + ) + + # 测试标准化关闭 + manager_no_norm = FAISSManager( + dimension=self.dimension, + normalize_vectors=False + ) + + # 添加相同的向量 + vector = np.random.random(self.dimension).astype(np.float32) * 10 # 放大向量 + + manager_norm.add_vector(vector, "norm_test") + manager_no_norm.add_vector(vector, "no_norm_test") + + # 两个管理器都应该成功添加 + self.assertEqual(manager_norm.index.ntotal, 1) + self.assertEqual(manager_no_norm.index.ntotal, 1) + + +class TestFAISSManagerIntegration(unittest.TestCase): + """FAISS管理器集成测试""" + + def test_food_classification_scenario(self): + """测试食物分类场景""" + # 创建管理器 + manager = FAISSManager(dimension=128, index_type="IndexFlatIP") + + # 模拟食物向量数据 + food_data = [ + {"id": "sichuan_001", "name": "回锅肉", "category": "川菜", "spicy": 3}, + {"id": "sichuan_002", "name": "麻辣小面", "category": "川菜", "spicy": 4}, + {"id": "sichuan_003", "name": "宫保鸡丁", "category": "川菜", "spicy": 3}, + {"id": "home_001", "name": "西红柿鸡蛋", "category": "家常菜", "spicy": 0}, + {"id": "home_002", "name": "炒细面", "category": "家常菜", "spicy": 1}, + ] + + # 生成模拟向量(实际应用中这些是通过embedding模型生成的) + vectors = [] + for i, food in enumerate(food_data): + # 为川菜生成相似的向量,为家常菜生成另一类相似的向量 + if food["category"] == "川菜": + base_vector = np.array([1.0] * 64 + [0.0] * 64) + else: + base_vector = np.array([0.0] * 64 + [1.0] * 64) + + # 添加随机噪声 + noise = np.random.normal(0, 0.1, 128) + vector = (base_vector + noise).astype(np.float32) + vectors.append(vector) + + # 批量添加向量 + vector_ids = [food["id"] for food in food_data] + metadata_list = [ + {"name": food["name"], "category": food["category"], "spicy": food["spicy"]} + for food in food_data + ] + + results = manager.add_vectors_batch(vectors, vector_ids, metadata_list) + self.assertTrue(all(results)) + + # 测试川菜查询 + sichuan_query = np.array([1.0] * 64 + [0.0] * 64).astype(np.float32) + sichuan_results = manager.search_similar(sichuan_query, k=3) + + # 验证川菜结果 + self.assertGreater(len(sichuan_results), 0) + for vector_id, similarity, metadata in sichuan_results: + self.assertEqual(metadata["category"], "川菜") + + # 测试家常菜查询 + home_query = np.array([0.0] * 64 + [1.0] * 64).astype(np.float32) + home_results = manager.search_similar(home_query, k=3) + + # 验证家常菜结果 + self.assertGreater(len(home_results), 0) + for vector_id, similarity, metadata in home_results: + self.assertEqual(metadata["category"], "家常菜") + + # 测试保存和加载 + temp_dir = tempfile.mkdtemp() + try: + save_path = os.path.join(temp_dir, "food_index") + self.assertTrue(manager.save_index(save_path)) + + # 加载到新管理器 + new_manager = FAISSManager(dimension=128) + self.assertTrue(new_manager.load_index(save_path)) + + # 验证加载后的搜索功能 + loaded_results = new_manager.search_similar(sichuan_query, k=2) + self.assertGreater(len(loaded_results), 0) + + finally: + shutil.rmtree(temp_dir) + + +def run_tests(): + """运行所有测试""" + # 创建测试套件 + test_suite = unittest.TestSuite() + + # 添加基本功能测试 + test_suite.addTest(unittest.makeSuite(TestFAISSManager)) + + # 添加集成测试 + test_suite.addTest(unittest.makeSuite(TestFAISSManagerIntegration)) + + # 运行测试 + runner = unittest.TextTestRunner(verbosity=2) + result = runner.run(test_suite) + + return result.wasSuccessful() + + +if __name__ == "__main__": + print("开始运行FAISS管理器测试...") + success = run_tests() + + if success: + print("\n✅ 所有测试通过!") + else: + print("\n❌ 部分测试失败,请检查代码。") + sys.exit(1) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 1fbf9b7..497d155 100644 --- a/requirements.txt +++ b/requirements.txt @@ -25,3 +25,4 @@ torch==2.7.0 torchvision==0.22.0 tqdm==4.65.0 typing_extensions==4.15.0 +faiss-cpu>=1.7.0