#!/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 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()