"""VLM OCR 节点。
调用本地 Ollama 服务的多模态模型(默认 glm-ocr:latest)对图片做文字识别
(OCR),用于从视频帧中提取字幕文字,构建"真实音频 + 正确字幕"的测试数据。
与 llm-translate 节点相互独立:本节点只做视觉 OCR,不做翻译,避免把两类
职责混在一起。调用协议见 Ollama 官方文档:POST /api/chat。
请求约定:
- 流式传输(stream=True)逐行接收,命中终止序列立即停止拉取,避免模型重复
循环时无限输出;整体受超时截止约束(timeout_seconds / VLM_TIMEOUT_SECONDS,
默认 5 秒),超时即终止而不继续等待后续块。
- 采样 temperature 默认 0.3,并携带终止序列 `["\n", "\n答", "答"]`(输出
首个换行即停,同时阻止“答:”式重复循环),模型遇到任一标记即停止。
"""
from __future__ import annotations
import base64
import json
import os
import re
import time
import urllib.error
import urllib.request
from pathlib import Path
from wov_sdk.models import InvokeRequest, InvokeResponse
# 终止序列列表:模型输出一旦出现任一标记(如换行 "\n"、"答")立即停止生成
# (Ollama 的 stop 参数,同时作为流式读取时的本地截断标记)。
# 语义:输出第一个换行即停(单行字幕),"答" 系列阻止"答:"式重复循环。
STOP_SEQUENCE = ["\n","\n答","答"]
def _default_host() -> str:
"""默认 Ollama 服务地址,可通过环境变量覆盖。"""
return os.getenv("OLLAMA_HOST", "http://192.168.123.70:11434")
def _clean_ocr_text(content: str) -> str:
"""清洗模型输出:去掉 markdown 围栏与空行,只保留识别到的文字。
glm-ocr 的自定义解析器会在识别文本后追加大量 ``` 围栏,必须剔除,
否则会污染字幕文本。
"""
lines: list[str] = []
for line in content.splitlines():
stripped = line.strip()
# 整行只有反引号(可带语言标记,如 ```markdown)的围栏行丢弃。
if re.fullmatch(r"`+[A-Za-z0-9]*", stripped):
continue
if stripped:
lines.append(stripped)
return "\n".join(lines)
# 标签提取正则:匹配 ...(DOTALL 让 . 也能匹配换行)。
_GETTEXT_RE = re.compile(r"(.*?)", re.DOTALL)
def _extract_gettext(raw: str) -> str:
"""从模型输出中提取 标签内的内容。
提示词要求模型用该标签包裹识别结果;模型未按格式输出(找不到标签)
时回退返回原始文本,保持旧行为。多个标签只取第一个(防重复循环)。
"""
match = _GETTEXT_RE.search(raw)
if match is None:
return raw
return match.group(1)
def _truncate_at_stop(raw: str) -> str:
"""在最先命中的终止序列处截断文本。
STOP_SEQUENCE 为终止序列列表(如 ["\n", "\n答", "答"]):取最早出现的位置
截断,返回截断后的文本;未命中任何序列时原样返回。
"""
positions = [pos for seq in STOP_SEQUENCE if (pos := raw.find(seq)) != -1]
if not positions:
return raw
return raw[: min(positions)]
def _consume_stream(response, deadline: float) -> str:
"""逐行读取流式响应,直到命中终止序列 / 流结束 / 超过截止时间。
response 为 urllib 打开的响应对象,readline() 返回 bytes(每行一个
Ollama 流式 JSON 块)。返回拼接后的原始文本(未清洗、未截断终止序列)。
"""
parts: list[str] = []
while True:
# 每次调用整体 5 秒上限:超过立即终止,不再等待下一个流式块。
if time.monotonic() >= deadline:
raise TimeoutError("timed out")
line = response.readline()
if not line:
# 流式输出正常结束(模型未发出终止序列,直接收完)。
break
chunk = json.loads(line.decode("utf-8"))
message = chunk.get("message")
if message is None:
# done 行是流结束标记,可能不带 message;其余缺失视为格式错误。
if chunk.get("done"):
break
raise KeyError("missing message in stream chunk")
parts.append(str(message.get("content", "")))
# 命中任一终止序列即停止接收后续内容(模型应已结束生成)。
if any(seq in "".join(parts) for seq in STOP_SEQUENCE):
break
if chunk.get("done"):
# 模型侧完成(未命中终止序列也自然结束)。
break
return "".join(parts)
def invoke(request: InvokeRequest) -> InvokeResponse:
"""识别 image_uri 指向的图片中的文字,产物为 ocr.txt。
参数:model(默认 glm-ocr:latest)、ollama_host、prompt、
timeout_seconds(默认 5,每次调用整体上限);均可通过环境变量
VLM_MODEL / OLLAMA_HOST / VLM_TIMEOUT_SECONDS 覆盖。
"""
image_uri = request.inputs.get("image_uri")
if not image_uri:
return InvokeResponse(status="failed", error="image_uri is required")
image_path = Path(image_uri)
# 文件不存在时提前失败,避免无谓的网络请求。
if not image_path.is_file():
return InvokeResponse(status="failed", error="image file not found")
host = str(request.params.get("ollama_host") or _default_host())
model = str(request.params.get("model") or os.getenv("VLM_MODEL", "glm-ocr:latest"))
prompt = str(
request.params.get("prompt")
or os.getenv("VLM_PROMPT", "提取图像中的文字,不要描述图片中的内容")
)
timeout = float(
request.params.get("timeout_seconds")
or os.getenv("VLM_TIMEOUT_SECONDS", "5")
)
# 图片按 base64 随请求体发送(Ollama 多模态标准格式)。
image_b64 = base64.b64encode(image_path.read_bytes()).decode("ascii")
body = {
"model": model,
# 流式传输:逐行接收生成内容,命中终止序列或超时即停止。
"stream": True,
# keep_alive 让模型在服务端常驻,避免逐帧调用反复加载模型。
"keep_alive": str(request.params.get("keep_alive", "5m")),
# 采样选项:temperature 默认 0.3(可参数覆盖);glm-ocr 在大图上有已知
# 重复循环 bug,repeat_penalty 惩罚重复 token、num_predict 限制输出上限。
"options": {
"temperature": float(request.params.get("temperature", 0.3)),
"repeat_penalty": float(request.params.get("repeat_penalty", 1)),
"num_predict": int(request.params.get("num_predict", 256)),
},
# /api/chat 的输入结构:识别指令放系统提示词,用户消息只携带图片
# (content 为空、images 传 base64,与 glm-ocr 期望结构一致)。
"messages": [
{"role": "system", "content": prompt},
{"role": "user", "content": "", "images": [image_b64]},
],
# 模型遇到任一终止序列即停止生成,遏制重复循环;同时作为流式读取截断点。
"stop": STOP_SEQUENCE,
}
request_url = f"{host.rstrip('/')}/api/chat"
http_request = urllib.request.Request(
request_url,
data=json.dumps(body).encode("utf-8"),
headers={"Content-Type": "application/json"},
method="POST",
)
try:
# timeout 同时作用于连接与每次 socket 读取;deadline 保证整体 5 秒上限。
deadline = time.monotonic() + timeout
with urllib.request.urlopen(http_request, timeout=timeout) as response:
raw = _consume_stream(response, deadline)
# 终止序列可能随最后一个流式块一起返回,在最先命中的序列处截断再清洗。
raw = _truncate_at_stop(raw)
# 提取提示词约定的 标签内容;模型未按格式输出时回退原始文本。
text = _clean_ocr_text(_extract_gettext(raw))
output_dir = Path(request.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
output_path = output_dir / "ocr.txt"
output_path.write_text(text + "\n", encoding="utf-8")
return InvokeResponse(
status="completed",
outputs={"text": text, "text_uri": str(output_path)},
)
except (urllib.error.URLError, KeyError, ValueError, OSError) as exc:
# 网络失败、响应格式异常、超时等统一转换为 failed 响应。
return InvokeResponse(status="failed", error=str(exc))