refactor(question_gen): slim calibrate to compare two existing runs

This commit is contained in:
2026-07-11 07:45:42 -04:00
parent 307c64c388
commit da70eb6e23
+74 -303
View File
@@ -4,8 +4,11 @@
用法: 用法:
conda activate Video-Tree-TRM conda activate Video-Tree-TRM
python tools/generate_questions.py generate --store-dir store ... python tools/generate_questions.py generate --store-dir store ...
python tools/generate_questions.py calibrate ... 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 此脚本。 app/core/adapters 不 import 此脚本。
""" """
@@ -17,7 +20,6 @@ import json
import os import os
import random import random
import sys import sys
import uuid
from collections import defaultdict from collections import defaultdict
from pathlib import Path from pathlib import Path
@@ -32,7 +34,6 @@ load_dotenv(PROJECT_ROOT / ".env")
import numpy as np import numpy as np
from scipy.stats import fisher_exact from scipy.stats import fisher_exact
from app.harness.log import HarnessLog
from app.question_gen.loader import load_benchmark from app.question_gen.loader import load_benchmark
from app.question_gen.synthesizer import ( from app.question_gen.synthesizer import (
TASK_TYPE_LEVEL_MAP, TASK_TYPE_LEVEL_MAP,
@@ -405,23 +406,6 @@ def _judge_task_type(
return "WARN" 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: def _calibrate_exit_code(verdicts: dict[str, str]) -> int:
@@ -489,140 +473,6 @@ def _read_baseline_per_task_type(
return result 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( def _format_comparison_table(
@@ -665,98 +515,51 @@ def _format_comparison_table(
return "\n".join(lines) return "\n".join(lines)
async def _run_calibrate(args: argparse.Namespace) -> None: def _run_calibrate(args: argparse.Namespace) -> None:
"""calibrate 子命令主流程。 """calibrate 子命令主流程(纯对比,不跑推理)
对比 benchmark 和生成题在 Agent 推理下的正确率, 从两组已有的推理结果(harness.db + run_id)读取 per_task_type 正确率,
逐题型 Fisher 精确检验判定校准质量。 逐题型 Fisher 精确检验判定校准质量。
参数: 推理应事先通过 main.py --mode infer 完成,确保 baseline 和 target
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) args: CLI 参数(baseline_db, baseline_run_id, target_db, target_run_id,
store_dir = Path(args.store_dir) tolerance, alpha)。
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 tolerance = args.tolerance
alpha = args.alpha alpha = args.alpha
# Phase 1: 加载题目 # Phase 1: 从 DB 读取两组推理结果
logger.info("加载生成题目: {}", generated_dir) logger.info(
gen_questions = load_benchmark(generated_dir) "读取 baseline: db={}, run_id={}",
logger.info("加载 benchmark 题目: {}", benchmark_dir) args.baseline_db, args.baseline_run_id,
bench_questions = load_benchmark(benchmark_dir) )
logger.info("生成题 {} 道, benchmark {}", len(gen_questions), len(bench_questions)) baseline_per_task = _read_baseline_per_task_type(
args.baseline_db, args.baseline_run_id,
# 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: 逐题型判定 logger.info(
all_types = sorted(set(bench_per_task) | set(gen_per_task)) "读取 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] = {} verdicts: dict[str, str] = {}
p_values: dict[str, float] = {} p_values: dict[str, float] = {}
for task_type in all_types: for task_type in all_types:
b = bench_per_task.get(task_type) b = baseline_per_task.get(task_type)
g = gen_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: if b is None or g is None or b["total"] == 0 or g["total"] == 0:
verdicts[task_type] = "WARN" verdicts[task_type] = "WARN"
p_values[task_type] = float("nan") p_values[task_type] = float("nan")
@@ -771,7 +574,6 @@ async def _run_calibrate(args: argparse.Namespace) -> None:
alpha=alpha, alpha=alpha,
) )
# 计算 p-value 供表格显示
table = [ table = [
[b["correct"], b["total"] - b["correct"]], [b["correct"], b["total"] - b["correct"]],
[g["correct"], g["total"] - g["correct"]], [g["correct"], g["total"] - g["correct"]],
@@ -779,11 +581,13 @@ async def _run_calibrate(args: argparse.Namespace) -> None:
_, p_val = fisher_exact(table) _, p_val = fisher_exact(table)
p_values[task_type] = p_val p_values[task_type] = p_val
# Phase 5: 输出比较表 # Phase 3: 输出比较表
table_str = _format_comparison_table(bench_per_task, gen_per_task, verdicts, p_values) table_str = _format_comparison_table(
baseline_per_task, target_per_task, verdicts, p_values,
)
logger.info("校准比较表:\n{}", table_str) logger.info("校准比较表:\n{}", table_str)
# Phase 6: 退出 # Phase 4: 退出
exit_code = _calibrate_exit_code(verdicts) exit_code = _calibrate_exit_code(verdicts)
if exit_code == 0: if exit_code == 0:
logger.info("校准通过: 所有题型 PASS 或 WARN") logger.info("校准通过: 所有题型 PASS 或 WARN")
@@ -1057,79 +861,46 @@ def _parse_args() -> argparse.Namespace:
help="随机种子", help="随机种子",
) )
# calibrate 子命令 # calibrate 子命令(纯对比,不跑推理)
cal_parser = subparsers.add_parser("calibrate", help="校准生成题与 benchmark 难度一致性") cal_parser = subparsers.add_parser(
cal_parser.add_argument( "calibrate",
"--generated-dir", help="对比两组已有推理结果的正确率(推理请先用 main.py --mode infer",
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( cal_parser.add_argument(
"--baseline-db", "--baseline-db",
type=str, type=str,
default=None, required=True,
help="基线数据库路径(可选,须与 --baseline-run-id 成对)", help="基线推理结果的 SQLite 数据库路径",
) )
cal_parser.add_argument( cal_parser.add_argument(
"--baseline-run-id", "--baseline-run-id",
type=str, type=str,
default=None, required=True,
help="基线运行标识(可选,须与 --baseline-db 成对)", 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() return parser.parse_args()
@@ -1143,7 +914,7 @@ def main() -> None:
if args.command == "generate": if args.command == "generate":
asyncio.run(_run_generate(args)) asyncio.run(_run_generate(args))
elif args.command == "calibrate": elif args.command == "calibrate":
asyncio.run(_run_calibrate(args)) _run_calibrate(args)
if __name__ == "__main__": if __name__ == "__main__":