feat(repair): 断点续跑 progress 文件管理

load_progress / save_progress(asyncio.Lock + os.replace 原子写入)
/ should_skip_video。支持并发安全的读改写和 --reaggregate-all 兜底。

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-09 00:23:04 -04:00
parent 8182cb86b1
commit afe80a8b32
2 changed files with 501 additions and 0 deletions
+71
View File
@@ -0,0 +1,71 @@
"""修复管线断点续跑 progress 管理测试。"""
from __future__ import annotations
import asyncio
import json
import pytest
def test_load_progress_missing_file(tmp_path):
"""progress 文件不存在时返回空集合。"""
from tools.repair_trees import load_progress
result = load_progress(tmp_path / "nonexistent.json")
assert result == set()
def test_load_progress_valid_file(tmp_path):
"""正常读取已有 progress 文件。"""
from tools.repair_trees import load_progress
path = tmp_path / "progress.json"
path.write_text(json.dumps({"finished_video_ids": ["vid_a", "vid_b"]}))
result = load_progress(path)
assert result == {"vid_a", "vid_b"}
def test_load_progress_corrupted_file(tmp_path):
"""损坏的 JSON 文件返回空集合(不抛异常)。"""
from tools.repair_trees import load_progress
path = tmp_path / "progress.json"
path.write_text("{invalid json")
result = load_progress(path)
assert result == set()
@pytest.mark.asyncio
async def test_save_progress_atomic(tmp_path):
"""save_progress 原子写入,并发调用不丢失更新。"""
from tools.repair_trees import save_progress
path = tmp_path / "progress.json"
lock = asyncio.Lock()
await save_progress(path, lock, "vid_a")
await save_progress(path, lock, "vid_b")
data = json.loads(path.read_text())
assert set(data["finished_video_ids"]) == {"vid_a", "vid_b"}
@pytest.mark.asyncio
async def test_save_progress_concurrent(tmp_path):
"""16 路并发 save_progress 不丢失更新。"""
from tools.repair_trees import save_progress
path = tmp_path / "progress.json"
lock = asyncio.Lock()
tasks = [save_progress(path, lock, f"vid_{i}") for i in range(16)]
await asyncio.gather(*tasks)
data = json.loads(path.read_text())
assert len(data["finished_video_ids"]) == 16
def test_should_skip_finished():
"""已在 finished 集合中的视频应跳过。"""
from tools.repair_trees import should_skip_video
finished = {"vid_a", "vid_b"}
assert should_skip_video("vid_a", finished, reaggregate_all=False) is True
assert should_skip_video("vid_c", finished, reaggregate_all=False) is False
def test_should_skip_reaggregate_all_forces_rerun():
"""--reaggregate-all 标志强制不跳过。"""
from tools.repair_trees import should_skip_video
finished = {"vid_a"}
assert should_skip_video("vid_a", finished, reaggregate_all=True) is False
+430
View File
@@ -0,0 +1,430 @@
#!/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():
"""构建 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_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)
# 扫描所有视频
vid_dirs = sorted(
d for d in videos_dir.iterdir()
if d.is_dir() and (d / "tree.json").exists()
)
logger.info("发现 {} 个视频待修复", len(vid_dirs))
if args.dry_run:
logger.info("=== DRY RUN 模式:仅检测不修复 ===")
# 构建客户端(dry_run 模式不需要)
llm, vlm = (None, None) if args.dry_run else _build_clients()
# 逐视频修复
all_stats = []
start_time = time.time()
for idx, vid_dir in enumerate(vid_dirs):
vid = vid_dir.name
tree_path = vid_dir / "tree.json"
frames_dir = vid_dir # frame_path 已含 "frames/" 前缀,不再嵌套
logger.info(
"[{}/{}] 开始修复 {}",
idx + 1, len(vid_dirs), vid,
)
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)
# 每 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,
)
# 最终汇总
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(" 问题总数: {}", total_issues)
logger.info(" L3 修复数: {}", total_repaired)
logger.info(" 事实注入数: {}", total_injected)
logger.info(" 失败数: {}", total_errors)
logger.info(" 总耗时: {:.0f}s", elapsed)
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",
)
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()