Files
Video-Tree-TRM5/tools/generate_questions.py
T
iomgaa 25c8d5ec42 fix(question_gen): 每次重试换视频避免同视频反复去重失败
Information Synopsis 163 道 benchmark 题,同一视频反复出题
必然相似。改为每次重试随机选不同视频。

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-07-10 08:15:03 -04:00

1146 lines
36 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 ...
app/core/adapters 不 import 此脚本。
"""
from __future__ import annotations
import argparse
import asyncio
import json
import os
import random
import sys
import uuid
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.harness.log import HarnessLog
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 _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_id": question.question_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,
}
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 _validate_calibrate_args(
baseline_db: str | None,
baseline_run_id: str | None,
) -> None:
"""校验 baseline 参数必须成对出现。
参数:
baseline_db: 基线数据库路径。
baseline_run_id: 基线运行标识。
异常:
ValueError: 只提供了一个而非两个参数。
"""
has_db = baseline_db is not None
has_run_id = baseline_run_id is not None
if has_db != has_run_id:
raise ValueError("--baseline-db 和 --baseline-run-id 必须成对出现")
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 _build_llm_client():
"""构建 GovernedLLMClient(推理用 LLM)。
从 .env 读取 SEARCH_LLM_MODEL / SEARCH_LLM_BASE_URL / SEARCH_LLM_API_KEY
和 LLM 韧性参数。
返回:
GovernedLLMClient 实例。
"""
from adapters.breaker import CircuitBreaker
from adapters.llm import GovernedLLMClient
from adapters.telemetry import SQLiteTelemetryRecorder
(PROJECT_ROOT / "logs").mkdir(exist_ok=True)
telemetry = SQLiteTelemetryRecorder(str(PROJECT_ROOT / "logs" / "calibrate_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"))
return 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,
)
async def _run_inference_for_questions(
questions: list,
*,
store_dir: Path,
prompts_dir: Path,
db_path: str,
run_id: str,
concurrency: int,
max_steps: int,
skill_mode: str,
llm,
vlm,
embed_provider,
) -> dict[str, dict]:
"""对题目列表运行推理,返回 per_task_type 指标。
按 video_id 分组,逐组构建推理依赖并执行推理,
最后合并所有组的 per_task_type 结果。
参数:
questions: 待推理的题目列表。
store_dir: store 根目录。
prompts_dir: prompt 文件目录。
db_path: SQLite 数据库路径。
run_id: 运行标识。
concurrency: 最大并发数。
max_steps: AgentLoop 单题最大步数。
skill_mode: skill 模式。
llm: LLMProvider 实例。
vlm: VLMProvider 实例。
embed_provider: EmbeddingProvider 实例。
返回:
{task_type: {"accuracy": float, "total": int, "correct": int}}。
"""
from app.harness.factory import build_inference_deps
from app.harness.inference import run_inference
# Phase 1: 按 video_id 分组
by_video: dict[str, list] = defaultdict(list)
for q in questions:
by_video[q.video_id].append(q)
# Phase 2: 逐组推理
all_per_task: dict[str, dict] = {}
skills_dir = store_dir / "skills"
if not skills_dir.exists():
skills_dir = None
with HarnessLog(db_path, run_id) as log:
for video_id, group in by_video.items():
deps = build_inference_deps(
store_dir=store_dir,
video_id=video_id,
prompts_dir=prompts_dir,
skills_dir=skills_dir,
skill_mode=skill_mode,
embed_provider=embed_provider,
llm=llm,
vlm=vlm,
ocr=None,
verify_vision=False,
anchor=False,
assemble_mode="ids",
)
result = await run_inference(
group,
llm=deps.llm,
tool_dispatch_fn=deps.tool_dispatch_fn,
prompt_builder=deps.prompt_builder,
log=log,
run_id=run_id,
concurrency=concurrency,
max_steps=max_steps,
skill_mode=skill_mode,
)
# Phase 3: 合并 per_task_type
for task_type, metrics in result.per_task_type.items():
if task_type in all_per_task:
existing = all_per_task[task_type]
merged_total = existing["total"] + metrics["total"]
merged_correct = existing["correct"] + metrics["correct"]
all_per_task[task_type] = {
"accuracy": merged_correct / merged_total,
"total": merged_total,
"correct": merged_correct,
}
else:
all_per_task[task_type] = dict(metrics)
return all_per_task
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)
async def _run_calibrate(args: argparse.Namespace) -> None:
"""calibrate 子命令主流程。
对比 benchmark 和生成题在 Agent 推理下的正确率,
逐题型 Fisher 精确检验判定校准质量。
参数:
args: CLI 参数。
"""
_validate_calibrate_args(
getattr(args, "baseline_db", None),
getattr(args, "baseline_run_id", None),
)
generated_dir = Path(args.generated_dir)
benchmark_dir = Path(args.benchmark_dir)
store_dir = Path(args.store_dir)
db_path = args.db_path
prompts_dir = Path(args.prompts_dir)
concurrency = args.concurrency
max_steps = args.max_steps
skill_mode = args.skill_mode
tolerance = args.tolerance
alpha = args.alpha
# Phase 1: 加载题目
logger.info("加载生成题目: {}", generated_dir)
gen_questions = load_benchmark(generated_dir)
logger.info("加载 benchmark 题目: {}", benchmark_dir)
bench_questions = load_benchmark(benchmark_dir)
logger.info("生成题 {} 道, benchmark {} 道", len(gen_questions), len(bench_questions))
# Phase 2: 获取 benchmark baseline
baseline_db = getattr(args, "baseline_db", None)
baseline_run_id = getattr(args, "baseline_run_id", None)
if baseline_db and baseline_run_id:
logger.info("从基线 DB 读取 benchmark 指标: db={}, run_id={}", baseline_db, baseline_run_id)
bench_per_task = _read_baseline_per_task_type(baseline_db, baseline_run_id)
else:
logger.info("运行 benchmark 推理以获取基线指标")
llm = _build_llm_client()
vlm = _build_vlm_client()
embed_provider = _build_embed_provider()
bench_run_id = f"calibrate-bench-{uuid.uuid4().hex[:8]}"
bench_per_task = await _run_inference_for_questions(
bench_questions,
store_dir=store_dir,
prompts_dir=prompts_dir,
db_path=db_path,
run_id=bench_run_id,
concurrency=concurrency,
max_steps=max_steps,
skill_mode=skill_mode,
llm=llm,
vlm=vlm,
embed_provider=embed_provider,
)
# Phase 3: 运行生成题推理
logger.info("运行生成题推理")
if not baseline_db:
# 客户端已在 Phase 2 构建
pass
else:
llm = _build_llm_client()
vlm = _build_vlm_client()
embed_provider = _build_embed_provider()
gen_run_id = f"calibrate-gen-{uuid.uuid4().hex[:8]}"
gen_per_task = await _run_inference_for_questions(
gen_questions,
store_dir=store_dir,
prompts_dir=prompts_dir,
db_path=db_path,
run_id=gen_run_id,
concurrency=concurrency,
max_steps=max_steps,
skill_mode=skill_mode,
llm=llm,
vlm=vlm,
embed_provider=embed_provider,
)
# Phase 4: 逐题型判定
all_types = sorted(set(bench_per_task) | set(gen_per_task))
verdicts: dict[str, str] = {}
p_values: dict[str, float] = {}
for task_type in all_types:
b = bench_per_task.get(task_type)
g = gen_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,
)
# 计算 p-value 供表格显示
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 5: 输出比较表
table_str = _format_comparison_table(bench_per_task, gen_per_task, verdicts, p_values)
logger.info("校准比较表:\n{}", table_str)
# Phase 6: 退出
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
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, 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 _parse_args() -> argparse.Namespace:
"""解析命令行参数。"""
parser = argparse.ArgumentParser(description="赛题生成工具:generate + calibrate")
subparsers = parser.add_subparsers(dest="command", required=True)
# generate 子命令
gen_parser = subparsers.add_parser("generate", help="生成新题目")
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="校准生成题与 benchmark 难度一致性")
cal_parser.add_argument(
"--generated-dir",
type=str,
required=True,
help="生成题目目录",
)
cal_parser.add_argument(
"--benchmark-dir",
type=str,
required=True,
help="benchmark 题目目录",
)
cal_parser.add_argument(
"--store-dir",
type=str,
required=True,
help="store 根目录",
)
cal_parser.add_argument(
"--db-path",
type=str,
required=True,
help="校准 SQLite 数据库路径",
)
cal_parser.add_argument(
"--prompts-dir",
type=str,
required=True,
help="prompt 文件目录",
)
cal_parser.add_argument(
"--concurrency",
type=int,
required=True,
help="推理并发数",
)
cal_parser.add_argument(
"--max-steps",
type=int,
required=True,
help="AgentLoop 单题最大步数",
)
cal_parser.add_argument(
"--skill-mode",
type=str,
required=True,
help="skill 模式 (auto/manual/none)",
)
cal_parser.add_argument(
"--tolerance",
type=float,
required=True,
help="正确率差值容忍阈值",
)
cal_parser.add_argument(
"--alpha",
type=float,
required=True,
help="Fisher 检验显著性水平",
)
cal_parser.add_argument(
"--baseline-db",
type=str,
default=None,
help="基线数据库路径(可选,须与 --baseline-run-id 成对)",
)
cal_parser.add_argument(
"--baseline-run-id",
type=str,
default=None,
help="基线运行标识(可选,须与 --baseline-db 成对)",
)
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 == "calibrate":
asyncio.run(_run_calibrate(args))
if __name__ == "__main__":
main()