#!/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, candidate_pool_size=config.candidate_pool_size, selector_delta_low=config.selector_delta_low, selector_delta_high=config.selector_delta_high, ) 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 _load_trees_abs(videos_dir: Path, video_ids: list[str]) -> dict: """加载视频树并将相对帧路径解析为绝对路径(复用 generate-v2 逻辑)。 参数: videos_dir: store/videos 目录。 video_ids: 待加载的 video_id 列表。 返回: video_id → TreeIndex 映射(加载失败的视频被跳过)。 """ from app.tree.index import TreeIndex trees: dict = {} for vid in video_ids: tree_path = videos_dir / vid / "tree.json" try: tree = TreeIndex.load_json(str(tree_path)) except (OSError, ValueError, KeyError) as exc: logger.warning("加载树 {} 失败,跳过: {}", tree_path, exc) continue 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) trees[vid] = tree return trees def _add_adversarial_filter_parser(subparsers: argparse._SubParsersAction) -> None: """注册 adversarial-filter 子命令(Phase B 后置对抗过滤 CLI 入口)。 参数: subparsers: argparse 子命令注册器。 """ p = subparsers.add_parser( "adversarial-filter", help="Phase B 后置对抗过滤(作弊门 + 翻转门 + 缺额补生成)", ) p.add_argument("--config", type=Path, required=True, help="YAML 配置(含 question_gen_v2 / adversarial_filter / embed 段)") p.add_argument("--store-dir", type=Path, required=True, help="store 根目录(含 videos/ prompts/ skills/)") p.add_argument("--accepted-path", type=Path, required=True, help="Phase A 产物 accepted_questions.json 路径(只读)") p.add_argument("--final-path", type=Path, default=None, help="最终题库输出路径(默认 accepted 同目录 accepted_questions_final.json)") p.add_argument("--db-path", type=Path, default=Path("logs/question_gen.db"), help="QuestionGenStore SQLite 路径") p.add_argument("--harness-db", type=Path, default=Path("logs/adversarial_harness.db"), help="agent 推理 HarnessLog SQLite 路径") p.add_argument("--prompts-version", type=str, default="v1", help="推理 prompt 版本目录名(store/prompts/)") p.add_argument("--skills-version", type=str, default="v1", help="推理 skill 版本目录名(store/skills/)") p.add_argument("--skill-mode", type=str, choices=["auto", "manual", "none"], default="auto", help="skill 模式") p.add_argument("--concurrency", type=int, default=4, help="agent 推理并发数") p.add_argument("--session-id", type=str, default="adversarial", help="遥测会话 ID(派生各门 run_id)") async def _run_adversarial_filter(args: argparse.Namespace) -> None: """adversarial-filter 子命令主流程。 装配 adapters(同 main._build_adapters)、InferenceDepsRouter(同 main.py 参数)、 QuestionGenStore、视频树(帧路径绝对化)、真实 _RealAgentRunner 与真实 backfill, 调 run_adversarial_filter 跑两门 + 缺额补生成迭代。 参数: args: CLI 参数(config, store_dir, accepted_path, final_path, db_path, harness_db, prompts_version, skills_version, skill_mode, concurrency, session_id)。 """ import yaml from app.harness.deps_router import InferenceDepsRouter from app.question_gen.adversarial_config import load_adversarial_config from app.question_gen.adversarial_filter import ( _RealAgentRunner, build_backfill, run_adversarial_filter, ) from app.question_gen.pipeline_v2 import load_pipeline_config from app.question_gen.run_store import QuestionGenStore from main import InfraSettings, _build_adapters config_path = args.config.resolve() store_dir = args.store_dir.resolve() # Phase 1: 加载配置(对抗过滤 + 出题管线 + embed 段) with config_path.open(encoding="utf-8") as f: raw_yaml = yaml.safe_load(f) or {} embed_cfg = raw_yaml.get("embed", {}) filter_config = load_adversarial_config(config_path) pipeline_config = load_pipeline_config(config_path) # Phase 2: 装配 adapters settings = InfraSettings() adapters = _build_adapters(settings, embed_cfg) # Phase 3: 加载视频树(帧路径绝对化) 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() ) trees = _load_trees_abs(videos_dir, video_ids) if not trees: logger.error("所有视频树加载失败,无法继续") sys.exit(1) logger.info("成功加载 {} / {} 棵视频树", len(trees), len(video_ids)) # Phase 4: InferenceDepsRouter(同 main.py 参数) router = InferenceDepsRouter( store_dir=store_dir, embed_provider=adapters.embed, llm=adapters.llm, vlm=adapters.vlm, ocr=adapters.ocr, default_prompts_dir=store_dir / "prompts" / args.prompts_version, default_skills_dir=store_dir / "skills" / args.skills_version, skill_mode=args.skill_mode, verify_vision=True, anchor=True, assemble_mode="ids_expand", ) # Phase 5: QuestionGenStore + 真实 agent + 真实 backfill db_path = args.db_path.resolve() db_path.parent.mkdir(parents=True, exist_ok=True) store = QuestionGenStore(str(db_path)) harness_db = args.harness_db.resolve() harness_db.parent.mkdir(parents=True, exist_ok=True) agent = _RealAgentRunner( llm=adapters.llm, tool_dispatch_fn=router.create_dispatch(), prompt_builder=router.create_prompt_builder(), db_path=str(harness_db), concurrency=args.concurrency, skill_mode=args.skill_mode, model=settings.search_llm_model, ) backfill = build_backfill( trees=trees, vlm=adapters.vlm, llm=adapters.llm, embed_fn=adapters.embed.embed, store=store, pipeline_config=pipeline_config, filter_task_types=filter_config.filter_task_types, ) # Phase 6: 运行对抗过滤 accepted_path = args.accepted_path.resolve() final_path = ( args.final_path.resolve() if args.final_path is not None else accepted_path.parent / "accepted_questions_final.json" ) await run_adversarial_filter( accepted_path=accepted_path, final_path=final_path, agent=agent, vlm=adapters.vlm, trees=trees, store=store, filter_config=filter_config, backfill=backfill, session_id=args.session_id, ) store.close() logger.info("对抗过滤完成,final 已写入: {}", final_path) 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) # adversarial-filter 子命令(Phase B 后置对抗过滤) _add_adversarial_filter_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 == "adversarial-filter": asyncio.run(_run_adversarial_filter(args)) elif args.command == "calibrate": _run_calibrate(args) if __name__ == "__main__": main()