Files
Video-Tree-TRM5/tools/generate_questions.py
T

1286 lines
42 KiB
Python
Raw 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
"""赛题生成工具:generate + calibrate。
用法:
conda activate Video-Tree-TRM
python tools/generate_questions.py generate --store-dir store ...
python tools/generate_questions.py calibrate \
--baseline-db workspaces/default/harness.db --baseline-run-id infer_adhoc \
--target-db workspaces/default/harness.db --target-run-id infer_gen240
calibrate 是纯对比工具,不跑推理。推理统一走 main.py --mode infer。
app/core/adapters 不 import 此脚本。
"""
from __future__ import annotations
import argparse
import asyncio
import json
import os
import random
import sys
from collections import defaultdict
from pathlib import 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")
import numpy as np
from scipy.stats import fisher_exact
from app.question_gen.loader import load_benchmark
from app.question_gen.synthesizer import (
TASK_TYPE_LEVEL_MAP,
generate_one,
is_duplicate,
)
from core.types import GeneratedQuestion # noqa: TCH001 — runtime use in _append_to_json
# ---------------------------------------------------------------------------
# 日志配置:不缓存,立即输出
# ---------------------------------------------------------------------------
logger.remove()
logger.add(
sys.stderr,
format="{time:HH:mm:ss} | {level:<7} | {message}",
level="DEBUG",
colorize=True,
)
logger.add(
PROJECT_ROOT / "logs" / "generate_questions.log",
format="{time:YYYY-MM-DD HH:mm:ss} | {level:<7} | {message}",
level="DEBUG",
rotation="50 MB",
enqueue=False,
)
# ---------------------------------------------------------------------------
# 断点续跑 — progress 文件管理
# ---------------------------------------------------------------------------
def _load_or_init_progress(output_dir: Path) -> dict:
"""加载 progress.json,不存在则返回初始结构。
参数:
output_dir: 输出目录路径。
返回:
{"completed": {task_type: [question_id, ...]}, "output_dir": str}。
文件损坏时返回初始结构并记录警告。
"""
progress_path = output_dir / "progress.json"
if progress_path.exists():
try:
data = json.loads(progress_path.read_text(encoding="utf-8"))
if not isinstance(data.get("completed"), dict):
raise ValueError("completed 字段不是 dict")
return data
except (json.JSONDecodeError, ValueError, KeyError, TypeError) as exc:
logger.warning("progress.json 损坏,重新初始化: {}", exc)
return {"completed": {}, "output_dir": str(output_dir)}
def _save_progress(output_dir: Path, progress: dict) -> None:
"""原子写入 progress.json。
参数:
output_dir: 输出目录路径。
progress: 进度数据。
"""
tmp = output_dir / "progress.json.tmp"
tmp.write_text(
json.dumps(progress, ensure_ascii=False, indent=2),
encoding="utf-8",
)
os.replace(str(tmp), str(output_dir / "progress.json"))
# ---------------------------------------------------------------------------
# Embedding 池重建(断点续跑时从已生成 JSON 重建)
# ---------------------------------------------------------------------------
def _rebuild_embedding_pool(
output_dir: Path,
embed_fn,
benchmark_questions: list[GeneratedQuestion],
) -> dict[str, np.ndarray]:
"""从已生成 JSON + benchmark 题目重建每个题型的 embedding 池。
断点续跑时调用,确保去重池包含所有已有题目。
参数:
output_dir: 包含 {video_id}.json 的输出目录。
embed_fn: 文本嵌入函数(str | list[str] → [N, D] ndarray)。
benchmark_questions: benchmark 题目列表(也要加入去重池)。
返回:
{task_type: [N, D] ndarray},空题型的 ndarray 为 shape (0,)。
"""
pools: dict[str, list[str]] = {}
# Phase 1: 收集 benchmark 题目文本
for q in benchmark_questions:
pools.setdefault(q.task_type, []).append(q.question)
# Phase 2: 收集已生成题目文本
for json_path in sorted(output_dir.glob("*.json")):
if json_path.name == "progress.json":
continue
try:
items = json.loads(json_path.read_text(encoding="utf-8"))
except (json.JSONDecodeError, OSError) as exc:
logger.warning("跳过损坏文件 {}: {}", json_path, exc)
continue
if not isinstance(items, list):
continue
for item in items:
task_type = item.get("task_type", "")
question = item.get("question", "")
if task_type and question:
pools.setdefault(task_type, []).append(question)
# Phase 3: 批量嵌入
result: dict[str, np.ndarray] = {}
for task_type, texts in pools.items():
if texts:
result[task_type] = embed_fn(texts)
else:
result[task_type] = np.empty(0)
# Phase 4: 确保所有 12 题型都有条目
for task_type in TASK_TYPE_LEVEL_MAP:
if task_type not in result:
result[task_type] = np.empty(0)
logger.info(
"embedding 池重建完成: {}",
{k: v.shape[0] if v.ndim == 2 else 0 for k, v in result.items()},
)
return result
# ---------------------------------------------------------------------------
# Exemplar 选取
# ---------------------------------------------------------------------------
def _select_exemplars(
benchmark: list[GeneratedQuestion],
task_type: str,
n: int,
rng: random.Random,
) -> list[GeneratedQuestion]:
"""从 benchmark 中选取同题型示例,优先跨视频多样性。
参数:
benchmark: benchmark 题目全集。
task_type: 目标题型。
n: 期望选取数量。
rng: 可控随机数生成器。
返回:
min(n, 可用数) 个示例,尽量来自不同 video_id。
"""
# Phase 1: 过滤同题型
candidates = [q for q in benchmark if q.task_type == task_type]
if not candidates:
return []
take = min(n, len(candidates))
# Phase 2: 按 video_id 分桶,轮询取样保证跨视频多样性
by_video: dict[str, list[GeneratedQuestion]] = {}
for q in candidates:
by_video.setdefault(q.video_id, []).append(q)
# 每桶内部打乱
for bucket in by_video.values():
rng.shuffle(bucket)
# 轮询选取
video_ids = list(by_video.keys())
rng.shuffle(video_ids)
selected: list[GeneratedQuestion] = []
idx = 0
while len(selected) < take:
vid = video_ids[idx % len(video_ids)]
bucket = by_video[vid]
if bucket:
selected.append(bucket.pop(0))
else:
# 桶空了,从 video_ids 中移除
video_ids.remove(vid)
if not video_ids:
break
# 不递增 idx,因为移除后当前位置是下一个
continue
idx += 1
return selected
# ---------------------------------------------------------------------------
# 客户端构建(从 .env
# ---------------------------------------------------------------------------
def _build_vlm_client():
"""构建 GovernedVLMClient,复用 repair_trees.py 的模式。
从 .env 读取 VL_LLM_MODEL / VL_LLM_BASE_URL / VL_LLM_API_KEY
和 LLM 韧性参数,构造治理栈。
返回:
GovernedVLMClient 实例。
"""
from adapters.breaker import CircuitBreaker
from adapters.llm import GovernedLLMClient
from adapters.telemetry import SQLiteTelemetryRecorder
from adapters.vlm import GovernedVLMClient
(PROJECT_ROOT / "logs").mkdir(exist_ok=True)
telemetry = SQLiteTelemetryRecorder(
str(PROJECT_ROOT / "logs" / "generate_questions_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"))
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,
)
return GovernedVLMClient(vlm_base)
def _build_embed_provider():
"""构建 EmbeddingProvider,从 .env 决定 local 或 remote。
环境变量:
EMBED_API_KEY + EMBED_API_URL 都非空 → RemoteEmbeddingProvider
否则 → LocalEmbeddingProvider
模型名称和维度通过 EMBED_MODEL / EMBED_DIM 环境变量配置。
返回:
LocalEmbeddingProvider 或 RemoteEmbeddingProvider 实例。
"""
from adapters.embedding import LocalEmbeddingProvider, RemoteEmbeddingProvider
model_name = os.environ.get("EMBED_MODEL", "BAAI/bge-base-zh-v1.5")
embed_dim = int(os.environ.get("EMBED_DIM", "768"))
api_key = os.environ.get("EMBED_API_KEY", "")
api_url = os.environ.get("EMBED_API_URL", "")
if api_key and api_url:
logger.info("使用远程嵌入: model={}, url={}", model_name, api_url)
return RemoteEmbeddingProvider(
model_name=model_name,
embed_dim=embed_dim,
api_key=api_key,
api_url=api_url,
)
logger.info("使用本地嵌入: model={}, dim={}", model_name, embed_dim)
device = os.environ.get("EMBED_DEVICE", "cpu")
return LocalEmbeddingProvider(
model_name=model_name,
embed_dim=embed_dim,
device=device,
)
# ---------------------------------------------------------------------------
# JSON 追加写入
# ---------------------------------------------------------------------------
def _question_to_entry(question: GeneratedQuestion) -> dict:
"""将题目序列化为 JSON entry_append_to_json 与 _on_accept 共用)。"""
return {
"question_id": question.question_id,
"video_id": question.video_id,
"task_type": question.task_type,
"question": question.question,
"options": list(question.options),
"answer": question.answer,
"source_nodes": list(question.source_nodes),
"difficulty": question.difficulty,
"family": question.family,
"skill_target": question.skill_target,
"sub_pattern": question.sub_pattern,
}
def _append_to_json(output_dir: Path, question: GeneratedQuestion) -> None:
"""将生成的题目追加到对应 video_id 的 JSON 文件。
文件格式:[{...}, {...}, ...],每个 video_id 一个文件。
参数:
output_dir: 输出目录。
question: 待写入的题目。
"""
json_path = output_dir / f"{question.video_id}.json"
existing: list[dict] = []
if json_path.exists():
try:
existing = json.loads(json_path.read_text(encoding="utf-8"))
except (json.JSONDecodeError, OSError):
logger.warning("读取 {} 失败,覆盖写入", json_path)
existing = []
entry = _question_to_entry(question)
existing.append(entry)
# 原子写入
tmp = json_path.with_suffix(".json.tmp")
tmp.write_text(
json.dumps(existing, ensure_ascii=False, indent=2),
encoding="utf-8",
)
os.replace(str(tmp), str(json_path))
# ---------------------------------------------------------------------------
# calibrate 辅助函数
# ---------------------------------------------------------------------------
def _judge_task_type(
bench_correct: int,
bench_total: int,
gen_correct: int,
gen_total: int,
tolerance: float,
alpha: float,
) -> str:
"""判定单个题型的校准结果。
根据 benchmark 和生成题正确率差值 + Fisher 精确检验决定判定。
参数:
bench_correct: benchmark 答对数。
bench_total: benchmark 总题数。
gen_correct: 生成题答对数。
gen_total: 生成题总题数。
tolerance: 正确率差值容忍阈值。
alpha: Fisher 检验显著性水平。
返回:
"PASS" — 差值在容忍范围内。
"FAIL" — 差值超阈值且统计显著。
"WARN" — 差值超阈值但不显著。
"""
delta = abs(gen_correct / gen_total - bench_correct / bench_total)
if delta <= tolerance:
return "PASS"
table = [
[bench_correct, bench_total - bench_correct],
[gen_correct, gen_total - gen_correct],
]
_, p = fisher_exact(table)
if p < alpha and delta > tolerance:
return "FAIL"
return "WARN"
def _calibrate_exit_code(verdicts: dict[str, str]) -> int:
"""根据所有题型的判定结果决定进程退出码。
参数:
verdicts: {题型: "PASS"|"WARN"|"FAIL"} 映射。
返回:
存在任一 FAIL → 1,否则 → 0。
"""
if any(v == "FAIL" for v in verdicts.values()):
return 1
return 0
def _read_baseline_per_task_type(
db_path: str,
run_id: str,
) -> dict[str, dict]:
"""从已有 HarnessLog DB 中读取指定 run 的 per_task_type 正确率。
参数:
db_path: SQLite 数据库路径。
run_id: 运行标识。
返回:
{task_type: {"accuracy": float, "total": int, "correct": int}}。
异常:
FileNotFoundError: 数据库文件不存在。
ValueError: 未找到指定 run_id 的预测记录。
"""
import sqlite3
if not Path(db_path).exists():
raise FileNotFoundError(f"基线数据库不存在: {db_path}")
conn = sqlite3.connect(db_path)
conn.row_factory = sqlite3.Row
try:
rows = conn.execute(
"SELECT task_type, prediction, answer FROM predictions WHERE run_id = ?",
(run_id,),
).fetchall()
finally:
conn.close()
if not rows:
raise ValueError(f"未找到 run_id={run_id} 的预测记录")
groups: dict[str, list[dict]] = defaultdict(list)
for row in rows:
groups[dict(row)["task_type"]].append(dict(row))
result: dict[str, dict] = {}
for task_type, records in groups.items():
total = len(records)
correct = sum(1 for r in records if r["prediction"] == r["answer"])
result[task_type] = {
"accuracy": correct / total,
"total": total,
"correct": correct,
}
return result
def _format_comparison_table(
bench_per_task: dict[str, dict],
gen_per_task: dict[str, dict],
verdicts: dict[str, str],
p_values: dict[str, float],
) -> str:
"""格式化校准比较表。
参数:
bench_per_task: benchmark 各题型指标。
gen_per_task: 生成题各题型指标。
verdicts: 各题型判定结果。
p_values: 各题型 Fisher 检验 p 值。
返回:
格式化的比较表字符串。
"""
verdict_symbols = {"PASS": "✓ PASS", "WARN": "⚠ WARN", "FAIL": "✗ FAIL"}
all_types = sorted(set(bench_per_task) | set(gen_per_task))
header = f"{'题型':<20s} | {'bench':>6s} | {'gen':>6s} | {'Δ':>7s} | {'p-value':>7s} | 判定"
sep = "-" * 19 + "-|" + "-" * 8 + "|" + "-" * 8 + "|" + "-" * 9 + "|" + "-" * 9 + "|" + "-" * 8
lines = [header, sep]
for task_type in all_types:
b = bench_per_task.get(task_type, {"accuracy": 0.0, "total": 0, "correct": 0})
g = gen_per_task.get(task_type, {"accuracy": 0.0, "total": 0, "correct": 0})
delta = g["accuracy"] - b["accuracy"]
p_val = p_values.get(task_type, float("nan"))
verdict = verdicts.get(task_type, "N/A")
symbol = verdict_symbols.get(verdict, verdict)
lines.append(
f"{task_type:<20s} | {b['accuracy']:>5.1%} | {g['accuracy']:>5.1%} "
f"| {delta:>+6.1%} | {p_val:>7.3f} | {symbol}"
)
return "\n".join(lines)
def _run_calibrate(args: argparse.Namespace) -> None:
"""calibrate 子命令主流程(纯对比,不跑推理)。
从两组已有的推理结果(harness.db + run_id)读取 per_task_type 正确率,
逐题型 Fisher 精确检验判定校准质量。
推理应事先通过 main.py --mode infer 完成,确保 baseline 和 target
使用完全相同的推理管线。
参数:
args: CLI 参数(baseline_db, baseline_run_id, target_db, target_run_id,
tolerance, alpha)。
"""
tolerance = args.tolerance
alpha = args.alpha
# Phase 1: 从 DB 读取两组推理结果
logger.info(
"读取 baseline: db={}, run_id={}",
args.baseline_db,
args.baseline_run_id,
)
baseline_per_task = _read_baseline_per_task_type(
args.baseline_db,
args.baseline_run_id,
)
logger.info(
"读取 target: db={}, run_id={}",
args.target_db,
args.target_run_id,
)
target_per_task = _read_baseline_per_task_type(
args.target_db,
args.target_run_id,
)
baseline_total = sum(v["total"] for v in baseline_per_task.values())
target_total = sum(v["total"] for v in target_per_task.values())
logger.info("baseline {} 道, target {} 道", baseline_total, target_total)
# Phase 2: 逐题型判定
all_types = sorted(set(baseline_per_task) | set(target_per_task))
verdicts: dict[str, str] = {}
p_values: dict[str, float] = {}
for task_type in all_types:
b = baseline_per_task.get(task_type)
g = target_per_task.get(task_type)
if b is None or g is None or b["total"] == 0 or g["total"] == 0:
verdicts[task_type] = "WARN"
p_values[task_type] = float("nan")
continue
verdicts[task_type] = _judge_task_type(
bench_correct=b["correct"],
bench_total=b["total"],
gen_correct=g["correct"],
gen_total=g["total"],
tolerance=tolerance,
alpha=alpha,
)
table = [
[b["correct"], b["total"] - b["correct"]],
[g["correct"], g["total"] - g["correct"]],
]
_, p_val = fisher_exact(table)
p_values[task_type] = p_val
# Phase 3: 输出比较表
table_str = _format_comparison_table(
baseline_per_task,
target_per_task,
verdicts,
p_values,
)
logger.info("校准比较表:\n{}", table_str)
# Phase 4: 退出
exit_code = _calibrate_exit_code(verdicts)
if exit_code == 0:
logger.info("校准通过: 所有题型 PASS 或 WARN")
else:
logger.error("校准失败: 存在 FAIL 题型")
sys.exit(exit_code)
# ---------------------------------------------------------------------------
# generate 主流程
# ---------------------------------------------------------------------------
async def _run_generate(args: argparse.Namespace) -> None:
"""generate 子命令主流程。
按题型顺序生成题目,每个 slot 串行生成并去重,
断点续跑通过 progress.json 跳过已完成 slot。
参数:
args: CLI 参数(store_dir, output_dir, per_type, similarity_threshold,
max_retries, concurrency, seed)。
"""
store_dir = Path(args.store_dir)
output_dir = Path(args.output_dir)
per_type = args.per_type
similarity_threshold = args.similarity_threshold
# Information Synopsis 题目天然高相似(均值 0.66,中位数 0.67),
# 用通用阈值会导致几乎所有新题被误判重复。按题型放宽。
_PER_TYPE_THRESHOLD = {
"Information Synopsis": max(similarity_threshold, 0.90),
}
max_retries = args.max_retries
concurrency = args.concurrency
seed = args.seed
output_dir.mkdir(parents=True, exist_ok=True)
# Phase 1: 加载视频列表
videos_dir = store_dir / "videos"
if not videos_dir.exists():
logger.error("视频目录不存在: {}", videos_dir)
sys.exit(1)
video_ids = sorted(
d.name for d in videos_dir.iterdir() if d.is_dir() and (d / "tree.json").exists()
)
if not video_ids:
logger.error("未找到任何有 tree.json 的视频目录")
sys.exit(1)
logger.info("发现 {} 个视频", len(video_ids))
# Phase 2: 加载 benchmark 题目(用于 exemplars + 去重池初始化)
benchmark_dir = store_dir / "questions" / "benchmarks" / "Video-MME"
benchmark: list[GeneratedQuestion] = []
if benchmark_dir.exists():
benchmark = load_benchmark(benchmark_dir)
logger.info("加载 {} 道 benchmark 题目", len(benchmark))
else:
logger.warning("benchmark 目录不存在: {}", benchmark_dir)
# Phase 3: 构建客户端
vlm = _build_vlm_client()
embed_provider = _build_embed_provider()
embed_fn = embed_provider.embed
# Phase 4: 初始化或恢复 progress
progress = _load_or_init_progress(output_dir)
# Phase 5: 重建 embedding 池
pools = _rebuild_embedding_pool(output_dir, embed_fn, benchmark)
# Phase 6: 加载视频树索引(延迟按需加载)
from app.tree.index import TreeIndex
rng = random.Random(seed)
sem = asyncio.Semaphore(concurrency)
task_types = list(TASK_TYPE_LEVEL_MAP.keys())
total_generated = 0
total_failed = 0
async def _generate_with_sem(
vlm_client,
tree,
video_id,
task_type,
seq,
*,
exemplars,
used_node_ids,
max_retries_inner,
rng_inner,
session_id,
):
"""Semaphore 包装的 generate_one 调用。"""
async with sem:
return await generate_one(
vlm_client,
tree,
video_id,
task_type,
seq,
exemplars=exemplars,
used_node_ids=used_node_ids,
max_retries=max_retries_inner,
rng=rng_inner,
session_id=session_id,
)
# Phase 7: 逐题型、逐 slot 生成
for task_type in task_types:
completed_ids = set(progress["completed"].get(task_type, []))
start_seq = len(completed_ids)
if start_seq >= per_type:
logger.info("题型 {} 已完成 {}/{}", task_type, start_seq, per_type)
continue
logger.info(
"题型 {} 开始生成: 已完成 {}, 目标 {}",
task_type,
start_seq,
per_type,
)
# 选取 exemplars
exemplars = _select_exemplars(benchmark, task_type, 3, rng)
for seq in range(start_seq, per_type):
used_node_ids: set[str] = set()
session_id = f"gen-{task_type}-{seq}"
generated = False
for _attempt in range(max_retries):
# 每次重试换一个视频,避免同视频反复去重失败
video_id = rng.choice(video_ids)
tree_path = videos_dir / video_id / "tree.json"
try:
tree = TreeIndex.load_json(str(tree_path))
except Exception as exc:
logger.warning("加载树 {} 失败: {}", tree_path, exc)
continue
video_dir = videos_dir / video_id
for l1 in tree.roots:
for l2 in l1.children:
for l3 in l2.children:
if l3.frame_path and not Path(l3.frame_path).is_absolute():
l3.frame_path = str(video_dir / l3.frame_path)
candidate = await _generate_with_sem(
vlm,
tree,
video_id,
task_type,
seq,
exemplars=exemplars,
used_node_ids=used_node_ids,
max_retries_inner=1,
rng_inner=rng,
session_id=session_id,
)
if candidate is None:
continue
# 去重检查(单线程原子操作)
pool = pools.get(task_type, np.empty(0))
if (
pool.ndim == 2
and pool.shape[0] > 0
and is_duplicate(
candidate.question,
pool,
embed_fn,
_PER_TYPE_THRESHOLD.get(task_type, similarity_threshold),
)
):
logger.warning("去重: {} 与池中题目相似", candidate.question_id)
continue
# 原子操作:更新池 + 写 JSON + 更新 progress
new_emb = embed_fn(candidate.question) # [1, D]
if pool.ndim == 2 and pool.shape[0] > 0:
pools[task_type] = np.vstack([pool, new_emb])
else:
pools[task_type] = new_emb
_append_to_json(output_dir, candidate)
progress["completed"].setdefault(task_type, []).append(candidate.question_id)
_save_progress(output_dir, progress)
total_generated += 1
generated = True
logger.debug(
"生成: {} (题型={}, 序号={})",
candidate.question_id,
task_type,
seq,
)
break
if not generated:
logger.error("题型 {} seq {} 耗尽 {} 次重试", task_type, seq, max_retries)
total_failed += 1
# Phase 8: 汇总
logger.info("=" * 60)
logger.info("生成完成: 成功 {}, 失败 {}", total_generated, total_failed)
logger.info("=" * 60)
if total_failed > 0:
logger.error("{} 个 slot 生成失败", total_failed)
sys.exit(1)
# 全部完成,删除 progress.json
progress_path = output_dir / "progress.json"
if progress_path.exists():
progress_path.unlink()
logger.info("已删除 progress.json(全部完成)")
# ---------------------------------------------------------------------------
# CLI 解析
# ---------------------------------------------------------------------------
def _add_generate_v2_parser(subparsers: argparse._SubParsersAction) -> None:
"""注册 generate-v2 子命令(v2 出题管线 CLI 入口)。
参数:
subparsers: argparse 子命令注册器。
"""
p = subparsers.add_parser("generate-v2", help="v2 出题管线(家族特化 + 四门质量检查)")
p.add_argument(
"--config",
type=Path,
default=Path("config/default.yaml"),
help="管线配置 YAML 文件路径(默认 config/default.yaml",
)
p.add_argument(
"--store-dir",
type=Path,
required=True,
help="store 根目录(包含 videos/ 子目录)",
)
p.add_argument(
"--db-path",
type=Path,
default=Path("logs/question_gen.db"),
help="QuestionGenStore SQLite 数据库路径(默认 logs/question_gen.db",
)
p.add_argument(
"--seed",
type=int,
default=None,
help="随机种子(覆盖配置文件中的 seed)",
)
p.add_argument(
"--task-types",
type=str,
nargs="+",
default=None,
help="只生成指定的 task_type(默认全部 12 类)",
)
p.add_argument(
"--dry-run",
action="store_true",
help="仅加载配置并计算 slot 分配,不调用 LLM/VLM",
)
async def _run_generate_v2(args: argparse.Namespace) -> None:
"""generate-v2 子命令主流程。
流程:
1. 加载 PipelineConfig
2. 如有 --seed,覆盖配置 seed
3. 发现视频列表
4. dry-run 模式下输出统计后返回
5. 构建 VLM/LLM/embedding 客户端(DI
6. 加载 TreeIndex
7. 初始化 QuestionGenStore
8. 加载断点续跑进度
9. 调用 run_pipeline_v2
10. 保存输出
参数:
args: CLI 参数(config, store_dir, db_path, seed, dry_run)。
"""
from app.question_gen.pipeline_v2 import (
PipelineConfig,
load_pipeline_config,
run_pipeline_v2,
)
# Phase 1: 加载配置
config_path = args.config.resolve()
config = load_pipeline_config(config_path)
logger.info("配置加载完成: {}", config_path)
# Phase 2: 覆盖 seed
if args.seed is not None:
config = PipelineConfig(
per_type=config.per_type,
retry_limit=config.retry_limit,
heavy_sample_rate=config.heavy_sample_rate,
dedup_threshold=config.dedup_threshold,
concurrency=config.concurrency,
seed=args.seed,
output_dir=config.output_dir,
)
logger.info("seed 覆盖为: {}", args.seed)
# Phase 3: 发现视频列表
store_dir = args.store_dir.resolve()
videos_dir = store_dir / "videos"
if not videos_dir.exists():
logger.error("视频目录不存在: {}", videos_dir)
sys.exit(1)
video_ids = sorted(
d.name for d in videos_dir.iterdir() if d.is_dir() and (d / "tree.json").exists()
)
if not video_ids:
logger.error("未找到任何有 tree.json 的视频目录")
sys.exit(1)
logger.info("发现 {} 个视频: {}", len(video_ids), video_ids[:5])
# Phase 4: dry-run 模式
all_task_types = [
"Action Recognition",
"Action Reasoning",
"Attribute Perception",
"Counting Problem",
"Information Synopsis",
"Object Recognition",
"Object Reasoning",
"OCR Problems",
"Spatial Perception",
"Spatial Reasoning",
"Temporal Perception",
"Temporal Reasoning",
]
task_types = args.task_types if args.task_types else all_task_types
total_slots = len(task_types) * config.per_type
if args.dry_run:
logger.info("[dry-run] 管线配置摘要:")
logger.info("[dry-run] 视频数: {}", len(video_ids))
logger.info("[dry-run] 任务类型: {} 种", len(task_types))
logger.info("[dry-run] 每类目标: {} 题", config.per_type)
logger.info("[dry-run] 总 slot 数: {}", total_slots)
logger.info("[dry-run] 并发: {}", config.concurrency)
logger.info("[dry-run] 去重阈值: {}", config.dedup_threshold)
logger.info("[dry-run] 退出(不调用 LLM/VLM")
return
# Phase 5: 构建客户端(DI
from adapters.breaker import CircuitBreaker
from adapters.llm import GovernedLLMClient
from adapters.telemetry import SQLiteTelemetryRecorder
from adapters.vlm import GovernedVLMClient
(PROJECT_ROOT / "logs").mkdir(exist_ok=True)
telemetry = SQLiteTelemetryRecorder(str(PROJECT_ROOT / "logs" / "generate_v2_telemetry.db"))
breaker_threshold = int(os.getenv("LLM_CIRCUIT_BREAKER_THRESHOLD", "5"))
# 并发 slot 共享客户端,熔断阈值需按并发度缩放(与 build_trees 一致)
breaker_threshold = max(breaker_threshold, config.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"))
def _make_breaker() -> CircuitBreaker:
return CircuitBreaker(fail_threshold=breaker_threshold, cooldown_s=breaker_cooldown)
# VLM 客户端
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=_make_breaker(),
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)
# LLM 客户端(门控用,复用 JUDGE_LLM 配置)
llm = GovernedLLMClient(
model=os.environ.get("JUDGE_LLM_MODEL", "gpt-4.1-mini"),
base_url=os.environ["JUDGE_LLM_BASE_URL"],
api_key=os.environ["JUDGE_LLM_API_KEY"],
provider="openai",
thinking=False,
breaker=_make_breaker(),
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,
)
# Embedding
from adapters.embedding import LocalEmbeddingProvider, RemoteEmbeddingProvider
embed_api_key = os.environ.get("EMBED_API_KEY", "")
embed_api_url = os.environ.get("EMBED_API_URL", "")
embed_model = os.environ.get("EMBED_MODEL", "BAAI/bge-base-zh-v1.5")
embed_dim = int(os.environ.get("EMBED_DIM", "768"))
if embed_api_key and embed_api_url:
embed_provider = RemoteEmbeddingProvider(
model_name=embed_model,
embed_dim=embed_dim,
api_key=embed_api_key,
api_url=embed_api_url,
)
else:
embed_device = os.environ.get("EMBED_DEVICE", "cpu")
embed_provider = LocalEmbeddingProvider(
model_name=embed_model,
embed_dim=embed_dim,
device=embed_device,
)
embed_fn = embed_provider.embed
# Phase 6: 加载 TreeIndex
from app.tree.index import TreeIndex
trees: dict[str, TreeIndex] = {}
for vid in video_ids:
tree_path = videos_dir / vid / "tree.json"
try:
trees[vid] = TreeIndex.load_json(str(tree_path))
except Exception as exc:
logger.warning("加载树 {} 失败,跳过: {}", tree_path, exc)
if not trees:
logger.error("所有视频的树加载均失败,无法继续")
sys.exit(1)
# 将相对帧路径解析为绝对路径(与 v1 generate 一致)
for vid, tree in trees.items():
video_dir = videos_dir / vid
for l1 in tree.roots:
for l2 in l1.children:
for l3 in l2.children:
if l3.frame_path and not Path(l3.frame_path).is_absolute():
l3.frame_path = str(video_dir / l3.frame_path)
logger.info("成功加载 {} / {} 棵视频树", len(trees), len(video_ids))
# Phase 7: 初始化 QuestionGenStore
from app.question_gen.run_store import QuestionGenStore
db_path = args.db_path.resolve()
db_path.parent.mkdir(parents=True, exist_ok=True)
store = QuestionGenStore(str(db_path))
# Phase 8: 加载断点续跑进度(DB 进度 + 已有 JSON 中已满的 slot
progress: dict[str, str] = store.load_progress()
output_path_check = config.output_dir / "accepted_questions.json"
if output_path_check.exists():
try:
existing_qs = json.loads(output_path_check.read_text(encoding="utf-8"))
filled_slots = 0
for q in existing_qs:
qid = q.get("question_id", "")
vid = q.get("video_id", "")
if not qid or not vid or not qid.startswith(vid + "_"):
logger.warning("跳过格式异常的已有题目: question_id={}", qid)
continue
slot_id = qid[len(vid) + 1 :]
if slot_id not in progress:
progress[slot_id] = "accepted"
filled_slots += 1
logger.info(
"从已有 JSON 补充 progress: {} 个 slot 标记为已完成 (已有 {} 题)",
filled_slots,
len(existing_qs),
)
except (json.JSONDecodeError, OSError, ValueError) as exc:
logger.warning("已有 JSON 读取失败,跳过 progress 补充: {}", exc)
# Phase 9: 准备实时持久化回调(v1 式逐题追加,崩溃最多丢一题)
output_dir = config.output_dir
output_dir.mkdir(parents=True, exist_ok=True)
output_path = output_dir / "accepted_questions.json"
_accepted_texts: set[str] = set()
if output_path.exists():
try:
_existing = json.loads(output_path.read_text(encoding="utf-8"))
_accepted_texts = {q["question"] for q in _existing}
except (json.JSONDecodeError, OSError):
pass
def _on_accept(q: GeneratedQuestion) -> None:
"""每接受一题立即追加到 JSON(原子写入)。"""
if q.question in _accepted_texts:
return
_accepted_texts.add(q.question)
existing: list[dict] = []
if output_path.exists():
try:
existing = json.loads(output_path.read_text(encoding="utf-8"))
except (json.JSONDecodeError, OSError):
existing = []
existing.append(_question_to_entry(q))
tmp_path = output_path.with_suffix(".tmp")
tmp_path.write_text(
json.dumps(existing, ensure_ascii=False, indent=2),
encoding="utf-8",
)
os.replace(str(tmp_path), str(output_path))
# Phase 10: 运行管线(on_accept 实时持久化每道接受的题)
active_video_ids = [vid for vid in video_ids if vid in trees]
result = await run_pipeline_v2(
video_ids=active_video_ids,
trees=trees,
vlm=vlm,
llm=llm,
embed_fn=embed_fn,
store=store,
config=config,
task_types=task_types,
progress=progress,
on_accept=_on_accept,
)
final_count = (
len(json.loads(output_path.read_text(encoding="utf-8"))) if output_path.exists() else 0
)
logger.info("输出已保存: {} (总计 {} 题)", output_path, final_count)
# 统计报告
logger.info("=" * 60)
logger.info(
"管线完成: accepted={}, rejected={}, heavy_sampled={}",
len(result.accepted),
result.rejected_count,
len(result.heavy_sampled),
)
logger.info("=" * 60)
if result.rejected_count > total_slots * 0.5:
logger.warning("超过 50% 的 slot 被拒绝,建议检查 VLM/门控配置")
def _parse_args() -> argparse.Namespace:
"""解析命令行参数。"""
parser = argparse.ArgumentParser(description="赛题生成工具:generate + calibrate + generate-v2")
subparsers = parser.add_subparsers(dest="command", required=True)
# generate-v2 子命令
_add_generate_v2_parser(subparsers)
# generate 子命令
gen_parser = subparsers.add_parser("generate", help="生成新题目(v1 传统模式)")
gen_parser.add_argument(
"--store-dir",
type=str,
required=True,
help="store 根目录(包含 videos/ 和 questions/",
)
gen_parser.add_argument(
"--output-dir",
type=str,
required=True,
help="输出目录(生成的 JSON 写入此处)",
)
gen_parser.add_argument(
"--per-type",
type=int,
required=True,
help="每种题型生成数量",
)
gen_parser.add_argument(
"--similarity-threshold",
type=float,
required=True,
help="embedding 去重阈值(余弦相似度)",
)
gen_parser.add_argument(
"--max-retries",
type=int,
required=True,
help="每个 slot 最大重试次数",
)
gen_parser.add_argument(
"--concurrency",
type=int,
required=True,
help="VLM 调用并发数(Semaphore 容量)",
)
gen_parser.add_argument(
"--seed",
type=int,
required=True,
help="随机种子",
)
# calibrate 子命令(纯对比,不跑推理)
cal_parser = subparsers.add_parser(
"calibrate",
help="对比两组已有推理结果的正确率(推理请先用 main.py --mode infer",
)
cal_parser.add_argument(
"--baseline-db",
type=str,
required=True,
help="基线推理结果的 SQLite 数据库路径",
)
cal_parser.add_argument(
"--baseline-run-id",
type=str,
required=True,
help="基线推理的 run_id",
)
cal_parser.add_argument(
"--target-db",
type=str,
required=True,
help="待对比推理结果的 SQLite 数据库路径",
)
cal_parser.add_argument(
"--target-run-id",
type=str,
required=True,
help="待对比推理的 run_id",
)
cal_parser.add_argument(
"--tolerance",
type=float,
default=0.10,
help="正确率差值容忍阈值(默认 0.10)",
)
cal_parser.add_argument(
"--alpha",
type=float,
default=0.05,
help="Fisher 检验显著性水平(默认 0.05",
)
return parser.parse_args()
def main() -> None:
"""同步入口。"""
args = _parse_args()
(PROJECT_ROOT / "logs").mkdir(exist_ok=True)
if args.command == "generate":
asyncio.run(_run_generate(args))
elif args.command == "generate-v2":
asyncio.run(_run_generate_v2(args))
elif args.command == "calibrate":
_run_calibrate(args)
if __name__ == "__main__":
main()