"""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}, )