Files
FoodClassifier/faiss_vector_db/test_faiss_manager.py
T

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)