增加FAISS向量数据库检索存储等功能。
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user