Files
vrsub/nodes/llm_filter.py
T
cat-shark 4746e0363f feat: VRSub 单体应用(WOV 单机版)初始提交
为视频生成 VR 双眼字幕的单体实现:FastAPI 后端、调度器与全部节点
(提音/转写/翻译/ASS/抽帧/OCR/LLM 过滤)在单进程内运行。

- 节点协议(wov_sdk 数据模型)与分布式版保持一致,预留回退桥梁
- 工作流即数据:DAG 存于 workflows/*.json,模型/链路改动只改数据
- 调度器:拓扑顺序执行、断点续跑(产物重建)、任务暂停/继续
- 抽帧按帧间隔(select 按帧号精确取帧),VLM OCR 与 LLM 过滤使用
  自适应线程池弹性并发,并打印数据处理速度进度日志
- 100% 行覆盖率(pytest --cov-fail-under=100)
2026-08-16 23:58:25 +08:00

172 lines
7.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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},
)