feat(question_gen): add generate-v2 CLI subcommand and experiment script

- 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>
This commit is contained in:
2026-07-12 00:04:42 -04:00
parent f46e87258c
commit eecb86e27a
6 changed files with 533 additions and 14 deletions
+318 -12
View File
@@ -406,8 +406,6 @@ def _judge_task_type(
return "WARN"
def _calibrate_exit_code(verdicts: dict[str, str]) -> int:
"""根据所有题型的判定结果决定进程退出码。
@@ -473,8 +471,6 @@ def _read_baseline_per_task_type(
return result
def _format_comparison_table(
bench_per_task: dict[str, dict],
gen_per_task: dict[str, dict],
@@ -534,18 +530,22 @@ def _run_calibrate(args: argparse.Namespace) -> None:
# Phase 1: 从 DB 读取两组推理结果
logger.info(
"读取 baseline: db={}, run_id={}",
args.baseline_db, args.baseline_run_id,
args.baseline_db,
args.baseline_run_id,
)
baseline_per_task = _read_baseline_per_task_type(
args.baseline_db, args.baseline_run_id,
args.baseline_db,
args.baseline_run_id,
)
logger.info(
"读取 target: db={}, run_id={}",
args.target_db, args.target_run_id,
args.target_db,
args.target_run_id,
)
target_per_task = _read_baseline_per_task_type(
args.target_db, args.target_run_id,
args.target_db,
args.target_run_id,
)
baseline_total = sum(v["total"] for v in baseline_per_task.values())
@@ -583,7 +583,10 @@ def _run_calibrate(args: argparse.Namespace) -> None:
# Phase 3: 输出比较表
table_str = _format_comparison_table(
baseline_per_task, target_per_task, verdicts, p_values,
baseline_per_task,
target_per_task,
verdicts,
p_values,
)
logger.info("校准比较表:\n{}", table_str)
@@ -760,7 +763,12 @@ async def _run_generate(args: argparse.Namespace) -> None:
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))
and is_duplicate(
candidate.question,
pool,
embed_fn,
_PER_TYPE_THRESHOLD.get(task_type, similarity_threshold),
)
):
logger.warning("去重: {} 与池中题目相似", candidate.question_id)
continue
@@ -811,13 +819,309 @@ async def _run_generate(args: argparse.Namespace) -> None:
# ---------------------------------------------------------------------------
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")
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="生成新题目")
gen_parser = subparsers.add_parser("generate", help="生成新题目v1 传统模式)")
gen_parser.add_argument(
"--store-dir",
type=str,
@@ -913,6 +1217,8 @@ def main() -> None:
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)