Files
FoodClassifier/exp_multimodal/exp_openset_vlm.py
T

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