增加多模态模型开放集识别。
This commit is contained in:
@@ -0,0 +1,158 @@
|
||||
"""
|
||||
构建菜品名向量索引
|
||||
将labels.py中的菜品名编码为向量并构建FAISS索引(一次性任务)
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import List
|
||||
|
||||
import faiss
|
||||
import numpy as np
|
||||
|
||||
# 兼容:支持直接运行脚本或用 -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.labels import build_labels
|
||||
from exp_multimodal.text_embedder import OllamaEmbedder
|
||||
|
||||
|
||||
def build_index(
|
||||
dish_names: List[str],
|
||||
embedder: OllamaEmbedder,
|
||||
output_dir: str,
|
||||
batch_size: int = 100,
|
||||
) -> None:
|
||||
"""
|
||||
构建FAISS索引
|
||||
|
||||
参数:
|
||||
dish_names: 菜品名列表
|
||||
embedder: OllamaEmbedder实例
|
||||
output_dir: 输出目录
|
||||
batch_size: 批量编码大小
|
||||
"""
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
print(f"[BuildIndex] Total dishes={len(dish_names)} batch_size={batch_size}")
|
||||
|
||||
# 批量编码
|
||||
all_embeddings = []
|
||||
for i in range(0, len(dish_names), batch_size):
|
||||
batch = dish_names[i:i+batch_size]
|
||||
print(f"[BuildIndex] Encoding batch {i//batch_size + 1}/{(len(dish_names)-1)//batch_size + 1} (size={len(batch)})...")
|
||||
|
||||
try:
|
||||
batch_embs = embedder.encode(batch)
|
||||
all_embeddings.append(batch_embs)
|
||||
except Exception as e:
|
||||
print(f"[BuildIndex] Error encoding batch {i//batch_size + 1}: {e}")
|
||||
raise
|
||||
|
||||
# 合并所有向量
|
||||
embeddings = np.vstack(all_embeddings)
|
||||
print(f"[BuildIndex] Concatenated embeddings shape={embeddings.shape}")
|
||||
|
||||
# 归一化向量(用于余弦相似度)
|
||||
norms = np.linalg.norm(embeddings, axis=1, keepdims=True)
|
||||
embeddings = embeddings / (norms + 1e-8)
|
||||
print(f"[BuildIndex] Normalized embeddings")
|
||||
|
||||
# 构建FAISS索引(IndexFlatIP = 内积索引,适合归一化后的向量)
|
||||
dim = embeddings.shape[1]
|
||||
index = faiss.IndexFlatIP(dim)
|
||||
index.add(embeddings.astype(np.float32))
|
||||
print(f"[BuildIndex] Built FAISS index dim={dim} ntotal={index.ntotal}")
|
||||
|
||||
# 保存文件
|
||||
names_path = os.path.join(output_dir, "dish_names.json")
|
||||
embeddings_path = os.path.join(output_dir, "dish_embeddings.npy")
|
||||
index_path = os.path.join(output_dir, "faiss_index.bin")
|
||||
|
||||
with open(names_path, "w", encoding="utf-8") as f:
|
||||
json.dump(dish_names, f, ensure_ascii=False, indent=2)
|
||||
print(f"[BuildIndex] Saved dish names to {names_path}")
|
||||
|
||||
np.save(embeddings_path, embeddings)
|
||||
print(f"[BuildIndex] Saved embeddings to {embeddings_path}")
|
||||
|
||||
faiss.write_index(index, index_path)
|
||||
print(f"[BuildIndex] Saved FAISS index to {index_path}")
|
||||
|
||||
print(f"[BuildIndex] ✅ Index build complete! Output dir: {output_dir}")
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description="构建菜品名向量索引")
|
||||
ap.add_argument(
|
||||
"--mode",
|
||||
choices=["dish", "whole", "processed"],
|
||||
default="dish",
|
||||
help="数据集模式(默认: dish)"
|
||||
)
|
||||
ap.add_argument(
|
||||
"--output",
|
||||
default=None,
|
||||
help="输出目录(默认: faiss_vector_db/<mode>_names)"
|
||||
)
|
||||
ap.add_argument(
|
||||
"--embedder_url",
|
||||
default="http://192.168.1.250:11434",
|
||||
help="Ollama服务地址"
|
||||
)
|
||||
ap.add_argument(
|
||||
"--embedder_model",
|
||||
default="quentinz/bge-large-zh-v1.5",
|
||||
help="Embedding模型名称"
|
||||
)
|
||||
ap.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=100,
|
||||
help="批量编码大小"
|
||||
)
|
||||
|
||||
args = ap.parse_args()
|
||||
|
||||
# 确定输出目录
|
||||
if args.output:
|
||||
output_dir = args.output
|
||||
else:
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
output_dir = os.path.join(project_root, "faiss_vector_db", f"{args.mode}_names")
|
||||
|
||||
print(f"[Main] mode={args.mode} output_dir={output_dir}")
|
||||
print(f"[Main] embedder_url={args.embedder_url} model={args.embedder_model}")
|
||||
|
||||
# 构建菜品名列表
|
||||
print(f"[Main] Building labels from mode={args.mode}...")
|
||||
dish_names = build_labels(args.mode, alias_map_path=None)
|
||||
print(f"[Main] Built {len(dish_names)} dish names")
|
||||
|
||||
if not dish_names:
|
||||
print("[Main] ❌ No dish names found, abort")
|
||||
return
|
||||
|
||||
# 初始化Embedder
|
||||
print(f"[Main] Initializing OllamaEmbedder...")
|
||||
embedder = OllamaEmbedder(
|
||||
base_url=args.embedder_url,
|
||||
model=args.embedder_model,
|
||||
)
|
||||
|
||||
# 构建索引
|
||||
print(f"[Main] Building FAISS index...")
|
||||
build_index(
|
||||
dish_names=dish_names,
|
||||
embedder=embedder,
|
||||
output_dir=output_dir,
|
||||
batch_size=args.batch_size,
|
||||
)
|
||||
|
||||
print(f"[Main] ✅ All done!")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,102 @@
|
||||
"""
|
||||
菜品名清洗工具
|
||||
用于将VLM输出的菜品名进行标准化处理,提升向量匹配准确率
|
||||
"""
|
||||
import re
|
||||
from typing import Dict
|
||||
|
||||
|
||||
# 同义词映射表(可根据实际情况扩展)
|
||||
SYNONYM_MAP: Dict[str, str] = {
|
||||
"西红柿": "番茄",
|
||||
"土豆": "马铃薯",
|
||||
"洋芋": "马铃薯",
|
||||
"青椒": "柿子椒",
|
||||
# 可继续添加...
|
||||
}
|
||||
|
||||
|
||||
def clean_dish_name(name: str) -> str:
|
||||
"""
|
||||
清洗菜品名:去除括号注释、前缀、英文等干扰信息
|
||||
|
||||
示例:
|
||||
"宫保鸡丁(川菜)" -> "宫保鸡丁"
|
||||
"川菜-麻婆豆腐" -> "麻婆豆腐"
|
||||
"红烧肉 Braised Pork" -> "红烧肉"
|
||||
"""
|
||||
if not isinstance(name, str):
|
||||
return ""
|
||||
|
||||
# 1. 去除括号及内容(中英文括号)
|
||||
name = re.sub(r"[((].*?[))]", "", name)
|
||||
name = re.sub(r"\[.*?\]", "", name)
|
||||
|
||||
# 2. 去除常见前缀(菜系、地域等)
|
||||
prefixes = ["川菜", "粤菜", "鲁菜", "苏菜", "浙菜", "闽菜", "湘菜", "徽菜",
|
||||
"东北", "西北", "西南", "华南", "华北"]
|
||||
for prefix in prefixes:
|
||||
if name.startswith(prefix):
|
||||
name = name[len(prefix):]
|
||||
break
|
||||
|
||||
# 3. 去除分隔符后的前缀(如 "川菜-宫保鸡丁")
|
||||
name = re.sub(r"^[^-—]*[-—]", "", name)
|
||||
|
||||
# 4. 去除英文部分(保留中文)
|
||||
name = re.sub(r"[a-zA-Z\s]+", "", name)
|
||||
|
||||
# 5. 去除多余空格和标点
|
||||
name = re.sub(r"[,,、·\s]+", "", name)
|
||||
|
||||
# 6. 去除可能的烹饪方式后缀(如果VLM违规输出)
|
||||
cooking_suffixes = ["炒制", "烹饪", "料理", "做法"]
|
||||
for suffix in cooking_suffixes:
|
||||
if name.endswith(suffix):
|
||||
name = name[:-len(suffix)]
|
||||
|
||||
return name.strip()
|
||||
|
||||
|
||||
def normalize_dish_name(name: str) -> str:
|
||||
"""
|
||||
标准化菜品名:应用同义词映射
|
||||
|
||||
示例:
|
||||
"西红柿炒鸡蛋" -> "番茄炒鸡蛋"
|
||||
"""
|
||||
cleaned = clean_dish_name(name)
|
||||
|
||||
# 应用同义词替换
|
||||
for synonym, standard in SYNONYM_MAP.items():
|
||||
cleaned = cleaned.replace(synonym, standard)
|
||||
|
||||
return cleaned
|
||||
|
||||
|
||||
def extract_main_dish_name(text: str) -> str:
|
||||
"""
|
||||
从VLM返回文本中提取主要菜品名(兼容多种输出格式)
|
||||
|
||||
示例:
|
||||
"这是宫保鸡丁" -> "宫保鸡丁"
|
||||
"菜品:红烧肉" -> "红烧肉"
|
||||
"""
|
||||
if not text:
|
||||
return ""
|
||||
|
||||
# 尝试匹配常见模式
|
||||
patterns = [
|
||||
r"菜品[::]\s*([^,,。\n]+)",
|
||||
r"识别为[::]\s*([^,,。\n]+)",
|
||||
r"这是\s*([^,,。\n]+)",
|
||||
r"应该是\s*([^,,。\n]+)",
|
||||
]
|
||||
|
||||
for pattern in patterns:
|
||||
match = re.search(pattern, text)
|
||||
if match:
|
||||
return match.group(1).strip()
|
||||
|
||||
# 如果没有匹配到,返回清洗后的整个文本
|
||||
return clean_dish_name(text)
|
||||
@@ -10,8 +10,11 @@ if __name__ == "__main__" and __package__ is None:
|
||||
|
||||
# 统一使用绝对导入,避免相对导入在脚本直跑时失败
|
||||
from exp_multimodal.labels import build_labels
|
||||
from exp_multimodal.vlm_classifier import classify_image
|
||||
from exp_multimodal.vlm_classifier import classify_image, classify_image_openset
|
||||
from exp_multimodal.ollama_client import OLLAMA_URL, DEFAULT_MODEL
|
||||
from exp_multimodal.text_embedder import OllamaEmbedder
|
||||
from exp_multimodal.vector_matcher import DishNameMatcher
|
||||
from exp_multimodal.vlm_providers.ollama_provider import OllamaVLMProvider
|
||||
|
||||
|
||||
FEWSHOT_HINTS: Dict[str, str] = {
|
||||
@@ -24,9 +27,14 @@ FEWSHOT_HINTS: Dict[str, str] = {
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--mode", choices=["dish", "whole", "processed"], default="dish")
|
||||
ap.add_argument("--mode", choices=["dish", "whole", "processed", "openset_dish"], default="openset_dish")
|
||||
ap.add_argument("--image", default=r"D:\MyProjects\PythonProjects\FoodClassifier\dataset\DishClassification\test\红烧肉\img04.png")
|
||||
ap.add_argument("--alias_map", default=None)
|
||||
|
||||
# 开放式识别专用参数
|
||||
ap.add_argument("--openset-top-k", type=int, default=3, help="开放式识别:向量检索Top-K候选数")
|
||||
ap.add_argument("--openset-min-score", type=float, default=0.5, help="开放式识别:最低匹配分数阈值")
|
||||
ap.add_argument("--openset-index-path", default="exp_multimodal/dish_name_index", help="开放式识别:向量索引路径")
|
||||
|
||||
args = ap.parse_args()
|
||||
|
||||
@@ -37,24 +45,59 @@ def main() -> None:
|
||||
)
|
||||
print(f"[Main] Image exists={os.path.exists(args.image)} size={os.path.getsize(args.image) if os.path.exists(args.image) else 'N/A'}")
|
||||
|
||||
labels = build_labels(args.mode, args.alias_map)
|
||||
ingredient_only = args.mode in {"whole", "processed"}
|
||||
fewshot = FEWSHOT_HINTS if args.mode == "dish" else None
|
||||
print(f"[Main] Built labels count={len(labels)} ingredient_only={ingredient_only} fewshot={bool(fewshot)}")
|
||||
# 开放式识别分支
|
||||
if args.mode == "openset_dish":
|
||||
print(f"[Main] Openset mode: top_k={args.openset_top_k} min_score={args.openset_min_score} index_path={args.openset_index_path}")
|
||||
|
||||
# 初始化Provider、Embedder和Matcher
|
||||
provider = OllamaVLMProvider(base_url=OLLAMA_URL, model_name=DEFAULT_MODEL)
|
||||
embedder = OllamaEmbedder(base_url=OLLAMA_URL, model_name="bge-large-zh-v1.5")
|
||||
matcher = DishNameMatcher(index_dir=args.openset_index_path)
|
||||
|
||||
print("[Main] Loading openset components...")
|
||||
matcher.load()
|
||||
print(f"[Main] Loaded {len(matcher.dish_names)} dish names from index")
|
||||
|
||||
res = classify_image_openset(
|
||||
image_path=args.image,
|
||||
provider=provider,
|
||||
embedder=embedder,
|
||||
matcher=matcher,
|
||||
top_k=args.openset_top_k,
|
||||
min_match_score=args.openset_min_score,
|
||||
)
|
||||
|
||||
print(
|
||||
json.dumps(
|
||||
{"mode": args.mode, "result": res},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
)
|
||||
|
||||
# 封闭式识别分支(原有逻辑)
|
||||
else:
|
||||
labels = build_labels(args.mode, args.alias_map)
|
||||
ingredient_only = args.mode in {"whole", "processed"}
|
||||
fewshot = FEWSHOT_HINTS if args.mode == "dish" else None
|
||||
print(f"[Main] Built labels count={len(labels)} ingredient_only={ingredient_only} fewshot={bool(fewshot)}")
|
||||
|
||||
res = classify_image(
|
||||
args.image,
|
||||
labels,
|
||||
fewshot_hints=fewshot,
|
||||
ingredient_only=ingredient_only,
|
||||
)
|
||||
provider = OllamaVLMProvider(base_url=OLLAMA_URL, model_name=DEFAULT_MODEL)
|
||||
|
||||
res = classify_image(
|
||||
args.image,
|
||||
labels,
|
||||
provider=provider,
|
||||
fewshot_hints=fewshot,
|
||||
ingredient_only=ingredient_only,
|
||||
)
|
||||
|
||||
print(
|
||||
json.dumps(
|
||||
{"mode": args.mode, "result": res, "num_labels": len(labels)},
|
||||
ensure_ascii=False,
|
||||
print(
|
||||
json.dumps(
|
||||
{"mode": args.mode, "result": res, "num_labels": len(labels)},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -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()
|
||||
@@ -38,3 +38,31 @@ def build_closedset_prompt(
|
||||
)
|
||||
|
||||
|
||||
def build_openset_prompt() -> str:
|
||||
"""
|
||||
构建开放式识别提示词(不提供候选列表,让VLM自由识别)
|
||||
用于"开放识别+向量匹配"方案
|
||||
"""
|
||||
return (
|
||||
"【重要】你必须且只能输出JSON格式,严禁输出任何其他内容(包括英文描述、解释、分析等)。\n\n"
|
||||
"你是专业的中国菜品识别专家。任务:识别图片中的菜品并输出标准中文名称。\n\n"
|
||||
"要求:\n"
|
||||
"1. 只输出菜品的标准中文名称,例如:宫保鸡丁、红烧肉、西红柿炒鸡蛋\n"
|
||||
"2. 禁止输出地域/菜系分类(如川菜、粤菜)\n"
|
||||
"3. 禁止输出烹饪方式后缀(如炒制、烹饪),除非该方式是菜名的一部分(如红烧肉)\n"
|
||||
"4. 禁止输出括号注释(如(川菜)、(辣味))\n"
|
||||
"5. 禁止输出英文翻译\n"
|
||||
"6. 给出0-1的置信度分数,表示你对识别结果的确信程度\n"
|
||||
"7. 严格按以下JSON格式输出,不允许有任何偏差\n\n"
|
||||
'输出格式(必须严格遵守):\n{"dish":"<菜品标准中文名>","confidence":0.95}\n\n'
|
||||
"示例输出:\n"
|
||||
'{"dish":"宫保鸡丁","confidence":0.92}\n'
|
||||
'{"dish":"红烧肉","confidence":0.88}\n'
|
||||
'{"dish":"西红柿炒鸡蛋","confidence":0.95}\n\n'
|
||||
"【再次强调】\n"
|
||||
"- 只输出上述JSON格式,不要有任何额外文字\n"
|
||||
"- 菜品名必须是标准中文,不带任何前缀、后缀、括号、英文\n"
|
||||
"- 如果无法识别,输出confidence接近0的值"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
import json
|
||||
import re
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
from .vlm_providers.base import VLMProvider
|
||||
from .prompts import build_closedset_prompt
|
||||
from .prompts import build_closedset_prompt, build_openset_prompt
|
||||
from .dish_name_cleaner import normalize_dish_name
|
||||
from .text_embedder import OllamaEmbedder
|
||||
from .vector_matcher import DishNameMatcher
|
||||
|
||||
|
||||
def _extract_json_obj(text: str):
|
||||
@@ -93,3 +96,133 @@ def classify_image(
|
||||
return {"label": "Unknown", "confidence": 0.0}
|
||||
|
||||
|
||||
def classify_image_openset(
|
||||
image_path: str,
|
||||
provider: VLMProvider,
|
||||
embedder: OllamaEmbedder,
|
||||
matcher: DishNameMatcher,
|
||||
top_k: int = 3,
|
||||
min_match_score: float = 0.5,
|
||||
) -> Dict:
|
||||
"""
|
||||
开放式菜品识别 + 向量匹配方案
|
||||
|
||||
流程:
|
||||
1. VLM开放式识别(不给候选列表)
|
||||
2. 清洗菜品名
|
||||
3. 文本向量化
|
||||
4. FAISS检索Top-K候选
|
||||
5. 综合置信度排序
|
||||
|
||||
参数:
|
||||
image_path: 图片路径
|
||||
provider: VLM Provider实例
|
||||
embedder: Ollama Embedder实例
|
||||
matcher: 菜品名匹配器实例
|
||||
top_k: 向量检索返回的候选数
|
||||
min_match_score: 最低匹配分数阈值
|
||||
|
||||
返回:
|
||||
{
|
||||
"raw_dish": "VLM原始输出",
|
||||
"cleaned_dish": "清洗后的菜品名",
|
||||
"vlm_confidence": 0.95,
|
||||
"candidates": [
|
||||
{"dish": "匹配的标准菜名", "match_score": 0.88, "final_score": 0.83},
|
||||
...
|
||||
],
|
||||
"best_match": "最佳匹配菜名",
|
||||
"final_confidence": 0.83
|
||||
}
|
||||
"""
|
||||
print(f"[VLM_Openset] Start classify image={image_path} top_k={top_k}")
|
||||
|
||||
# 步骤1: VLM开放式识别
|
||||
prompt = build_openset_prompt()
|
||||
print(f"[VLM_Openset] Prompt length={len(prompt)}")
|
||||
|
||||
print(f"[VLM_Openset] Calling {provider}.chat_vision()...")
|
||||
text = provider.chat_vision(prompt, [image_path], temperature=0.1)
|
||||
_preview = (str(text)[:200]).replace("\n", " ")
|
||||
print(f"[VLM_Openset] Received text length={len(str(text))} preview={_preview}...")
|
||||
|
||||
# 步骤2: 解析VLM输出
|
||||
obj = _extract_json_obj(str(text))
|
||||
if not obj or "dish" not in obj:
|
||||
print("[VLM_Openset] JSON parse failed or missing 'dish' field")
|
||||
return {
|
||||
"raw_dish": str(text),
|
||||
"cleaned_dish": "",
|
||||
"vlm_confidence": 0.0,
|
||||
"candidates": [],
|
||||
"best_match": "Unknown",
|
||||
"final_confidence": 0.0,
|
||||
}
|
||||
|
||||
raw_dish = obj.get("dish", "")
|
||||
vlm_conf = obj.get("confidence", 0.0)
|
||||
try:
|
||||
vlm_conf = float(vlm_conf)
|
||||
except Exception:
|
||||
vlm_conf = 0.0
|
||||
|
||||
print(f"[VLM_Openset] Parsed raw_dish='{raw_dish}' vlm_confidence={vlm_conf}")
|
||||
|
||||
# 步骤3: 清洗菜品名
|
||||
cleaned_dish = normalize_dish_name(raw_dish)
|
||||
print(f"[VLM_Openset] Cleaned dish='{cleaned_dish}'")
|
||||
|
||||
if not cleaned_dish:
|
||||
print("[VLM_Openset] Cleaned dish is empty, return Unknown")
|
||||
return {
|
||||
"raw_dish": raw_dish,
|
||||
"cleaned_dish": cleaned_dish,
|
||||
"vlm_confidence": vlm_conf,
|
||||
"candidates": [],
|
||||
"best_match": "Unknown",
|
||||
"final_confidence": 0.0,
|
||||
}
|
||||
|
||||
# 步骤4: 向量匹配
|
||||
print(f"[VLM_Openset] Embedding text='{cleaned_dish}'...")
|
||||
query_emb = embedder.encode_single(cleaned_dish)
|
||||
|
||||
print(f"[VLM_Openset] Matching with FAISS index...")
|
||||
matches = matcher.match(query_emb, top_k)
|
||||
|
||||
# 步骤5: 构建候选列表(综合置信度 = VLM置信度 × 匹配分数)
|
||||
candidates = []
|
||||
for dish_name, match_score in matches:
|
||||
if match_score < min_match_score:
|
||||
print(f"[VLM_Openset] Skip candidate '{dish_name}' (score={match_score:.4f} < threshold={min_match_score})")
|
||||
continue
|
||||
|
||||
final_score = vlm_conf * match_score
|
||||
candidates.append({
|
||||
"dish": dish_name,
|
||||
"match_score": match_score,
|
||||
"final_score": final_score,
|
||||
})
|
||||
|
||||
# 按综合分数降序排序
|
||||
candidates.sort(key=lambda x: x["final_score"], reverse=True)
|
||||
|
||||
# 最佳匹配
|
||||
if candidates:
|
||||
best = candidates[0]
|
||||
best_match = best["dish"]
|
||||
final_conf = best["final_score"]
|
||||
print(f"[VLM_Openset] Best match='{best_match}' final_confidence={final_conf:.4f}")
|
||||
else:
|
||||
best_match = "Unknown"
|
||||
final_conf = 0.0
|
||||
print("[VLM_Openset] No valid candidates, return Unknown")
|
||||
|
||||
return {
|
||||
"raw_dish": raw_dish,
|
||||
"cleaned_dish": cleaned_dish,
|
||||
"vlm_confidence": vlm_conf,
|
||||
"candidates": candidates,
|
||||
"best_match": best_match,
|
||||
"final_confidence": final_conf,
|
||||
}
|
||||
|
||||
@@ -354,8 +354,8 @@ def grid_search_main(
|
||||
if __name__ == '__main__':
|
||||
import argparse
|
||||
parser = argparse.ArgumentParser(description='CosFace超参数网格搜索')
|
||||
# parser.add_argument('--task', choices=list(TASKS.keys()), default='dish', help='任务名称')
|
||||
parser.add_argument('--task', choices=list(TASKS.keys()), default='whole_ingredient', help='任务名称')
|
||||
parser.add_argument('--task', choices=list(TASKS.keys()), default='dish', help='任务名称')
|
||||
# parser.add_argument('--task', choices=list(TASKS.keys()), default='whole_ingredient', help='任务名称')
|
||||
parser.add_argument('--max_configs', type=int, default=None, help='最大配置数(用于测试)')
|
||||
parser.add_argument('--epochs', type=int, default=100, help='每个配置的最大训练轮数')
|
||||
parser.add_argument('--patience', type=int, default=10, help='早停容忍轮数')
|
||||
|
||||
Reference in New Issue
Block a user