Files
iomgaa d6bcf41336 fix(repair): Codex 审查修正 — progress 完成判据加严 + load_progress 防御增强
1. 保存 progress 前重新 detect_issues,只有 empty_field 清零才记 finished
   (修复 L2/L1 空字段被检测但未修复仍写 finished 的问题)
2. load_progress 增加 AttributeError 捕获(防 JSON 非 dict 形态崩溃)
2026-07-09 00:30:23 -04:00

493 lines
16 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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, AttributeError):
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):
"""构建 GovernedLLMClientLLM + 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
# 修复后重新检测,关键 issue 清零才记 finished
if stats["error"] is None and not args.dry_run:
index = TreeIndex.load_json(str(tree_path))
remaining = [i for i in detect_issues(index) if i.issue_type == "empty_field"]
if not remaining:
await save_progress(progress_path, progress_lock, vid)
else:
logger.warning(
"[{}] 修复后仍有 {} 个 empty_field,不计入 finished",
vid,
len(remaining),
)
# 进度日志
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()