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:
@@ -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
|
||||
@@ -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():
|
||||
"""构建 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_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()
|
||||
Reference in New Issue
Block a user