增加多模态模型开放集识别。

This commit is contained in:
2025-11-13 11:56:51 +08:00
parent b4b2b18ccd
commit ced0e4f7ef
7 changed files with 628 additions and 21 deletions
+143
View File
@@ -0,0 +1,143 @@
"""
开放式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()