229 lines
7.4 KiB
Python
229 lines
7.4 KiB
Python
import json
|
|
import re
|
|
from typing import Dict, List, Optional, Tuple
|
|
|
|
from .vlm_providers.base import VLMProvider
|
|
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):
|
|
"""尽量从返回文本中提取出合法 JSON(支持代码块/前后说明)。"""
|
|
if not isinstance(text, str):
|
|
return None
|
|
s = text.strip()
|
|
# 去掉 Markdown 代码块包裹
|
|
if s.startswith("```"):
|
|
s = re.sub(r"^```[a-zA-Z]*\n|\n```$", "", s).strip()
|
|
# 直接解析
|
|
try:
|
|
obj = json.loads(s)
|
|
if isinstance(obj, dict):
|
|
return obj
|
|
except Exception:
|
|
pass
|
|
# 提取第一个花括号对象
|
|
m = re.search(r"\{[\s\S]*?\}", s)
|
|
if m:
|
|
frag = m.group(0)
|
|
try:
|
|
obj = json.loads(frag)
|
|
if isinstance(obj, dict):
|
|
return obj
|
|
except Exception:
|
|
pass
|
|
return None
|
|
|
|
|
|
def classify_image(
|
|
image_path: str,
|
|
labels: List[str],
|
|
provider: VLMProvider,
|
|
fewshot_hints: Optional[Dict[str, str]] = None,
|
|
ingredient_only: bool = False,
|
|
) -> Dict:
|
|
if not labels:
|
|
print("[VLM] Empty labels provided, return Unknown")
|
|
return {"label": "Unknown", "confidence": 0.0}
|
|
|
|
print(
|
|
f"[VLM] Start classify image={image_path} labels={len(labels)} "
|
|
f"ingredient_only={ingredient_only} fewshot={bool(fewshot_hints)} provider={provider}"
|
|
)
|
|
prompt = build_closedset_prompt(labels, fewshot_hints, ingredient_only)
|
|
# _preview_prompt = (prompt[:120]).replace("\n", " ")
|
|
final_prompt = prompt.replace("\n", " ")
|
|
# print(f"[VLM] Prompt length={len(prompt)} preview={_preview_prompt}...")
|
|
print(f"[VLM] Prompt length={len(prompt)} preview={final_prompt}")
|
|
|
|
print(f"[VLM] Calling {provider}.chat_vision()...")
|
|
text = provider.chat_vision(prompt, [image_path], temperature=0.1)
|
|
_preview_text = (str(text)[:200]).replace("\n", " ")
|
|
print(f"[VLM] Received text length={len(str(text))} preview={_preview_text}...")
|
|
|
|
# 解析 JSON(更健壮)
|
|
obj = _extract_json_obj(str(text))
|
|
if obj and "label" in obj:
|
|
label = obj.get("label", "")
|
|
# 若输出不在封闭集,尝试包含匹配;否则 Unknown
|
|
if label not in labels:
|
|
for lb in labels:
|
|
if lb in str(label):
|
|
label = lb
|
|
break
|
|
else:
|
|
label = "Unknown"
|
|
conf = obj.get("confidence", 0.0)
|
|
try:
|
|
conf = float(conf)
|
|
except Exception:
|
|
conf = 0.0
|
|
result = {"label": label, "confidence": conf}
|
|
print(f"[VLM] Parsed JSON result={result}")
|
|
return result
|
|
|
|
print("[VLM] JSON parse failed or missing 'label', enter fallback")
|
|
|
|
# 回退:字符串包含匹配(无法确定置信度时返回 0.0)
|
|
for lb in labels:
|
|
if lb in str(text):
|
|
print(f"[VLM] Fallback match label={lb}")
|
|
return {"label": lb, "confidence": 0.0}
|
|
|
|
print("[VLM] No match, return Unknown")
|
|
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,
|
|
}
|