为视频生成 VR 双眼字幕的单体实现:FastAPI 后端、调度器与全部节点 (提音/转写/翻译/ASS/抽帧/OCR/LLM 过滤)在单进程内运行。 - 节点协议(wov_sdk 数据模型)与分布式版保持一致,预留回退桥梁 - 工作流即数据:DAG 存于 workflows/*.json,模型/链路改动只改数据 - 调度器:拓扑顺序执行、断点续跑(产物重建)、任务暂停/继续 - 抽帧按帧间隔(select 按帧号精确取帧),VLM OCR 与 LLM 过滤使用 自适应线程池弹性并发,并打印数据处理速度进度日志 - 100% 行覆盖率(pytest --cov-fail-under=100)
172 lines
7.2 KiB
Python
172 lines
7.2 KiB
Python
"""LLM 字幕过滤节点。
|
||
|
||
对 OCR 识别出的 SRT 字幕做二次过滤:为避免字幕上下文过长,把每条字幕连同其
|
||
前后各 context_size 条字幕(纯文本,**不含时间戳**)分批提供给 LLM,模型仅判断
|
||
目标字幕是否属于多余、无意义的字符(如重复、残缺、无实际语义的杂项);
|
||
判定无意义则删除该条(连同其时间戳),其余字幕保持原样并重新编号输出。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import os
|
||
import re
|
||
import urllib.request
|
||
from pathlib import Path
|
||
from nodes.adaptive_pool import AdaptiveThreadPool
|
||
from wov_app.logging import get_logger
|
||
from wov_sdk.models import InvokeRequest, InvokeResponse
|
||
|
||
logger = get_logger("llm-filter")
|
||
|
||
# 匹配 SRT 条目:时间轴行 + 文本(文本可多行),到下一个序号行或文末结束。
|
||
_SRT_BLOCK_RE = re.compile(
|
||
r"(\d{2}:\d{2}:\d{2},\d{3})\s*-->\s*(\d{2}:\d{2}:\d{2},\d{3})\s*\n(.*?)(?=\n\s*\d+\s*\n|\Z)",
|
||
re.DOTALL,
|
||
)
|
||
|
||
# 目标字幕标记:提示词用该标记指明需要判断的那一条字幕。
|
||
TARGET_MARK = "【目标】"
|
||
|
||
# 默认上下文窗口:目标字幕前后各取 10 条。
|
||
DEFAULT_CONTEXT_SIZE = 10
|
||
|
||
|
||
def parse_srt(text: str) -> list[dict]:
|
||
"""解析 SRT 文本为条目列表:[{"start", "end", "text"}]。"""
|
||
entries: list[dict] = []
|
||
for match in _SRT_BLOCK_RE.finditer(text):
|
||
entries.append(
|
||
{
|
||
"start": match.group(1),
|
||
"end": match.group(2),
|
||
"text": match.group(3).strip(),
|
||
}
|
||
)
|
||
return entries
|
||
|
||
|
||
def serialize_srt(entries: list[dict]) -> str:
|
||
"""把条目列表序列化为标准 SRT 文本(序号重新从 1 编号)。"""
|
||
blocks = [
|
||
f"{index}\n{entry['start']} --> {entry['end']}\n{entry['text']}"
|
||
for index, entry in enumerate(entries, start=1)
|
||
]
|
||
return "\n\n".join(blocks) + "\n"
|
||
|
||
|
||
def _judge_target(
|
||
entries: list[dict], index: int, context_size: int, params: dict
|
||
) -> bool:
|
||
"""调用 LLM 判断目标字幕是否多余/无意义;返回 True 表示应删除。
|
||
|
||
请求体只含目标字幕及其前后各 context_size 条字幕的纯文本(无时间戳),
|
||
目标字幕用 TARGET_MARK 标记;模型只需回答"保留"或"删除"。
|
||
"""
|
||
start = max(0, index - context_size)
|
||
end = min(len(entries), index + context_size + 1)
|
||
target_pos = index - start
|
||
lines = [
|
||
f"{TARGET_MARK}{text}" if pos == target_pos else text
|
||
for pos, text in enumerate(entry["text"] for entry in entries[start:end])
|
||
]
|
||
|
||
# LLM 兼容接口配置:地址/Key/模型/超时均可通过环境变量覆盖(默认 SiliconFlow)。
|
||
api_base = os.getenv(
|
||
"LLM_API_BASE",
|
||
"https://api.siliconflow.cn/v1/chat/completions",
|
||
)
|
||
api_key = os.getenv("LLM_API_KEY", "")
|
||
request_timeout = float(os.getenv("LLM_TIMEOUT_SECONDS", "60"))
|
||
model = str(params.get("model") or os.getenv("LLM_MODEL", "Qwen/Qwen3.6-35B-A3B"))
|
||
system_prompt = (
|
||
"你是字幕质量过滤器。用户会提供一段字幕序列(纯文本,不含时间戳),"
|
||
f"其中用{TARGET_MARK}标记的字幕是需要判断的目标。请判断该字幕是否属于"
|
||
"多余、无意义的字符(如重复、残缺、无实际语义的杂项)。"
|
||
"只回答两个字:保留 或 删除,不要输出其他内容。"
|
||
)
|
||
body = {
|
||
"model": model,
|
||
"messages": [
|
||
{"role": "system", "content": system_prompt},
|
||
{"role": "user", "content": "\n".join(lines)},
|
||
],
|
||
# 关闭推理模式:Qwen3 等模型默认会把思考过程写入 reasoning_content,
|
||
# 导致 content 为空或包含多余内容。
|
||
"enable_thinking": False,
|
||
# 只输出"保留/删除",输出上限给得很小即可。
|
||
"max_tokens": 16,
|
||
}
|
||
headers = {"Content-Type": "application/json"}
|
||
# 配置了 Key 时附带 Bearer 鉴权头。
|
||
if api_key:
|
||
headers["Authorization"] = f"Bearer {api_key}"
|
||
request = urllib.request.Request(
|
||
api_base,
|
||
data=json.dumps(body).encode("utf-8"),
|
||
headers=headers,
|
||
method="POST",
|
||
)
|
||
with urllib.request.urlopen(request, timeout=request_timeout) as response:
|
||
payload = json.loads(response.read().decode("utf-8"))
|
||
content = str(payload["choices"][0]["message"]["content"])
|
||
# 模型回答含"删除"即视为该条无意义;其余情况(保留/异常)一律保留,宁多勿删。
|
||
return "删除" in content
|
||
|
||
|
||
def invoke(request: InvokeRequest) -> InvokeResponse:
|
||
"""过滤 SRT 中多余/无意义的字幕,产物为 filtered.srt。"""
|
||
srt_uri = request.inputs.get("srt_uri")
|
||
if not srt_uri:
|
||
return InvokeResponse(status="failed", error="srt_uri is required")
|
||
srt_path = Path(srt_uri)
|
||
if not srt_path.is_file():
|
||
return InvokeResponse(status="failed", error="srt file not found")
|
||
|
||
entries = parse_srt(srt_path.read_text(encoding="utf-8"))
|
||
context_size = int(request.params.get("context_size", DEFAULT_CONTEXT_SIZE))
|
||
# 单条判断的工作函数:返回 True 表示该条应删除。
|
||
def judge_one(index) -> bool:
|
||
return _judge_target(entries, index, context_size, request.params)
|
||
|
||
# 进度日志:打印已判定条数、总数与平均处理速度(条/s)。
|
||
def log_progress(done: int, total: int, rate: float) -> None:
|
||
logger.info("字幕判定进度 %d/%d 条 (%.1f 条/s)", done, total, rate)
|
||
|
||
# 自适应并发调用 LLM:10s 窗口内平均响应 < 0.3s 则加 1 线程(上限
|
||
# pool_max_workers),> pool_slow_threshold 则减 1 线程(下限 1),
|
||
# 按实测负载弹性伸缩,避免盲目并发压垮 LLM 接口。
|
||
pool = AdaptiveThreadPool(
|
||
worker=judge_one,
|
||
on_progress=log_progress,
|
||
min_workers=int(request.params.get("pool_min_workers", 1)),
|
||
max_workers=int(request.params.get("pool_max_workers", 16)),
|
||
window_seconds=float(request.params.get("pool_window_seconds", 10.0)),
|
||
fast_threshold=float(request.params.get("pool_fast_threshold", 0.3)),
|
||
slow_threshold=float(request.params.get("pool_slow_threshold", 1.0)),
|
||
)
|
||
verdicts = pool.map(range(len(entries)))
|
||
|
||
kept: list[dict] = []
|
||
removed = 0
|
||
for index, (entry, verdict) in enumerate(zip(entries, verdicts)):
|
||
# 并行下 LLM 异常被线程池隔离为异常结果:任一条失败即整体失败,
|
||
# 避免静默输出未过滤结果。
|
||
if isinstance(verdict, Exception):
|
||
return InvokeResponse(status="failed", error=str(verdict))
|
||
if verdict:
|
||
removed += 1
|
||
logger.info("删除无意义字幕 %d: %r", index + 1, entry["text"][:40])
|
||
else:
|
||
kept.append(entry)
|
||
|
||
output_dir = Path(request.output_dir)
|
||
output_dir.mkdir(parents=True, exist_ok=True)
|
||
output_path = output_dir / "filtered.srt"
|
||
output_path.write_text(serialize_srt(kept), encoding="utf-8")
|
||
logger.info("字幕过滤完成: 保留 %d 条, 删除 %d 条", len(kept), removed)
|
||
return InvokeResponse(
|
||
status="completed",
|
||
outputs={"srt_uri": str(output_path), "kept": len(kept), "removed": removed},
|
||
)
|