eecb86e27a
- Add generate-v2 subparser with --config, --store-dir, --db-path, --seed, and --dry-run arguments to tools/generate_questions.py - Implement _run_generate_v2 async handler: config loading, video discovery, DI client construction, TreeIndex loading, pipeline invocation, and result persistence - Add scripts/generate_questions_v2.sh following build_trees.sh conventions (source .env, conda run python path, MODE=mock support) - Update app/question_gen/__init__.py to export full v2 public API: run_pipeline_v2, PipelineConfig, PipelineResult, QuestionFamilySpec, ALL_FAMILIES, CandidateQuestion, generate_one_v2, GateReport, run_gates - Add QuestionGenStore.load_progress() for pipeline resumption - Add integration tests for CLI help and dry-run behavior - Update test_question_gen_api to match expanded __all__ Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1228 lines
39 KiB
Python
1228 lines
39 KiB
Python
#!/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 _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 _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(
|
||
"--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(
|
||
family_ratios=config.family_ratios,
|
||
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,
|
||
gate_models=config.gate_models,
|
||
heavy_agent_model=config.heavy_agent_model,
|
||
)
|
||
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 模式
|
||
task_types = [
|
||
"Action Recognition",
|
||
"Action Reasoning",
|
||
"Action Prediction",
|
||
"Action Sequence",
|
||
"Object Recognition",
|
||
"Object Reasoning",
|
||
"Object Interaction",
|
||
"Scene Understanding",
|
||
"Event Reasoning",
|
||
"Causal Reasoning",
|
||
"Temporal Reasoning",
|
||
"Spatial Reasoning",
|
||
]
|
||
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] 家族权重: {}", config.family_ratios)
|
||
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"))
|
||
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 客户端(门控用)
|
||
llm = GovernedLLMClient(
|
||
model=os.environ.get("LLM_MODEL", "gpt-4.1-mini"),
|
||
base_url=os.environ["LLM_BASE_URL"],
|
||
api_key=os.environ["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)
|
||
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: 加载断点续跑进度
|
||
progress: dict[str, str] = store.load_progress()
|
||
|
||
# Phase 9: 运行管线
|
||
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,
|
||
)
|
||
|
||
# Phase 10: 保存输出
|
||
output_dir = config.output_dir
|
||
output_dir.mkdir(parents=True, exist_ok=True)
|
||
output_path = output_dir / "accepted_questions.json"
|
||
|
||
accepted_data = []
|
||
for q in result.accepted:
|
||
accepted_data.append(
|
||
{
|
||
"question_id": q.question_id,
|
||
"video_id": q.video_id,
|
||
"task_type": q.task_type,
|
||
"question": q.question,
|
||
"options": list(q.options),
|
||
"answer": q.answer,
|
||
"source_nodes": list(q.source_nodes),
|
||
"difficulty": q.difficulty,
|
||
"skill_target": q.skill_target,
|
||
}
|
||
)
|
||
|
||
output_path.write_text(
|
||
json.dumps(accepted_data, ensure_ascii=False, indent=2),
|
||
encoding="utf-8",
|
||
)
|
||
logger.info(
|
||
"输出已保存: {} ({} 题)",
|
||
output_path,
|
||
len(accepted_data),
|
||
)
|
||
|
||
# 统计报告
|
||
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()
|