847def4a03
--concurrency 默认 16,--reaggregate-all 强制全量重聚合。 Semaphore 限视频并发数,视频内四步串行。progress 文件 asyncio.Lock + os.replace 原子写入。熔断阈值 max(.env, concurrency*2)。 Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
480 lines
16 KiB
Python
480 lines
16 KiB
Python
#!/usr/bin/env python3
|
||
"""树修复管线:检测 + VLM 重生成 + 校验 + Q&A 反向补全。
|
||
|
||
对 store/videos/ 下所有已迁移的树执行完整修复流程:
|
||
1. detect_issues() — 扫描空字段/缺失帧
|
||
2. repair_tree() — VLM 重新描述 + 底向上级联(如有问题节点)
|
||
3. verify_tree() — 交叉校验删除幻觉
|
||
4. supplement_tree() — Q&A 反向补全注入缺失事实
|
||
5. save_json() — 覆盖保存
|
||
|
||
用法:
|
||
conda activate Video-Tree-TRM
|
||
python tools/repair_trees.py [--videos-dir store/videos] [--concurrency 4] [--dry-run]
|
||
|
||
app/core/adapters 不 import 此脚本。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import asyncio
|
||
import json
|
||
import os
|
||
import sys
|
||
import time
|
||
from pathlib import Path
|
||
|
||
# 确保项目根目录在 sys.path 中
|
||
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
||
sys.path.insert(0, str(PROJECT_ROOT))
|
||
|
||
from dotenv import load_dotenv
|
||
from loguru import logger
|
||
|
||
load_dotenv(PROJECT_ROOT / ".env")
|
||
|
||
from app.tree.index import TreeIndex
|
||
from app.tree.repair.detector import detect_issues
|
||
from app.tree.repair.regenerator import repair_tree
|
||
from app.tree.repair.supplement import supplement_tree
|
||
from app.tree.subtitle import SRTEntry, parse_srt
|
||
from app.tree.verify import verify_tree
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 日志配置:不缓存,立即输出
|
||
# ---------------------------------------------------------------------------
|
||
|
||
logger.remove()
|
||
logger.add(
|
||
sys.stderr,
|
||
format="{time:HH:mm:ss} | {level:<7} | {message}",
|
||
level="DEBUG",
|
||
colorize=True,
|
||
)
|
||
logger.add(
|
||
PROJECT_ROOT / "logs" / "repair_trees.log",
|
||
format="{time:YYYY-MM-DD HH:mm:ss} | {level:<7} | {message}",
|
||
level="DEBUG",
|
||
rotation="50 MB",
|
||
)
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 断点续跑 — progress 文件管理
|
||
# ---------------------------------------------------------------------------
|
||
|
||
PROGRESS_FILE = "repair_progress.json"
|
||
|
||
|
||
def load_progress(path: Path) -> set[str]:
|
||
"""读取 progress 文件,返回已完成视频 ID 集合。
|
||
|
||
参数:
|
||
path: progress JSON 文件路径。
|
||
|
||
返回:
|
||
已完成视频 ID 集合。文件不存在或损坏时返回空集。
|
||
"""
|
||
if not path.exists():
|
||
return set()
|
||
try:
|
||
data = json.loads(path.read_text(encoding="utf-8"))
|
||
return set(data.get("finished_video_ids", []))
|
||
except (json.JSONDecodeError, KeyError, TypeError):
|
||
logger.warning("progress 文件损坏,忽略: {}", path)
|
||
return set()
|
||
|
||
|
||
async def save_progress(path: Path, lock: asyncio.Lock, vid: str) -> None:
|
||
"""原子追加一个视频 ID 到 progress 文件。
|
||
|
||
参数:
|
||
path: progress JSON 文件路径。
|
||
lock: asyncio.Lock,防并发读改写丢更新。
|
||
vid: 要追加的视频 ID。
|
||
"""
|
||
async with lock:
|
||
finished = load_progress(path)
|
||
finished.add(vid)
|
||
tmp = path.with_suffix(".tmp")
|
||
tmp.write_text(
|
||
json.dumps({"finished_video_ids": sorted(finished)}, ensure_ascii=False, indent=2),
|
||
encoding="utf-8",
|
||
)
|
||
os.replace(str(tmp), str(path))
|
||
|
||
|
||
def should_skip_video(vid: str, finished: set[str], *, reaggregate_all: bool) -> bool:
|
||
"""判断是否跳过该视频。
|
||
|
||
参数:
|
||
vid: 视频 ID。
|
||
finished: progress 中已完成的视频 ID 集合。
|
||
reaggregate_all: --reaggregate-all 标志。
|
||
|
||
返回:
|
||
True 表示跳过。
|
||
"""
|
||
if reaggregate_all:
|
||
return False
|
||
return vid in finished
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# LLM/VLM 客户端构建
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _build_clients(concurrency: int = 16):
|
||
"""构建 GovernedLLMClient(LLM + VLM)。
|
||
|
||
返回:
|
||
(llm_client, vlm_client) 元组。
|
||
"""
|
||
from adapters.breaker import CircuitBreaker
|
||
from adapters.llm import GovernedLLMClient
|
||
from adapters.telemetry import SQLiteTelemetryRecorder
|
||
from adapters.vlm import GovernedVLMClient
|
||
|
||
# 遥测记录器(GovernedLLMClient 要求非 None)
|
||
(PROJECT_ROOT / "logs").mkdir(exist_ok=True)
|
||
telemetry = SQLiteTelemetryRecorder(str(PROJECT_ROOT / "logs" / "repair_telemetry.db"))
|
||
|
||
breaker_threshold = int(os.getenv("LLM_CIRCUIT_BREAKER_THRESHOLD", "5"))
|
||
breaker_threshold = max(breaker_threshold, concurrency * 2)
|
||
breaker_cooldown = int(os.getenv("LLM_CIRCUIT_BREAKER_COOLDOWN", "60"))
|
||
timeout_s = float(os.getenv("LLM_TIMEOUT", "120"))
|
||
max_retries = int(os.getenv("LLM_MAX_RETRIES", "3"))
|
||
base_delay = float(os.getenv("LLM_RETRY_BASE_DELAY", "2.0"))
|
||
max_delay = float(os.getenv("LLM_RETRY_MAX_DELAY", "30.0"))
|
||
ttft = float(os.getenv("LLM_TTFT_TIMEOUT", "30"))
|
||
inter_token = float(os.getenv("LLM_INTER_TOKEN_TIMEOUT", "15"))
|
||
|
||
# LLM 客户端(用于 supplement 和 L2/L1 重生成)
|
||
llm = GovernedLLMClient(
|
||
model=os.environ["SEARCH_LLM_MODEL"],
|
||
base_url=os.environ["SEARCH_LLM_BASE_URL"],
|
||
api_key=os.environ["SEARCH_LLM_API_KEY"],
|
||
provider="deepseek",
|
||
thinking=False,
|
||
breaker=CircuitBreaker(fail_threshold=breaker_threshold, cooldown_s=breaker_cooldown),
|
||
cache=None,
|
||
telemetry=telemetry,
|
||
timeout_s=timeout_s,
|
||
ttft_timeout_s=ttft,
|
||
inter_token_timeout_s=inter_token,
|
||
max_retries=max_retries,
|
||
retry_base_delay_s=base_delay,
|
||
retry_max_delay_s=max_delay,
|
||
)
|
||
|
||
# VLM 客户端(用于 L3 帧重新描述)
|
||
vlm_base = GovernedLLMClient(
|
||
model=os.environ["VL_LLM_MODEL"],
|
||
base_url=os.environ["VL_LLM_BASE_URL"],
|
||
api_key=os.environ["VL_LLM_API_KEY"],
|
||
provider="qwen",
|
||
thinking=False,
|
||
breaker=CircuitBreaker(fail_threshold=breaker_threshold, cooldown_s=breaker_cooldown),
|
||
cache=None,
|
||
telemetry=telemetry,
|
||
timeout_s=timeout_s,
|
||
ttft_timeout_s=ttft,
|
||
inter_token_timeout_s=inter_token,
|
||
max_retries=max_retries,
|
||
retry_base_delay_s=base_delay,
|
||
retry_max_delay_s=max_delay,
|
||
)
|
||
vlm = GovernedVLMClient(vlm_base)
|
||
|
||
return llm, vlm
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 单视频修复
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
async def _repair_one_video(
|
||
vid: str,
|
||
tree_path: Path,
|
||
frames_dir: Path,
|
||
srt_dir: Path,
|
||
questions_dir: Path,
|
||
llm,
|
||
vlm,
|
||
*,
|
||
dry_run: bool = False,
|
||
) -> dict:
|
||
"""修复单个视频的树。
|
||
|
||
参数:
|
||
vid: 视频 ID。
|
||
tree_path: tree.json 路径。
|
||
frames_dir: 帧文件目录。
|
||
srt_dir: SRT 字幕目录。
|
||
questions_dir: 问题 JSON 目录。
|
||
llm: LLMProvider 实例。
|
||
vlm: VLMProvider 实例。
|
||
dry_run: 仅检测不修复。
|
||
|
||
返回:
|
||
统计 dict。
|
||
"""
|
||
stats = {
|
||
"vid": vid,
|
||
"issues_found": 0,
|
||
"l3_repaired": 0,
|
||
"l2_regenerated": 0,
|
||
"l1_regenerated": 0,
|
||
"verify_removed": 0,
|
||
"facts_injected": 0,
|
||
"error": None,
|
||
}
|
||
|
||
try:
|
||
# 加载树
|
||
index = TreeIndex.load_json(str(tree_path))
|
||
|
||
# Step 1: 检测问题
|
||
issues = detect_issues(index, frames_dir=frames_dir)
|
||
stats["issues_found"] = len(issues)
|
||
|
||
if issues:
|
||
logger.info("[{}] 发现 {} 个问题", vid, len(issues))
|
||
for issue in issues[:5]:
|
||
logger.debug(" {} [L{}] {}", issue.node_id, issue.level, issue.details)
|
||
if len(issues) > 5:
|
||
logger.debug(" ... 还有 {} 个", len(issues) - 5)
|
||
|
||
if dry_run:
|
||
return stats
|
||
|
||
# Step 2: VLM 修复(如有 empty_field 问题)
|
||
empty_issues = [i for i in issues if i.issue_type == "empty_field"]
|
||
if empty_issues:
|
||
srt_entries = None
|
||
srt_path = srt_dir / f"{vid}.srt"
|
||
if srt_path.exists():
|
||
srt_entries = parse_srt(str(srt_path))
|
||
|
||
repair_stats = await repair_tree(
|
||
index, empty_issues, vlm, llm, frames_dir, srt_entries
|
||
)
|
||
stats["l3_repaired"] = repair_stats.l3_repaired
|
||
stats["l2_regenerated"] = repair_stats.l2_regenerated
|
||
stats["l1_regenerated"] = repair_stats.l1_regenerated
|
||
logger.info(
|
||
"[{}] 修复完成: L3={}, L2={}, L1={}",
|
||
vid, repair_stats.l3_repaired, repair_stats.l2_regenerated,
|
||
repair_stats.l1_regenerated,
|
||
)
|
||
|
||
# Step 3: 质量校验
|
||
verify_stats = verify_tree(index)
|
||
total_removed = (
|
||
verify_stats.l2_entities_removed
|
||
+ verify_stats.l2_visible_text_removed
|
||
+ verify_stats.l1_visible_text_removed
|
||
+ verify_stats.l1_key_entities_removed
|
||
)
|
||
stats["verify_removed"] = total_removed
|
||
if total_removed > 0:
|
||
logger.info("[{}] 校验删除 {} 项不可靠内容", vid, total_removed)
|
||
|
||
# Step 4: Q&A 反向补全
|
||
questions_path = questions_dir / f"{vid}.json"
|
||
if questions_path.exists():
|
||
with open(questions_path, encoding="utf-8") as f:
|
||
questions = json.load(f)
|
||
if isinstance(questions, list) and questions:
|
||
logger.info("[{}] 开始 Q&A 补全 ({} 道题)...", vid, len(questions))
|
||
srt_text = ""
|
||
srt_path = srt_dir / f"{vid}.srt"
|
||
if srt_path.exists():
|
||
srt_text = srt_path.read_text(encoding="utf-8", errors="ignore")
|
||
|
||
try:
|
||
supplement_stats = await supplement_tree(
|
||
index, questions, llm, srt_text=srt_text
|
||
)
|
||
stats["facts_injected"] = supplement_stats.facts_injected
|
||
if supplement_stats.facts_injected > 0:
|
||
logger.info(
|
||
"[{}] 补全注入 {} 个事实", vid, supplement_stats.facts_injected
|
||
)
|
||
except Exception as exc:
|
||
logger.error("[{}] Q&A 补全失败: {}", vid, exc)
|
||
logger.info("[{}] Q&A 补全完成", vid)
|
||
|
||
# Step 5: 保存
|
||
index.save_json(str(tree_path))
|
||
logger.info("[{}] 已保存", vid)
|
||
|
||
except Exception as exc:
|
||
stats["error"] = str(exc)
|
||
logger.error("[{}] 修复失败: {}", vid, exc)
|
||
|
||
return stats
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 主流程
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
async def main_async(args: argparse.Namespace) -> None:
|
||
"""异步主流程:并发修复视频。"""
|
||
videos_dir = Path(args.videos_dir)
|
||
srt_dir = Path(args.srt_dir)
|
||
questions_dir = Path(args.questions_dir)
|
||
concurrency = args.concurrency
|
||
reaggregate_all = args.reaggregate_all
|
||
|
||
# 扫描所有视频
|
||
vid_dirs = sorted(
|
||
d for d in videos_dir.iterdir()
|
||
if d.is_dir() and (d / "tree.json").exists()
|
||
)
|
||
logger.info("发现 {} 个视频", len(vid_dirs))
|
||
|
||
# 加载 progress
|
||
progress_path = PROJECT_ROOT / "logs" / PROGRESS_FILE
|
||
finished = load_progress(progress_path)
|
||
if finished:
|
||
logger.info("已完成 {} 个视频(从 progress 文件加载)", len(finished))
|
||
|
||
if args.dry_run:
|
||
logger.info("=== DRY RUN 模式:仅检测不修复 ===")
|
||
|
||
# 构建客户端(dry_run 模式不需要)
|
||
llm, vlm = (None, None) if args.dry_run else _build_clients(concurrency)
|
||
|
||
# 过滤跳过的视频
|
||
pending = []
|
||
skipped_count = 0
|
||
for vid_dir in vid_dirs:
|
||
vid = vid_dir.name
|
||
if should_skip_video(vid, finished, reaggregate_all=reaggregate_all):
|
||
skipped_count += 1
|
||
continue
|
||
pending.append(vid_dir)
|
||
|
||
if skipped_count:
|
||
logger.info("跳过 {} 个已完成视频,待处理 {} 个", skipped_count, len(pending))
|
||
|
||
# 并发编排
|
||
sem = asyncio.Semaphore(concurrency)
|
||
progress_lock = asyncio.Lock()
|
||
all_stats: list[dict] = []
|
||
stats_lock = asyncio.Lock()
|
||
start_time = time.time()
|
||
completed = 0
|
||
|
||
async def _process(vid_dir: Path) -> None:
|
||
nonlocal completed
|
||
async with sem:
|
||
vid = vid_dir.name
|
||
tree_path = vid_dir / "tree.json"
|
||
frames_dir = vid_dir
|
||
|
||
logger.info("开始修复 {}", vid)
|
||
|
||
stats = await _repair_one_video(
|
||
vid, tree_path, frames_dir, srt_dir, questions_dir,
|
||
llm, vlm, dry_run=args.dry_run,
|
||
)
|
||
|
||
async with stats_lock:
|
||
all_stats.append(stats)
|
||
completed += 1
|
||
|
||
# 无 error 且非 dry_run 才记 finished
|
||
if stats["error"] is None and not args.dry_run:
|
||
await save_progress(progress_path, progress_lock, vid)
|
||
|
||
# 进度日志
|
||
if completed % 10 == 0:
|
||
elapsed = time.time() - start_time
|
||
rate = completed / elapsed * 60 if elapsed > 0 else 0
|
||
logger.info(
|
||
"进度: {}/{}, 已用 {:.0f}s, 速率 {:.1f} 视频/分钟",
|
||
completed, len(pending), elapsed, rate,
|
||
)
|
||
|
||
tasks = [asyncio.create_task(_process(vd)) for vd in pending]
|
||
await asyncio.gather(*tasks)
|
||
|
||
# 最终汇总
|
||
elapsed = time.time() - start_time
|
||
total_issues = sum(s["issues_found"] for s in all_stats)
|
||
total_repaired = sum(s["l3_repaired"] for s in all_stats)
|
||
total_injected = sum(s["facts_injected"] for s in all_stats)
|
||
total_errors = sum(1 for s in all_stats if s["error"])
|
||
|
||
logger.info("=" * 60)
|
||
logger.info("修复完成")
|
||
logger.info(" 视频总数: {}", len(all_stats))
|
||
logger.info(" 跳过数: {}", skipped_count)
|
||
logger.info(" 问题总数: {}", total_issues)
|
||
logger.info(" L3 修复数: {}", total_repaired)
|
||
logger.info(" 事实注入数: {}", total_injected)
|
||
logger.info(" 失败数: {}", total_errors)
|
||
logger.info(" 总耗时: {:.0f}s", elapsed)
|
||
logger.info(" 并发数: {}", concurrency)
|
||
logger.info("=" * 60)
|
||
|
||
if total_errors > 0:
|
||
logger.warning("以下视频修复失败:")
|
||
for s in all_stats:
|
||
if s["error"]:
|
||
logger.warning(" {}: {}", s["vid"], s["error"])
|
||
|
||
|
||
def parse_args() -> argparse.Namespace:
|
||
"""解析命令行参数。"""
|
||
parser = argparse.ArgumentParser(description="树修复管线")
|
||
parser.add_argument(
|
||
"--videos-dir",
|
||
default="store/videos",
|
||
help="视频目录(默认: store/videos)",
|
||
)
|
||
parser.add_argument(
|
||
"--srt-dir",
|
||
default="data/Video-MME/subtitle",
|
||
help="SRT 字幕目录(默认: data/Video-MME/subtitle)",
|
||
)
|
||
parser.add_argument(
|
||
"--questions-dir",
|
||
default="store/questions/benchmarks/Video-MME",
|
||
help="问题 JSON 目录(默认: store/questions/benchmarks/Video-MME)",
|
||
)
|
||
parser.add_argument(
|
||
"--dry-run",
|
||
action="store_true",
|
||
help="仅检测不修复,不调用 VLM/LLM",
|
||
)
|
||
parser.add_argument(
|
||
"--concurrency",
|
||
type=int,
|
||
default=16,
|
||
help="并发修复视频数(默认: 16)",
|
||
)
|
||
parser.add_argument(
|
||
"--reaggregate-all",
|
||
action="store_true",
|
||
help="强制全量重聚合,忽略 progress 文件",
|
||
)
|
||
return parser.parse_args()
|
||
|
||
|
||
def main() -> None:
|
||
"""同步入口。"""
|
||
args = parse_args()
|
||
(PROJECT_ROOT / "logs").mkdir(exist_ok=True)
|
||
asyncio.run(main_async(args))
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|