Files
contentm_agent/ocr_subs.py
T

98 lines
3.6 KiB
Python
Raw Normal View History

# -*- coding: utf-8 -*-
"""对关键帧字幕区做 OCR,按时间整理成字幕文档。
依赖:rapidocr-onnxruntime(中文 OCR)。复用 cluster_subs 的字幕区二值哈希聚类,只 OCR 代表帧。
用法:python ocr_subs.py --frames <frames目录> --video <源视频名/说明> --out <输出.md> [--thresh 25] [--crop 0.15]
"""
import argparse, glob, os, re, sys, difflib
import numpy as np
from PIL import Image
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from cluster_subs import bin_hash, hamming, crop_subtitle, ts_from_name
def fmt_ts(t):
m = int(t // 60); s = t - m * 60
return "%02d:%06.3f" % (m, s)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--frames", required=True)
ap.add_argument("--video", required=True, help="源视频说明(写入文档头)")
ap.add_argument("--out", required=True)
ap.add_argument("--thresh", type=int, default=25)
ap.add_argument("--crop", type=float, default=0.15)
args = ap.parse_args()
from rapidocr_onnxruntime import RapidOCR
engine = RapidOCR()
files = sorted(glob.glob(os.path.join(args.frames, "*.jpg")))
print("总帧数: %d" % len(files))
# 聚类取代表帧
clusters = []
for f in files:
sub = crop_subtitle(Image.open(f), args.crop)
hsh = bin_hash(sub)
best, bestd = None, 999
for ci, c in enumerate(clusters):
d = hamming(hsh, c["hash"])
if d < bestd:
bestd, best = d, ci
ts = ts_from_name(os.path.basename(f))
if best is not None and bestd <= args.thresh:
clusters[best]["items"].append((f, ts))
else:
clusters.append({"hash": hsh, "items": [(f, ts)]})
reps = []
for c in clusters:
it = sorted(c["items"], key=lambda x: x[1])
mid = it[len(it) // 2]
reps.append((mid[0], mid[1]))
reps.sort(key=lambda x: x[1])
print("代表帧(待 OCR): %d" % len(reps))
lines = []
for f, ts in reps:
sub = crop_subtitle(Image.open(f), args.crop)
sub = sub.resize((sub.width * 2, sub.height * 2), Image.LANCZOS)
res = engine(np.array(sub))
if not res or not res[0]:
continue
txt = " ".join([t[1] for t in res[0] if t[1].strip()]).strip()
txt = re.sub(r"\s+", " ", txt)
if not txt:
continue
# 丢弃纯符号/数字垃圾行(几乎无中文字符)
cjk = len(re.findall(r"[一-鿿]", txt))
if cjk < 2:
continue
lines.append((ts, txt))
# 模糊合并连续近重复句(OCR 抖动导致同句多版本)
def sim(a, b):
return difflib.SequenceMatcher(None, a, b).ratio()
merged = []
for ts, txt in lines:
if merged and sim(txt, merged[-1][1]) >= 0.82:
# 保留更长更完整的一句,时间取最早
if len(txt) > len(merged[-1][1]):
merged[-1] = (merged[-1][0], txt)
continue
merged.append((ts, txt))
lines = merged
with open(args.out, "w", encoding="utf-8") as fo:
fo.write("# 字幕文档(关键帧 OCR)\n\n")
fo.write("源视频:%s\n\n" % args.video)
fo.write("> 说明:视频无独立字幕轨,本字幕由关键帧画面字幕区 OCR 提取,按出现时间排序;连续近重复句(相似度≥0.82)已合并、纯符号垃圾行已剔除。\n\n")
for ts, txt in lines:
fo.write("- [%s] %s\n" % (fmt_ts(ts), txt))
print("已写出 %d 条字幕 → %s" % (len(lines), args.out))
if __name__ == "__main__":
main()