459 lines
17 KiB
Python
459 lines
17 KiB
Python
"""
|
|
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) |