refactor: self-contained two-phase video-split CLI entry

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-15 13:12:21 -04:00
parent 6a21d80313
commit 2844732126
3 changed files with 878 additions and 124 deletions
+682
View File
@@ -0,0 +1,682 @@
"""结果驱动视频级切分的自包含两阶段 CLI 入口。
把整条离线管线的编排从 shell 搬进 Python:一次调用内联串起
Phase 1 离线诊断(run_baseline_diagnosisLLM 重活,断点续跑幂等)→
Phase 2 冻结切分(build_split,纯 code-controlled,产出 pools.json + manifest)→
McNemar 功效护栏(validation 池错题数达阈校验)。
设计要点:
- 诊断口径指纹 = (诊断 prompt 版本, 模型名, git 短 SHA) 三分量合成,隔离不同
诊断配置的信号;换 prompt / 模型 / 代码实现即换指纹,旧信号不被覆盖。
- 真实依赖组装参考 app/harness/runner.py::_run_diagnosisGovernedLLMClient
(search llm, thinking=True) + RunLogImpl(harness.db) + VersionedSkillStore +
DiagnosePrompts(项目根 prompts/) + tree_data={}(由诊断管线内部按需加载)。
- 缺 .env / config 关键项一律 fail loud(P5),绝不静默兜底。
- `--dry-run` 用假 deps 跑通两阶段 wiring 不真调 LLM,打印将执行的步骤 + 指纹,
用于校验装配正确性(对齐 CLAUDE.md §2.5 smoke test)。
编排函数(run_pipeline)通过依赖注入接收 DiagnosisDeps / signal_store / wrong_ids /
questions,便于单测用假实现替换、不触真实 LLM 与 harness.db。
"""
from __future__ import annotations
import argparse
import asyncio
import datetime
import os
import sqlite3
import subprocess
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any
import yaml
from loguru import logger
from app.harness.baseline_diagnosis import DiagnosisDeps, run_baseline_diagnosis
from app.harness.build_split import SplitBuildConfig, SplitBuildResult, build_split
from app.harness.split_selection import diag_fingerprint
from app.question_gen.loader import load_benchmark
if TYPE_CHECKING:
from app.harness.pools import Pools
from core.evolution.protocols import DiagnosisSignalStore
from core.types import GeneratedQuestion
# 与 core.evolution.diagnose._INFRA_STOP_REASONS 对齐:执行/解析层失败排除出可诊断错题。
_INFRA_STOP_REASONS: frozenset[str] = frozenset({"error", "parse_error"})
# 工程路径默认值(少变;可经 CLI 单次覆盖)。诊断信号表建在 harness.db。
_DEFAULT_HARNESS_DB = Path("workspaces/default/harness.db")
_DEFAULT_QUESTIONS_DIR = Path("store/questions/benchmarks/Video-MME")
_DEFAULT_OUT_DIR = Path("workspaces/video-split")
# ---------------------------------------------------------------------------
# 配置解析(fail loud
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class VideoSplitConfig:
"""结果驱动视频级切分的科研旋钮快照(从 config/video_split.yaml 解析)。
字段:
baseline_run_id: 基线 run 标识(错题诊断与切分依据)。
n_trainval: trainval 目标视频数(多样性阶段填充上限)。
epsilon: test 相对全局最大允许分布偏差(题型 / 难度两维)。
report_floor: per-type 报告门限,题数 ≥ 此值的 task_type 才入 ε 约束。
val_wrong_min: validation 池最少错题数(McNemar 功效阈;0=不检查)。
val_ratio: validation 占 trainval 视频组总数的比例。
seed: 贪心选择器预洗牌 + 视频组题级切分种子。
floor_k: 各高信号 task_type 的 T2 defect 下限(硬约束)。
prompt_version: 诊断 prompt 版本标识(指纹分量)。
model: 执行诊断的模型名(指纹分量)。
"""
baseline_run_id: str
n_trainval: int
epsilon: float
report_floor: int
val_wrong_min: int
val_ratio: float
seed: int
floor_k: dict[str, int]
prompt_version: str
model: str
def _require(section: dict[str, Any], keys: tuple[str, ...], where: str) -> None:
"""校验 section 含全部必填键,缺任一即 fail loud(P5,不静默兜底)。
参数:
section: 待校验的配置子字典。
keys: 必填键元组。
where: 出错信息中标注的段名(如 "video_split")。
异常:
SystemExit: 存在缺失键。
"""
missing = [k for k in keys if k not in section]
if missing:
raise SystemExit(f"config {where} 段缺关键项 {missing},无法运行(P5 fail loud")
def parse_config(raw: dict[str, Any]) -> VideoSplitConfig:
"""把 yaml 原始字典解析为 VideoSplitConfig,缺关键项 fail loud。
参数:
raw: yaml.safe_load 的顶层字典,需含 video_split / diag 两段。
返回:
VideoSplitConfig 冻结快照。
异常:
SystemExit: 缺 video_split / diag 段或段内关键项。
"""
if "video_split" not in raw or "diag" not in raw:
raise SystemExit("config 缺 video_split / diag 段,无法运行(P5 fail loud")
vs = raw["video_split"]
dg = raw["diag"]
_require(
vs,
(
"baseline_run_id",
"n_trainval",
"epsilon",
"report_floor",
"val_wrong_min",
"val_ratio",
"seed",
"floor_k",
),
"video_split",
)
_require(dg, ("prompt_version", "model"), "diag")
return VideoSplitConfig(
baseline_run_id=vs["baseline_run_id"],
n_trainval=vs["n_trainval"],
epsilon=vs["epsilon"],
report_floor=vs["report_floor"],
val_wrong_min=vs["val_wrong_min"],
val_ratio=vs["val_ratio"],
seed=vs["seed"],
floor_k=dict(vs["floor_k"]),
prompt_version=dg["prompt_version"],
model=dg["model"],
)
def load_config(config_path: Path) -> VideoSplitConfig:
"""读取并解析 video_split yaml 配置文件(缺文件 / 关键项 fail loud)。
参数:
config_path: yaml 配置路径。
返回:
VideoSplitConfig。
异常:
SystemExit: 文件不存在或缺关键项。
"""
if not config_path.exists():
raise SystemExit(f"config 文件不存在: {config_path}P5 fail loud")
raw = yaml.safe_load(config_path.read_text(encoding="utf-8"))
return parse_config(raw)
def git_short_sha() -> str:
"""取当前 git 短 SHA 作为诊断口径指纹的代码分量(诊断代码变则指纹变)。
返回:
git rev-parse --short HEAD 输出(去空白)。
异常:
SystemExit: 非 git 仓库或 git 不可用(fail loud,指纹不可缺分量)。
"""
try:
out = subprocess.run(
["git", "rev-parse", "--short", "HEAD"],
capture_output=True,
text=True,
check=True,
)
except (subprocess.CalledProcessError, FileNotFoundError) as exc:
raise SystemExit(f"无法获取 git 短 SHA 作为诊断代码版本: {exc}P5 fail loud") from exc
sha = out.stdout.strip()
if not sha:
raise SystemExit("git rev-parse --short HEAD 返回空,诊断指纹缺代码分量(P5 fail loud")
return sha
# ---------------------------------------------------------------------------
# 真实依赖组装(参考 runner.py::_run_diagnosis
# ---------------------------------------------------------------------------
class _DiagLLMSettings:
"""诊断 LLM 的工程配置(从 .env 读取 search llm 凭证 + 韧性旋钮)。
仅承载诊断所需字段(搜索 LLM = 诊断 judge),不复用 main.InfraSettings 以免
构造整套适配器(embed / vlm)的重活;缺关键凭证 fail loud。
"""
def __init__(self) -> None:
from pydantic_settings import BaseSettings, SettingsConfigDict
class _Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env", extra="ignore")
search_llm_model: str = ""
search_llm_base_url: str = ""
search_llm_api_key: str = ""
redis_url: str = ""
redis_cache_ttl: int = 86400
llm_timeout: float = 300.0
llm_max_retries: int = 3
llm_retry_base_delay: float = 20.0
llm_retry_max_delay: float = 120.0
llm_circuit_breaker_threshold: int = 48
llm_circuit_breaker_cooldown: float = 60.0
llm_ttft_timeout: float = 30.0
llm_inter_token_timeout: float = 15.0
self._s = _Settings()
def __getattr__(self, name: str) -> Any:
return getattr(self._s, name)
def _build_redis_cache(settings: Any) -> Any | None:
"""按 .env redis_url 构建响应缓存(不可用则降级 None,与 main 一致)。"""
if not settings.redis_url:
return None
try:
import redis.asyncio as aioredis
from adapters.redis_cache import RedisResponseCache
redis_client = aioredis.from_url(settings.redis_url, decode_responses=True)
ttl_s = settings.redis_cache_ttl if settings.redis_cache_ttl > 0 else None
return RedisResponseCache(redis=redis_client, ttl_s=ttl_s)
except Exception:
logger.warning("Redis 缓存不可用,诊断降级为无缓存模式")
return None
def build_diagnosis_deps(*, harness_db: Path, concurrency: int) -> DiagnosisDeps:
"""组装 Phase 1 诊断的真实依赖束(GovernedLLMClient + RunLogImpl + prompts)。
与 runner.py::_run_diagnosis 对齐:search LLMthinking=True)作诊断 judge
RunLogImpl 只读读取 harness.db 的 predictions/tracesVersionedSkillStore 读技能,
DiagnosePrompts 从项目根 prompts/ 加载,tree_data={} 由诊断管线内部按需加载。
参数:
harness_db: harness.db 路径(诊断读预测 + 信号落库同库)。
concurrency: 诊断并发上限。
返回:
DiagnosisDeps 冻结依赖束。
异常:
SystemExit: .env 缺 search LLM 凭证(model / base_url / api_key 任一为空)。
"""
from adapters.breaker import CircuitBreaker
from adapters.llm import GovernedLLMClient
from adapters.telemetry import SQLiteTelemetryRecorder
from app.harness.log import RunLogImpl
from app.harness.workspace import VersionedSkillStore
settings = _DiagLLMSettings()
if not (
settings.search_llm_model and settings.search_llm_base_url and settings.search_llm_api_key
):
raise SystemExit(
"诊断 LLM 凭证缺失:.env 需配置 SEARCH_LLM_MODEL / SEARCH_LLM_BASE_URL / "
"SEARCH_LLM_API_KEYP5 fail loud,不静默兜底)"
)
telemetry_db = Path("logs/telemetry.db")
telemetry_db.parent.mkdir(parents=True, exist_ok=True)
telemetry = SQLiteTelemetryRecorder(telemetry_db)
breaker = CircuitBreaker(
fail_threshold=max(settings.llm_circuit_breaker_threshold, 1),
cooldown_s=settings.llm_circuit_breaker_cooldown,
)
llm = GovernedLLMClient(
model=settings.search_llm_model,
base_url=settings.search_llm_base_url,
api_key=settings.search_llm_api_key,
provider=settings.search_llm_model.split("-")[0],
thinking=True,
breaker=breaker,
cache=_build_redis_cache(settings),
telemetry=telemetry,
timeout_s=settings.llm_timeout,
ttft_timeout_s=settings.llm_ttft_timeout,
inter_token_timeout_s=settings.llm_inter_token_timeout,
max_retries=settings.llm_max_retries,
retry_base_delay_s=settings.llm_retry_base_delay,
retry_max_delay_s=settings.llm_retry_max_delay,
)
return DiagnosisDeps(
run_log=RunLogImpl(str(harness_db)),
llm=llm,
skill_store=VersionedSkillStore(_diagnosis_skills_dir()),
prompts=_load_diagnose_prompts(),
tree_data={},
concurrency=concurrency,
)
def _diagnosis_skills_dir() -> Path:
"""诊断用技能目录:种子 store 的当前技能版本(诊断读技能遵从判定)。
诊断只读技能内容判断"是否遵从技能",用 store 种子 v1 即可(与基线 run 一致)。
"""
return Path("store/skills/v1")
def _load_diagnose_prompts() -> Any:
"""加载诊断模板束(从项目根 prompts/ 读取;与 runner._load_diagnose_prompts 一致)。"""
from core.evolution.types import DiagnosePrompts
def _read(name: str) -> str:
p = Path("prompts") / name
return p.read_text(encoding="utf-8") if p.exists() else ""
return DiagnosePrompts(
defect_vs_lapse=_read("defect_vs_lapse.md"),
reasoning_sub=_read("reasoning_sub.md"),
span_eval_system=_read("span_eval_system.md"),
span_eval_user=_read("span_eval_user.md"),
missed_nodes=_read("missed_nodes.md"),
skill_adherence=_read("skill_adherence.md"),
confirmation_bias=_read("confirmation_bias.md"),
evidence_sufficiency=_read("evidence_sufficiency.md"),
)
def _normalize_choice(choice: str | None) -> str:
"""选项归一:strip → 大写 → 取首字母(None 归一为空串)。"""
return (choice or "").strip().upper()[:1]
def load_diagnosable_wrong_ids(harness_db: Path, baseline_run_id: str) -> list[str]:
"""从 harness.db 读 baseline run 的可诊断错题 question_id(保序、canonical 首行)。
可诊断错题判据:canonical 首行(rowid 最小)预测非空 且 stop_reason 非 INFRA
error / parse_error)且 归一后预测 != 答案。INFRA / 空预测题不进 wrong_ids
run_diagnosis 内部也会二次排除,此处前置过滤减少无谓 LLM 调用)。
参数:
harness_db: harness.db 路径(只读打开)。
baseline_run_id: 基线 run 标识。
返回:
可诊断错题 question_id 列表(按 rowid 升序 canonical 顺序,去重)。
异常:
SystemExit: 该 run 无任何预测行(fail loud)。
"""
conn = sqlite3.connect(f"file:{harness_db}?mode=ro", uri=True)
conn.row_factory = sqlite3.Row
try:
rows = conn.execute(
"SELECT question_id, prediction, answer, stop_reason "
"FROM predictions WHERE run_id = ? ORDER BY rowid",
(baseline_run_id,),
).fetchall()
finally:
conn.close()
if not rows:
raise SystemExit(
f"run_id={baseline_run_id}{harness_db} 无任何预测行,无法诊断(P5 fail loud"
)
seen: set[str] = set()
wrong_ids: list[str] = []
for row in rows:
qid = row["question_id"]
if qid in seen:
continue
seen.add(qid)
prediction = (row["prediction"] or "").strip()
if not prediction or row["stop_reason"] in _INFRA_STOP_REASONS:
continue
if _normalize_choice(row["prediction"]) != _normalize_choice(row["answer"]):
wrong_ids.append(qid)
return wrong_ids
def load_questions_by_id(questions_dir: Path) -> dict[str, GeneratedQuestion]:
"""加载 benchmark 全部题并建 question_id → GeneratedQuestion 映射。
覆盖 wrong_ids 与 run_diagnosis 返回的全部 infra/degraded 题(取 video_id/task_type)。
参数:
questions_dir: benchmark 题库目录。
返回:
question_id → GeneratedQuestion 映射。
"""
return {q.question_id: q for q in load_benchmark(questions_dir)}
# ---------------------------------------------------------------------------
# McNemar 功效护栏
# ---------------------------------------------------------------------------
def check_mcnemar_power(pools: Pools, val_wrong_min: int) -> int:
"""校验 validation 池错题数达 McNemar 功效阈,不足即 fail loud。
build_split 按契约用 split_by_video_assignment(val_wrong_min=0)(不破契约),
故功效护栏在 capstone 层单独核验:val 错题数 < 阈 → 验证信号不足以支撑可靠比较。
参数:
pools: 冻结三池(含 validation 与 correctness)。
val_wrong_min: 最少错题数阈(0 = 不检查)。
返回:
validation 池实际错题数(供日志)。
异常:
SystemExit: val_wrong_min > 0 且 val 错题数 < 阈(P5 fail loud,不静默放行)。
"""
val_wrong = sum(1 for q in pools.validation if not pools.correctness[q.question_id])
if val_wrong_min > 0 and val_wrong < val_wrong_min:
raise SystemExit(
f"validation 池错题数 {val_wrong} < val_wrong_min={val_wrong_min}"
"验证信号不足以支撑可靠比较(McNemar 检验功效不够)。"
"请放大 val_ratio / 调整旋钮后重跑,勿静默放行。"
)
return val_wrong
# ---------------------------------------------------------------------------
# 两阶段编排(依赖注入,便于单测)
# ---------------------------------------------------------------------------
async def run_pipeline(
*,
config: VideoSplitConfig,
fingerprint: str,
diagnosis_deps: DiagnosisDeps,
signal_store: DiagnosisSignalStore,
wrong_ids: list[str],
questions: dict[str, GeneratedQuestion],
harness_db: Path,
questions_dir: Path,
out_dir: Path,
generated_at: str,
) -> SplitBuildResult:
"""内联两阶段:Phase 1 诊断 → Phase 2 冻结切分 → McNemar 护栏。
参数:
config: 科研旋钮快照。
fingerprint: 诊断口径指纹(已合成,作诊断信号主键之一)。
diagnosis_deps: Phase 1 诊断依赖束(真实或假实现)。
signal_store: 诊断信号存储端口(Phase 1 写、Phase 2 读)。
wrong_ids: 待诊断的可诊断错题 question_id 列表。
questions: question_id → GeneratedQuestion 映射。
harness_db: harness.db 路径(Phase 2 读 canonical 预测)。
questions_dir: benchmark 题库目录(Phase 2 加载题库切池)。
out_dir: 冻结产物目录(pools.json + split_manifest.json)。
generated_at: 生成时间戳(ISO 字符串,由调用方传入保证可复现)。
返回:
SplitBuildResult(冻结三池 + manifest + assignment)。
"""
# Phase 1: 离线诊断(断点续跑幂等:done_question_ids 已完成题跳过)。
logger.info(
"Phase 1 离线诊断:baseline={} 待诊断错题 {}", config.baseline_run_id, len(wrong_ids)
)
await run_baseline_diagnosis(
baseline_run_id=config.baseline_run_id,
diag_fingerprint=fingerprint,
wrong_ids=wrong_ids,
questions=questions,
store=signal_store,
deps=diagnosis_deps,
)
# Phase 2: 冻结切分(读诊断信号 → 贪心选择 → 视频组原子切三池 → 冻结 + 六条断言)。
out_dir.mkdir(parents=True, exist_ok=True)
logger.info("Phase 2 冻结切分:out={}", out_dir)
result = build_split(
db_path=harness_db,
baseline_run_id=config.baseline_run_id,
signal_store=signal_store,
diag_fingerprint=fingerprint,
questions_dir=questions_dir,
config=SplitBuildConfig(
n_trainval=config.n_trainval,
floor_k=config.floor_k,
epsilon=config.epsilon,
report_floor=config.report_floor,
select_seed=config.seed,
val_ratio=config.val_ratio,
split_seed=config.seed,
),
out_path=out_dir / "pools.json",
manifest_path=out_dir / "split_manifest.json",
generated_at=generated_at,
)
# McNemar 功效护栏(build_split 契约外的 capstone 层校验)。
val_wrong = check_mcnemar_power(result.pools, config.val_wrong_min)
logger.info(
"切分冻结完成:pools={} manifest={} val错题={}/{}(阈)",
out_dir / "pools.json",
out_dir / "split_manifest.json",
val_wrong,
config.val_wrong_min,
)
return result
# ---------------------------------------------------------------------------
# 真实执行 / dry-run 入口
# ---------------------------------------------------------------------------
def _resolve_paths(args: argparse.Namespace) -> tuple[Path, Path, Path]:
"""解析 harness_db / questions_dir / out_dirCLI 覆盖默认工程路径)。"""
harness_db = args.harness_db or _DEFAULT_HARNESS_DB
questions_dir = args.questions_dir or _DEFAULT_QUESTIONS_DIR
out_dir = args.out_dir or _DEFAULT_OUT_DIR
return harness_db, questions_dir, out_dir
def _execute_real(config: VideoSplitConfig, fingerprint: str, args: argparse.Namespace) -> None:
"""真实执行两阶段管线:组装真实 deps、读错题、跑诊断 + 冻结切分。"""
harness_db, questions_dir, out_dir = _resolve_paths(args)
if not harness_db.exists():
raise SystemExit(f"harness.db 不存在: {harness_db}P5 fail loud")
wrong_ids = load_diagnosable_wrong_ids(harness_db, config.baseline_run_id)
questions = load_questions_by_id(questions_dir)
deps = build_diagnosis_deps(harness_db=harness_db, concurrency=args.concurrency)
from adapters.baseline_diagnosis_store import SqliteDiagnosisSignalStore
store = SqliteDiagnosisSignalStore(str(harness_db))
try:
asyncio.run(
run_pipeline(
config=config,
fingerprint=fingerprint,
diagnosis_deps=deps,
signal_store=store,
wrong_ids=wrong_ids,
questions=questions,
harness_db=harness_db,
questions_dir=questions_dir,
out_dir=out_dir,
generated_at=datetime.datetime.now(datetime.UTC).isoformat(),
)
)
finally:
store.close()
class _DryRunLLM:
"""dry-run 假 LLM:被真实调用即报错,保证不真调 LLM。"""
async def complete(self, *args: Any, **kwargs: Any) -> Any:
raise AssertionError("dry-run 不应真调 LLM.complete")
class _DryRunLog:
"""dry-run 假 RunLogpredictions/traces 均返回空,诊断不真正执行。"""
async def get_predictions(self, run_id: str, *, question_ids: list[str] | None = None) -> list:
return []
async def get_traces(self, run_id: str, *, question_ids: list[str] | None = None) -> list:
return []
def _execute_dry_run(config: VideoSplitConfig, fingerprint: str, args: argparse.Namespace) -> None:
"""dry-run:用假 deps 跑通 Phase 1 wiring(空错题 → 诊断早返回),打印步骤 + 指纹。
Phase 2 build_split 需真实诊断信号方能冻结,dry-run 不真实冻结,仅打印其计划;
Phase 1 用空 wrong_ids 走 run_baseline_diagnosis 早返回路径,验证装配可调用而不触 LLM。
"""
harness_db, questions_dir, out_dir = _resolve_paths(args)
logger.info("=== dry-run:校验两阶段装配(不真调 LLM / 不冻结产物)===")
logger.info(
"诊断口径指纹 diag_fingerprint={} (prompt={} model={})",
fingerprint,
config.prompt_version,
config.model,
)
logger.info(
"解析路径:harness_db={} questions_dir={} out_dir={}", harness_db, questions_dir, out_dir
)
logger.info(
"旋钮:n_trainval={} epsilon={} report_floor={} val_ratio={} seed={} "
"val_wrong_min={} floor_k={}",
config.n_trainval,
config.epsilon,
config.report_floor,
config.val_ratio,
config.seed,
config.val_wrong_min,
config.floor_k,
)
fake_deps = DiagnosisDeps(
run_log=_DryRunLog(),
llm=_DryRunLLM(),
skill_store=object(),
prompts=object(),
tree_data={},
concurrency=args.concurrency,
)
from adapters.baseline_diagnosis_store import SqliteDiagnosisSignalStore
dry_db = out_dir / "_dry_run_signals.db"
dry_db.parent.mkdir(parents=True, exist_ok=True)
store = SqliteDiagnosisSignalStore(str(dry_db))
try:
logger.info("Phase 1 装配 OKrun_baseline_diagnosis 以空错题走早返回路径(不触 LLM)")
asyncio.run(
run_baseline_diagnosis(
baseline_run_id=config.baseline_run_id,
diag_fingerprint=fingerprint,
wrong_ids=[],
questions={},
store=store,
deps=fake_deps,
)
)
finally:
store.close()
dry_db.unlink(missing_ok=True)
logger.info(
"Phase 2 装配 OK:真实执行将调 build_split 冻结 pools.json + manifestdry-run 跳过)"
)
logger.info("=== dry-run 通过:两阶段装配可调用,指纹已算出 ===")
def build_arg_parser() -> argparse.ArgumentParser:
"""构建 CLI 参数解析器。"""
parser = argparse.ArgumentParser(description="结果驱动视频级切分两阶段 CLI(诊断 → 冻结切分)")
parser.add_argument("--config", type=Path, default=Path("config/video_split.yaml"))
parser.add_argument("--dry-run", action="store_true", dest="dry_run")
parser.add_argument("--gpu", type=str, default=None, help="可选:设置 CUDA_VISIBLE_DEVICES")
parser.add_argument("--concurrency", type=int, default=8, help="诊断并发上限")
parser.add_argument("--harness-db", type=Path, default=None, dest="harness_db")
parser.add_argument("--questions-dir", type=Path, default=None, dest="questions_dir")
parser.add_argument("--out-dir", type=Path, default=None, dest="out_dir")
return parser
def main(argv: list[str] | None = None) -> None:
"""CLI 入口:解析参数 → 载配置 → 算指纹 → dry-run 或真实两阶段执行。
参数:
argv: 可选参数列表(默认 sys.argv[1:]),便于测试注入。
"""
from dotenv import load_dotenv
load_dotenv()
args = build_arg_parser().parse_args(argv)
if args.gpu is not None:
os.environ["CUDA_VISIBLE_DEVICES"] = args.gpu
logger.info("CUDA_VISIBLE_DEVICES={}", args.gpu)
config = load_config(args.config)
fingerprint = diag_fingerprint(config.prompt_version, config.model, git_short_sha())
if args.dry_run:
_execute_dry_run(config, fingerprint, args)
return
_execute_real(config, fingerprint, args)
if __name__ == "__main__":
main()