refactor(question_gen): slim calibrate to compare two existing runs
This commit is contained in:
+74
-303
@@ -4,8 +4,11 @@
|
||||
用法:
|
||||
conda activate Video-Tree-TRM
|
||||
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 此脚本。
|
||||
"""
|
||||
|
||||
@@ -17,7 +20,6 @@ import json
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
import uuid
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
@@ -32,7 +34,6 @@ load_dotenv(PROJECT_ROOT / ".env")
|
||||
import numpy as np
|
||||
from scipy.stats import fisher_exact
|
||||
|
||||
from app.harness.log import HarnessLog
|
||||
from app.question_gen.loader import load_benchmark
|
||||
from app.question_gen.synthesizer import (
|
||||
TASK_TYPE_LEVEL_MAP,
|
||||
@@ -405,23 +406,6 @@ def _judge_task_type(
|
||||
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:
|
||||
@@ -489,140 +473,6 @@ def _read_baseline_per_task_type(
|
||||
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(
|
||||
@@ -665,98 +515,51 @@ def _format_comparison_table(
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
async def _run_calibrate(args: argparse.Namespace) -> None:
|
||||
"""calibrate 子命令主流程。
|
||||
def _run_calibrate(args: argparse.Namespace) -> None:
|
||||
"""calibrate 子命令主流程(纯对比,不跑推理)。
|
||||
|
||||
对比 benchmark 和生成题在 Agent 推理下的正确率,
|
||||
从两组已有的推理结果(harness.db + run_id)读取 per_task_type 正确率,
|
||||
逐题型 Fisher 精确检验判定校准质量。
|
||||
|
||||
参数:
|
||||
args: CLI 参数。
|
||||
"""
|
||||
_validate_calibrate_args(
|
||||
getattr(args, "baseline_db", None),
|
||||
getattr(args, "baseline_run_id", None),
|
||||
)
|
||||
推理应事先通过 main.py --mode infer 完成,确保 baseline 和 target
|
||||
使用完全相同的推理管线。
|
||||
|
||||
generated_dir = Path(args.generated_dir)
|
||||
benchmark_dir = Path(args.benchmark_dir)
|
||||
store_dir = Path(args.store_dir)
|
||||
db_path = args.db_path
|
||||
prompts_dir = Path(args.prompts_dir)
|
||||
concurrency = args.concurrency
|
||||
max_steps = args.max_steps
|
||||
skill_mode = args.skill_mode
|
||||
参数:
|
||||
args: CLI 参数(baseline_db, baseline_run_id, target_db, target_run_id,
|
||||
tolerance, alpha)。
|
||||
"""
|
||||
tolerance = args.tolerance
|
||||
alpha = args.alpha
|
||||
|
||||
# Phase 1: 加载题目
|
||||
logger.info("加载生成题目: {}", generated_dir)
|
||||
gen_questions = load_benchmark(generated_dir)
|
||||
logger.info("加载 benchmark 题目: {}", benchmark_dir)
|
||||
bench_questions = load_benchmark(benchmark_dir)
|
||||
logger.info("生成题 {} 道, benchmark {} 道", len(gen_questions), len(bench_questions))
|
||||
|
||||
# Phase 2: 获取 benchmark baseline
|
||||
baseline_db = getattr(args, "baseline_db", None)
|
||||
baseline_run_id = getattr(args, "baseline_run_id", None)
|
||||
|
||||
if baseline_db and baseline_run_id:
|
||||
logger.info("从基线 DB 读取 benchmark 指标: db={}, run_id={}", baseline_db, baseline_run_id)
|
||||
bench_per_task = _read_baseline_per_task_type(baseline_db, baseline_run_id)
|
||||
else:
|
||||
logger.info("运行 benchmark 推理以获取基线指标")
|
||||
llm = _build_llm_client()
|
||||
vlm = _build_vlm_client()
|
||||
embed_provider = _build_embed_provider()
|
||||
bench_run_id = f"calibrate-bench-{uuid.uuid4().hex[:8]}"
|
||||
bench_per_task = await _run_inference_for_questions(
|
||||
bench_questions,
|
||||
store_dir=store_dir,
|
||||
prompts_dir=prompts_dir,
|
||||
db_path=db_path,
|
||||
run_id=bench_run_id,
|
||||
concurrency=concurrency,
|
||||
max_steps=max_steps,
|
||||
skill_mode=skill_mode,
|
||||
llm=llm,
|
||||
vlm=vlm,
|
||||
embed_provider=embed_provider,
|
||||
)
|
||||
|
||||
# Phase 3: 运行生成题推理
|
||||
logger.info("运行生成题推理")
|
||||
if not baseline_db:
|
||||
# 客户端已在 Phase 2 构建
|
||||
pass
|
||||
else:
|
||||
llm = _build_llm_client()
|
||||
vlm = _build_vlm_client()
|
||||
embed_provider = _build_embed_provider()
|
||||
|
||||
gen_run_id = f"calibrate-gen-{uuid.uuid4().hex[:8]}"
|
||||
gen_per_task = await _run_inference_for_questions(
|
||||
gen_questions,
|
||||
store_dir=store_dir,
|
||||
prompts_dir=prompts_dir,
|
||||
db_path=db_path,
|
||||
run_id=gen_run_id,
|
||||
concurrency=concurrency,
|
||||
max_steps=max_steps,
|
||||
skill_mode=skill_mode,
|
||||
llm=llm,
|
||||
vlm=vlm,
|
||||
embed_provider=embed_provider,
|
||||
# Phase 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,
|
||||
)
|
||||
|
||||
# Phase 4: 逐题型判定
|
||||
all_types = sorted(set(bench_per_task) | set(gen_per_task))
|
||||
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 = bench_per_task.get(task_type)
|
||||
g = gen_per_task.get(task_type)
|
||||
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")
|
||||
@@ -771,7 +574,6 @@ async def _run_calibrate(args: argparse.Namespace) -> None:
|
||||
alpha=alpha,
|
||||
)
|
||||
|
||||
# 计算 p-value 供表格显示
|
||||
table = [
|
||||
[b["correct"], b["total"] - b["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_values[task_type] = p_val
|
||||
|
||||
# Phase 5: 输出比较表
|
||||
table_str = _format_comparison_table(bench_per_task, gen_per_task, verdicts, p_values)
|
||||
# Phase 3: 输出比较表
|
||||
table_str = _format_comparison_table(
|
||||
baseline_per_task, target_per_task, verdicts, p_values,
|
||||
)
|
||||
logger.info("校准比较表:\n{}", table_str)
|
||||
|
||||
# Phase 6: 退出
|
||||
# Phase 4: 退出
|
||||
exit_code = _calibrate_exit_code(verdicts)
|
||||
if exit_code == 0:
|
||||
logger.info("校准通过: 所有题型 PASS 或 WARN")
|
||||
@@ -1057,79 +861,46 @@ def _parse_args() -> argparse.Namespace:
|
||||
help="随机种子",
|
||||
)
|
||||
|
||||
# calibrate 子命令
|
||||
cal_parser = subparsers.add_parser("calibrate", help="校准生成题与 benchmark 难度一致性")
|
||||
cal_parser.add_argument(
|
||||
"--generated-dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="生成题目目录",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--benchmark-dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="benchmark 题目目录",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--store-dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="store 根目录",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--db-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="校准 SQLite 数据库路径",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--prompts-dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="prompt 文件目录",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--concurrency",
|
||||
type=int,
|
||||
required=True,
|
||||
help="推理并发数",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--max-steps",
|
||||
type=int,
|
||||
required=True,
|
||||
help="AgentLoop 单题最大步数",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--skill-mode",
|
||||
type=str,
|
||||
required=True,
|
||||
help="skill 模式 (auto/manual/none)",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--tolerance",
|
||||
type=float,
|
||||
required=True,
|
||||
help="正确率差值容忍阈值",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--alpha",
|
||||
type=float,
|
||||
required=True,
|
||||
help="Fisher 检验显著性水平",
|
||||
# calibrate 子命令(纯对比,不跑推理)
|
||||
cal_parser = subparsers.add_parser(
|
||||
"calibrate",
|
||||
help="对比两组已有推理结果的正确率(推理请先用 main.py --mode infer)",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--baseline-db",
|
||||
type=str,
|
||||
default=None,
|
||||
help="基线数据库路径(可选,须与 --baseline-run-id 成对)",
|
||||
required=True,
|
||||
help="基线推理结果的 SQLite 数据库路径",
|
||||
)
|
||||
cal_parser.add_argument(
|
||||
"--baseline-run-id",
|
||||
type=str,
|
||||
default=None,
|
||||
help="基线运行标识(可选,须与 --baseline-db 成对)",
|
||||
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()
|
||||
@@ -1143,7 +914,7 @@ def main() -> None:
|
||||
if args.command == "generate":
|
||||
asyncio.run(_run_generate(args))
|
||||
elif args.command == "calibrate":
|
||||
asyncio.run(_run_calibrate(args))
|
||||
_run_calibrate(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user