diff --git a/.env.example b/.env.example index 2b45c02..a0c1ce0 100644 --- a/.env.example +++ b/.env.example @@ -41,7 +41,7 @@ REDIS_URL=redis://localhost:6379/0 LLM_TIMEOUT=120 LLM_MAX_RETRIES=3 LLM_RETRY_BASE_DELAY=2.0 -LLM_CIRCUIT_BREAKER_THRESHOLD=5 +LLM_CIRCUIT_BREAKER_THRESHOLD=5 # 实际阈值 = max(此值, concurrency*2) LLM_CIRCUIT_BREAKER_COOLDOWN=60 LLM_TTFT_TIMEOUT=30 LLM_INTER_TOKEN_TIMEOUT=15 diff --git a/tools/repair_trees.py b/tools/repair_trees.py index 13d287b..e87df42 100644 --- a/tools/repair_trees.py +++ b/tools/repair_trees.py @@ -125,7 +125,7 @@ def should_skip_video(vid: str, finished: set[str], *, reaggregate_all: bool) -> # --------------------------------------------------------------------------- -def _build_clients(): +def _build_clients(concurrency: int = 16): """构建 GovernedLLMClient(LLM + VLM)。 返回: @@ -141,6 +141,7 @@ def _build_clients(): 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")) @@ -323,52 +324,87 @@ async def _repair_one_video( 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)) + 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() + llm, vlm = (None, None) if args.dry_run else _build_clients(concurrency) - # 逐视频修复 - all_stats = [] - start_time = time.time() - for idx, vid_dir in enumerate(vid_dirs): + # 过滤跳过的视频 + pending = [] + skipped_count = 0 + for vid_dir in vid_dirs: vid = vid_dir.name - tree_path = vid_dir / "tree.json" - frames_dir = vid_dir # frame_path 已含 "frames/" 前缀,不再嵌套 + if should_skip_video(vid, finished, reaggregate_all=reaggregate_all): + skipped_count += 1 + continue + pending.append(vid_dir) - logger.info( - "[{}/{}] 开始修复 {}", - idx + 1, len(vid_dirs), vid, - ) + if skipped_count: + logger.info("跳过 {} 个已完成视频,待处理 {} 个", skipped_count, len(pending)) - stats = await _repair_one_video( - vid, tree_path, frames_dir, srt_dir, questions_dir, - llm, vlm, dry_run=args.dry_run, - ) - all_stats.append(stats) + # 并发编排 + sem = asyncio.Semaphore(concurrency) + progress_lock = asyncio.Lock() + all_stats: list[dict] = [] + stats_lock = asyncio.Lock() + start_time = time.time() + completed = 0 - # 每 10 个视频汇总一次 - if (idx + 1) % 10 == 0: - elapsed = time.time() - start_time - rate = (idx + 1) / elapsed * 60 - logger.info( - "进度: {}/{}, 已用 {:.0f}s, 速率 {:.1f} 视频/分钟", - idx + 1, len(vid_dirs), elapsed, rate, + 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) @@ -379,11 +415,13 @@ async def main_async(args: argparse.Namespace) -> None: 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: @@ -416,6 +454,17 @@ def parse_args() -> argparse.Namespace: 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()