增加多模态模型识别功能,可以进行识别了,程序跑通了
This commit is contained in:
@@ -25,7 +25,7 @@ FEWSHOT_HINTS: Dict[str, str] = {
|
|||||||
def main() -> None:
|
def main() -> None:
|
||||||
ap = argparse.ArgumentParser()
|
ap = argparse.ArgumentParser()
|
||||||
ap.add_argument("--mode", choices=["dish", "whole", "processed"], default="dish")
|
ap.add_argument("--mode", choices=["dish", "whole", "processed"], default="dish")
|
||||||
ap.add_argument("--image", default=r"D:\MyProjects\PythonProjects\FoodClassifier\dataset\DishClassification\test\回锅肉\img01.png")
|
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("--alias_map", default=None)
|
||||||
|
|
||||||
args = ap.parse_args()
|
args = ap.parse_args()
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import base64
|
import base64
|
||||||
|
import mimetypes
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from typing import Any, Dict, List
|
from typing import Any, Dict, List
|
||||||
@@ -11,9 +12,18 @@ DEFAULT_MODEL = os.getenv("VLM_MODEL", "qwen2.5vl:32b")
|
|||||||
|
|
||||||
|
|
||||||
def _img_to_data_uri(path: str) -> str:
|
def _img_to_data_uri(path: str) -> str:
|
||||||
|
mime, _ = mimetypes.guess_type(path)
|
||||||
|
if not mime:
|
||||||
|
ext = os.path.splitext(path)[1].lower()
|
||||||
|
if ext in {".jpg", ".jpeg"}:
|
||||||
|
mime = "image/jpeg"
|
||||||
|
elif ext == ".png":
|
||||||
|
mime = "image/png"
|
||||||
|
else:
|
||||||
|
mime = "application/octet-stream"
|
||||||
with open(path, "rb") as f:
|
with open(path, "rb") as f:
|
||||||
b64 = base64.b64encode(f.read()).decode("utf-8")
|
b64 = base64.b64encode(f.read()).decode("utf-8")
|
||||||
return f"data:image/jpeg;base64,{b64}"
|
return f"data:{mime};base64,{b64}"
|
||||||
|
|
||||||
|
|
||||||
def chat_vision(
|
def chat_vision(
|
||||||
@@ -21,7 +31,7 @@ def chat_vision(
|
|||||||
image_paths: List[str],
|
image_paths: List[str],
|
||||||
*,
|
*,
|
||||||
temperature: float = 0.1,
|
temperature: float = 0.1,
|
||||||
options: Dict[str, Any] | None = None,
|
options: Dict[str, Any] | None = None, # kept for API compatibility, not used in /v1 schema
|
||||||
model: str = DEFAULT_MODEL,
|
model: str = DEFAULT_MODEL,
|
||||||
timeout_sec: int = 120,
|
timeout_sec: int = 120,
|
||||||
) -> str:
|
) -> str:
|
||||||
@@ -34,39 +44,71 @@ def chat_vision(
|
|||||||
except Exception:
|
except Exception:
|
||||||
img_sizes.append(-1)
|
img_sizes.append(-1)
|
||||||
print(
|
print(
|
||||||
f"[Ollama] Prepare chat url={OLLAMA_URL}/api/chat model={model} temp={temperature} "
|
f"[Ollama] Prepare chat url={OLLAMA_URL}/v1/chat/completions model={model} temp={temperature} "
|
||||||
f"images={len(image_paths)} sizes={img_sizes} timeout={timeout_sec}s"
|
f"images={len(image_paths)} sizes={img_sizes} timeout={timeout_sec}s"
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
images = [_img_to_data_uri(p) for p in image_paths]
|
# OpenAI-compatible vision message format: content is a list of parts (text + image_url)
|
||||||
|
content_parts: List[Dict[str, Any]] = [{"type": "text", "text": prompt}]
|
||||||
|
for p in image_paths:
|
||||||
|
content_parts.append({
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {"url": _img_to_data_uri(p)}
|
||||||
|
})
|
||||||
|
|
||||||
payload: Dict[str, Any] = {
|
payload: Dict[str, Any] = {
|
||||||
"model": model,
|
"model": model,
|
||||||
"messages": [
|
"messages": [
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
"content": prompt,
|
"content": content_parts,
|
||||||
"images": images,
|
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"options": {"temperature": temperature, **(options or {})},
|
"temperature": temperature,
|
||||||
"stream": False,
|
"stream": False,
|
||||||
}
|
}
|
||||||
|
|
||||||
t0 = time.time()
|
t0 = time.time()
|
||||||
try:
|
try:
|
||||||
resp = requests.post(f"{OLLAMA_URL}/api/chat", json=payload, timeout=timeout_sec)
|
resp = requests.post(f"{OLLAMA_URL}/v1/chat/completions", json=payload, timeout=timeout_sec)
|
||||||
dt = time.time() - t0
|
dt = time.time() - t0
|
||||||
print(f"[Ollama] Response status={resp.status_code} elapsed={dt:.2f}s")
|
print(f"[Ollama] Response status={resp.status_code} elapsed={dt:.2f}s")
|
||||||
|
if resp.status_code >= 400:
|
||||||
|
body = resp.text
|
||||||
|
snippet = (body[:400]).replace("\n", " ")
|
||||||
|
print(f"[Ollama][ErrorBody] {snippet}...")
|
||||||
resp.raise_for_status()
|
resp.raise_for_status()
|
||||||
|
|
||||||
|
# Parse OpenAI-compatible response
|
||||||
|
# Expected: { choices: [ { message: { content: "..." } } ], ... }
|
||||||
data = resp.json()
|
data = resp.json()
|
||||||
text = (data.get("message") or {}).get("content", "")
|
if not isinstance(data, dict) or "choices" not in data or not data["choices"]:
|
||||||
|
snippet = (resp.text[:400]).replace("\n", " ")
|
||||||
|
print(f"[Ollama][Error] Unexpected response shape: {snippet}")
|
||||||
|
raise RuntimeError("Unexpected response from /v1/chat/completions")
|
||||||
|
|
||||||
|
choice0 = data["choices"][0]
|
||||||
|
msg = choice0.get("message") or {}
|
||||||
|
text = msg.get("content") or ""
|
||||||
|
if not text:
|
||||||
|
# Some servers might put text in 'choices[0].text' (rare for chat) – try fallback
|
||||||
|
text = choice0.get("text") or ""
|
||||||
_preview = (text[:120]).replace("\n", " ")
|
_preview = (text[:120]).replace("\n", " ")
|
||||||
print(f"[Ollama] Text length={len(text)} preview={_preview}...")
|
print(f"[Ollama] Text length={len(text)} preview={_preview}...")
|
||||||
|
if not text:
|
||||||
|
raise RuntimeError("Empty content in choices[0].message.content")
|
||||||
return text
|
return text
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
dt = time.time() - t0
|
dt = time.time() - t0
|
||||||
|
# Print body to help debugging when JSON parsing fails, etc.
|
||||||
|
try:
|
||||||
|
body = resp.text # type: ignore[name-defined]
|
||||||
|
snippet = (body[:400]).replace("\n", " ")
|
||||||
|
print(f"[Ollama][ErrorBody] {snippet}...")
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
print(f"[Ollama][Error] during request elapsed={dt:.2f}s: {e}")
|
print(f"[Ollama][Error] during request elapsed={dt:.2f}s: {e}")
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user