增加FAISS向量数据库检索存储等功能。
This commit is contained in:
@@ -1,3 +1,5 @@
|
||||
/dataset/
|
||||
/.idea/
|
||||
/model/
|
||||
/faiss_vector_db/demo_faiss_index/
|
||||
/faiss_vector_db/faiss_index/
|
||||
|
||||
@@ -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]]
|
||||
@@ -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
|
||||
|
||||
@@ -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 许可证。
|
||||
@@ -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()
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user