""" 开放式VLM识别 + 向量匹配 主程序 用于测试"开放识别+向量检索"方案的性能 """ import argparse import json import os import sys import time # 兼容:支持直接运行脚本或用 -m 模块方式运行 if __name__ == "__main__" and __package__ is None: sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from exp_multimodal.vlm_classifier import classify_image_openset from exp_multimodal.text_embedder import OllamaEmbedder from exp_multimodal.vector_matcher import DishNameMatcher from exp_multimodal.vlm_providers.ollama_provider import OllamaProvider def main(): ap = argparse.ArgumentParser(description="开放式菜品识别测试") ap.add_argument( "--image", default=r"D:\MyProjects\PythonProjects\FoodClassifier\dataset\DishClassification\test\红烧肉\img04.png", help="测试图片路径" ) ap.add_argument( "--index_dir", default=None, help="FAISS索引目录(默认: faiss_vector_db/dish_names)" ) ap.add_argument( "--vlm_url", default="http://192.168.1.250:11434", help="VLM服务地址(Ollama)" ) ap.add_argument( "--vlm_model", default="qwen2.5vl:32b", help="VLM模型名称" ) ap.add_argument( "--embedder_url", default="http://192.168.1.250:11434", help="Embedding服务地址(Ollama)" ) ap.add_argument( "--embedder_model", default="quentinz/bge-large-zh-v1.5", help="Embedding模型名称" ) ap.add_argument( "--top_k", type=int, default=3, help="向量检索返回的候选数" ) ap.add_argument( "--min_score", type=float, default=0.5, help="最低匹配分数阈值" ) args = ap.parse_args() # 确定索引目录 if args.index_dir: index_dir = args.index_dir else: project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) index_dir = os.path.join(project_root, "faiss_vector_db", "dish_names") print("=" * 80) print(f"[Main] 开放式VLM识别 + 向量匹配测试") print("=" * 80) print(f"[Main] Image: {args.image}") print(f"[Main] Image exists: {os.path.exists(args.image)}") if os.path.exists(args.image): print(f"[Main] Image size: {os.path.getsize(args.image)} bytes") print(f"[Main] Index dir: {index_dir}") print(f"[Main] VLM: {args.vlm_url} / {args.vlm_model}") print(f"[Main] Embedder: {args.embedder_url} / {args.embedder_model}") print(f"[Main] Top-K: {args.top_k}, Min score: {args.min_score}") print("=" * 80) # 初始化组件 print(f"\n[Main] 初始化VLM Provider...") vlm_provider = OllamaProvider( base_url=args.vlm_url, model=args.vlm_model, ) print(f"\n[Main] 初始化Embedder...") embedder = OllamaEmbedder( base_url=args.embedder_url, model=args.embedder_model, ) print(f"\n[Main] 加载FAISS索引...") matcher = DishNameMatcher(index_dir) # 执行识别 print(f"\n[Main] 开始识别...") print("=" * 80) t_start = time.time() result = classify_image_openset( image_path=args.image, provider=vlm_provider, embedder=embedder, matcher=matcher, top_k=args.top_k, min_match_score=args.min_score, ) t_elapsed = time.time() - t_start print("=" * 80) print(f"\n[Main] ✅ 识别完成!耗时: {t_elapsed:.2f}秒") print("=" * 80) print("\n识别结果:") print(json.dumps(result, ensure_ascii=False, indent=2)) print("=" * 80) # 简要总结 print(f"\n📊 性能总结:") print(f" - 总耗时: {t_elapsed:.2f}秒") print(f" - VLM原始输出: {result['raw_dish']}") print(f" - 清洗后菜名: {result['cleaned_dish']}") print(f" - VLM置信度: {result['vlm_confidence']:.4f}") print(f" - 最佳匹配: {result['best_match']}") print(f" - 最终置信度: {result['final_confidence']:.4f}") print(f" - 候选数量: {len(result['candidates'])}") if t_elapsed <= 3.0: print(f"\n✅ 性能达标!({t_elapsed:.2f}s ≤ 3.0s)") else: print(f"\n⚠️ 性能未达标 ({t_elapsed:.2f}s > 3.0s)") if __name__ == "__main__": main()