144 lines
4.3 KiB
Python
144 lines
4.3 KiB
Python
"""
|
||
开放式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()
|