增加FAISS向量数据库检索存储等功能。

This commit is contained in:
zhanghuan
2025-09-22 10:42:57 +08:00
parent e34ecbd94d
commit be284442e4
10 changed files with 1573 additions and 0 deletions
+2
View File
@@ -1,3 +1,5 @@
/dataset/
/.idea/
/model/
/faiss_vector_db/demo_faiss_index/
/faiss_vector_db/faiss_index/
+40
View File
@@ -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]]
+20
View File
@@ -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
+312
View File
@@ -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 许可证。
+223
View File
@@ -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()
+516
View File
@@ -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")
View File
+459
View File
@@ -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)
+1
View File
@@ -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