223 lines
7.8 KiB
Python
223 lines
7.8 KiB
Python
"""
|
||
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() |