Files
FoodClassifier/faiss_vector_db/example_usage.py
T

223 lines
7.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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()