Files
FoodClassifier/faiss_vector_db/faiss_similarity_search.py

318 lines
12 KiB
Python

"""
FAISS相似度检索演示程序
提供简单的命令行界面来测试相似度检索功能
功能:
1. 加载预构建的FAISS索引
2. 对指定图片进行相似度检索
3. 显示检索结果和统计信息
4. 支持批量检索和性能测试
"""
import os
import sys
import time
import argparse
from typing import List, Tuple
from PIL import Image
import matplotlib.pyplot as plt
# 导入构建的FAISS检索器
from faiss_vector_db.build_faiss_index import FAISSSearcher
class SimilaritySearchDemo:
"""相似度检索演示类"""
def __init__(self, index_dir: str, model_path: str):
"""
初始化演示程序
Args:
index_dir: FAISS索引目录
model_path: 模型路径
"""
self.searcher = FAISSSearcher(index_dir, model_path)
# 设置中文字体(用于matplotlib显示)
try:
# 尝试使用系统中文字体
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False
except:
print("警告: 无法设置中文字体,图表中的中文可能显示异常")
def search_and_display(self, query_image_path: str, k: int = 5,
show_images: bool = False, save_result: bool = False):
"""
搜索相似图片并显示结果
Args:
query_image_path: 查询图片路径
k: 返回最相似的k张图片
show_images: 是否显示图片
save_result: 是否保存结果图片
"""
print(f"\n查询图片: {query_image_path}")
print("-" * 50)
# 检查文件是否存在
if not os.path.exists(query_image_path):
print(f"错误: 文件不存在 - {query_image_path}")
return
# 开始计时
start_time = time.time()
try:
# 执行相似度搜索
results = self.searcher.search_similar(query_image_path, k)
# 计算耗时
search_time = time.time() - start_time
# 显示结果
print(f"搜索完成,耗时: {search_time:.3f}秒")
print(f"找到 {len(results)} 个相似结果:\n")
for i, (img_path, class_name, score) in enumerate(results, 1):
rel_path = os.path.relpath(img_path)
print(f"{i:2d}. {os.path.basename(img_path):20s} "
f"类别: {class_name:8s} "
f"相似度: {score:.4f}")
# 统计类别分布
class_counts = {}
for _, class_name, _ in results:
class_counts[class_name] = class_counts.get(class_name, 0) + 1
print(f"\n类别分布:")
for class_name, count in class_counts.items():
print(f" {class_name}: {count} 张")
# 显示图片
if show_images:
self.display_results(query_image_path, results, save_result)
except Exception as e:
print(f"搜索失败: {e}")
import traceback
traceback.print_exc()
def display_results(self, query_path: str, results: List[Tuple[str, str, float]],
save_result: bool = False):
"""
显示搜索结果图片
Args:
query_path: 查询图片路径
results: 搜索结果
save_result: 是否保存结果图片
"""
try:
# 计算子图布局
n_results = len(results)
n_cols = min(3, n_results + 1) # 包括查询图片
n_rows = (n_results + 1 + n_cols - 1) // n_cols
fig, axes = plt.subplots(n_rows, n_cols, figsize=(4*n_cols, 4*n_rows))
if n_rows == 1:
axes = axes.reshape(1, -1)
# 显示查询图片
query_img = Image.open(query_path).convert('RGB')
axes[0, 0].imshow(query_img)
axes[0, 0].set_title(f"查询图片\n{os.path.basename(query_path)}", fontsize=10)
axes[0, 0].axis('off')
# 显示搜索结果
for i, (img_path, class_name, score) in enumerate(results):
row = (i + 1) // n_cols
col = (i + 1) % n_cols
if row < n_rows and col < n_cols:
try:
img = Image.open(img_path).convert('RGB')
axes[row, col].imshow(img)
axes[row, col].set_title(
f"#{i+1} {class_name}\n"
f"{os.path.basename(img_path)}\n"
f"相似度: {score:.3f}",
fontsize=9
)
axes[row, col].axis('off')
except Exception as e:
axes[row, col].text(0.5, 0.5, f"加载失败\n{e}",
ha='center', va='center')
axes[row, col].axis('off')
# 隐藏多余的子图
for i in range(n_results + 1, n_rows * n_cols):
row = i // n_cols
col = i % n_cols
if row < n_rows and col < n_cols:
axes[row, col].axis('off')
plt.tight_layout()
# 保存结果
if save_result:
timestamp = time.strftime("%Y%m%d_%H%M%S")
save_path = f"similarity_search_result_{timestamp}.png"
plt.savefig(save_path, dpi=150, bbox_inches='tight')
print(f"结果图片已保存到: {save_path}")
plt.show()
except Exception as e:
print(f"显示图片失败: {e}")
def batch_search_test(self, test_dir: str, k: int = 5, max_tests: int = 10):
"""
批量搜索测试
Args:
test_dir: 测试图片目录
k: 每次搜索返回的结果数
max_tests: 最大测试数量
"""
print(f"\n批量搜索测试 - 目录: {test_dir}")
print("=" * 60)
# 获取测试图片
test_images = []
for root, dirs, files in os.walk(test_dir):
for file in files:
if file.lower().endswith(('.jpg', '.jpeg', '.png', '.bmp')):
test_images.append(os.path.join(root, file))
if len(test_images) >= max_tests:
break
if len(test_images) >= max_tests:
break
if not test_images:
print("未找到测试图片")
return
print(f"找到 {len(test_images)} 张测试图片")
# 执行批量测试
total_time = 0
success_count = 0
for i, img_path in enumerate(test_images, 1):
print(f"\n[{i}/{len(test_images)}] 测试: {os.path.basename(img_path)}")
try:
start_time = time.time()
results = self.searcher.search_similar(img_path, k)
search_time = time.time() - start_time
total_time += search_time
success_count += 1
print(f" 搜索耗时: {search_time:.3f}秒")
print(f" 前3个结果:")
for j, (result_path, class_name, score) in enumerate(results[:3], 1):
print(f" {j}. {class_name} (相似度: {score:.3f})")
except Exception as e:
print(f" 搜索失败: {e}")
# 显示统计信息
print("\n" + "=" * 60)
print("批量测试统计:")
print(f" 总测试数: {len(test_images)}")
print(f" 成功数: {success_count}")
print(f" 失败数: {len(test_images) - success_count}")
if success_count > 0:
print(f" 平均搜索时间: {total_time/success_count:.3f}秒")
print(f" 总耗时: {total_time:.3f}秒")
def show_index_info(self):
"""显示索引信息"""
print("\n" + "=" * 60)
print("FAISS索引信息")
print("=" * 60)
print(f"索引目录: {self.searcher.index_dir}")
print(f"总向量数: {self.searcher.index.ntotal}")
print(f"特征维度: {self.searcher.embedding_dim}")
print(f"类别数: {len(self.searcher.class_names)}")
print(f"类别列表: {', '.join(self.searcher.class_names)}")
# 显示每个类别的图片数量
print(f"\n各类别图片数量:")
class_counts = {}
for label in self.searcher.labels:
class_name = self.searcher.class_names[label]
class_counts[class_name] = class_counts.get(class_name, 0) + 1
for class_name in self.searcher.class_names:
count = class_counts.get(class_name, 0)
print(f" {class_name}: {count} 张")
print(f"\n构建时间: {self.searcher.class_info.get('build_time', '未知')}")
def main():
"""主函数"""
parser = argparse.ArgumentParser(description='FAISS相似度检索演示程序')
parser.add_argument('--index_dir', type=str, default='faiss_vector_db/faiss_index',
help='FAISS索引目录路径')
parser.add_argument('--model_path', type=str,
default='model/embedding_20250917_145342/best_embedding_model.pth',
help='模型文件路径')
parser.add_argument('--query', type=str, help='查询图片路径')
parser.add_argument('--k', type=int, default=5, help='返回最相似的k张图片')
parser.add_argument('--show_images', action='store_true', help='显示搜索结果图片')
parser.add_argument('--save_result', action='store_true', help='保存结果图片')
parser.add_argument('--batch_test', type=str, help='批量测试目录')
parser.add_argument('--max_tests', type=int, default=10, help='批量测试最大数量')
parser.add_argument('--info', action='store_true', help='显示索引信息')
args = parser.parse_args()
try:
# 检查索引目录
if not os.path.exists(args.index_dir):
print(f"错误: 索引目录不存在 - {args.index_dir}")
print("请先运行 build_faiss_index.py 构建索引")
return
# 检查模型文件
if not os.path.exists(args.model_path):
print(f"错误: 模型文件不存在 - {args.model_path}")
return
# 初始化演示程序
demo = SimilaritySearchDemo(args.index_dir, args.model_path)
# 显示索引信息
if args.info:
demo.show_index_info()
# 单张图片搜索
if args.query:
demo.search_and_display(args.query, args.k, args.show_images, args.save_result)
# 批量测试
if args.batch_test:
demo.batch_search_test(args.batch_test, args.k, args.max_tests)
# 如果没有指定任何操作,显示帮助信息
if not any([args.query, args.batch_test, args.info]):
print("使用示例:")
print(f" python {sys.argv[0]} --info # 显示索引信息")
print(f" python {sys.argv[0]} --query dataset/test/回锅肉/aug_000.jpg --show_images")
print(f" python {sys.argv[0]} --batch_test dataset/test --max_tests 5")
except Exception as e:
print(f"程序执行失败: {e}")
import traceback
traceback.print_exc()
if __name__ == "__main__":
main()