Files

338 lines
11 KiB
Python

"""CLI 入口 — Composition Root:构建适配器,注入 Runner,调度执行。
三层配置合并(YAML > .env > CLI)由 load_config 完成。
适配器参数通过 InfraSettings(BaseSettings) 从 .env 加载。
"""
from __future__ import annotations
import argparse
import asyncio
from pathlib import Path
from typing import NamedTuple
from dotenv import load_dotenv
from loguru import logger
from pydantic_settings import BaseSettings, SettingsConfigDict
class InfraSettings(BaseSettings):
"""工程配置(少变/敏感),从 .env 加载。"""
model_config = SettingsConfigDict(env_file=".env", extra="ignore")
search_llm_model: str = ""
search_llm_base_url: str = ""
search_llm_api_key: str = ""
vl_llm_model: str = ""
vl_llm_base_url: str = ""
vl_llm_api_key: str = ""
evolve_llm_model: str = ""
evolve_llm_base_url: str = ""
evolve_llm_api_key: str = ""
embed_api_key: str = ""
embed_api_url: str = ""
monkey_ocr_urls: 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
class _Adapters(NamedTuple):
"""全套适配器实例。"""
llm: object
evolve_llm: object
vlm: object
telemetry: object
embed: object
ocr: object
def _build_adapters(settings: InfraSettings, embed_cfg: dict) -> _Adapters:
"""从 InfraSettings 构建全套适配器。
参数:
settings: 工程配置。
embed_cfg: 嵌入配置字典(来自 YAML embed 段)。
返回:
_Adapters 命名元组。
"""
from adapters.breaker import CircuitBreaker
from adapters.embedding import LocalEmbeddingProvider
from adapters.llm import GovernedLLMClient
from adapters.telemetry import SQLiteTelemetryRecorder
from adapters.vlm import GovernedVLMClient
breaker = CircuitBreaker(
fail_threshold=max(settings.llm_circuit_breaker_threshold, 1),
cooldown_s=settings.llm_circuit_breaker_cooldown,
)
cache = None
if settings.redis_url:
from adapters.redis_cache import RedisResponseCache, _resolve_cache_ttl
# 配置校验 fail-loud(不属于 Redis 连接故障,不得被下方降级 except 吞掉)
ttl_s = _resolve_cache_ttl(settings.redis_cache_ttl)
try:
import redis.asyncio as aioredis
redis_client = aioredis.from_url(settings.redis_url, decode_responses=True)
cache = RedisResponseCache(redis=redis_client, ttl_s=ttl_s)
except Exception:
logger.warning("Redis 缓存不可用,降级为无缓存模式")
telemetry_db = Path("logs/telemetry.db")
telemetry_db.parent.mkdir(parents=True, exist_ok=True)
telemetry = SQLiteTelemetryRecorder(telemetry_db)
def _make_llm(model: str, base_url: str, api_key: str, *, thinking: bool) -> GovernedLLMClient:
"""构建单个 GovernedLLMClient 实例。"""
return GovernedLLMClient(
model=model,
base_url=base_url,
api_key=api_key,
provider=model.split("-")[0] if model else "unknown",
thinking=thinking,
breaker=breaker,
cache=cache,
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,
)
llm = _make_llm(
settings.search_llm_model,
settings.search_llm_base_url,
settings.search_llm_api_key,
thinking=True,
)
evolve_llm = llm
vl_llm = _make_llm(
settings.vl_llm_model,
settings.vl_llm_base_url,
settings.vl_llm_api_key,
thinking=False,
)
vlm = GovernedVLMClient(governed_llm=vl_llm)
backend = embed_cfg.get("backend", "local")
if backend == "local":
embed = LocalEmbeddingProvider(
model_name=embed_cfg.get("model_name", "BAAI/bge-base-zh-v1.5"),
embed_dim=embed_cfg.get("embed_dim", 768),
device=embed_cfg.get("device", "cpu"),
)
else:
from adapters.embedding import RemoteEmbeddingProvider
embed = RemoteEmbeddingProvider(
model_name=embed_cfg.get("model_name", ""),
embed_dim=embed_cfg.get("embed_dim", 768),
api_key=settings.embed_api_key,
api_url=settings.embed_api_url,
)
ocr = None
if settings.monkey_ocr_urls:
from adapters.ocr import MonkeyOCRClient
urls = [u.strip() for u in settings.monkey_ocr_urls.split(",") if u.strip()]
if urls:
ocr = MonkeyOCRClient(urls=urls)
return _Adapters(
llm=llm,
evolve_llm=evolve_llm,
vlm=vlm,
telemetry=telemetry,
embed=embed,
ocr=ocr,
)
def _build_parser() -> argparse.ArgumentParser:
"""构建 CLI 参数解析器。所有参数 default=None,未传入时使用 YAML 默认值。"""
parser = argparse.ArgumentParser(description="Video-Tree-TRM5 实验运行器")
parser.add_argument("--config", type=Path, default=Path("config/default.yaml"))
parser.add_argument("--workspace-dir", type=Path, dest="workspace_dir")
parser.add_argument("--store-dir", type=Path, dest="store_dir")
parser.add_argument(
"--mode",
choices=["infer", "train", "diagnose", "evolve", "eval", "promote"],
)
parser.add_argument("--run-id", type=str, dest="run_id")
parser.add_argument("--concurrency", type=int)
parser.add_argument("--max-steps", type=int, dest="max_steps")
parser.add_argument(
"--skill-mode",
choices=["auto", "manual", "none"],
dest="skill_mode",
)
parser.add_argument("--n-samples", type=int, dest="n_samples")
parser.add_argument("--questions", type=str)
parser.add_argument("--skills-version", type=str, dest="skills_version")
parser.add_argument("--prompts-version", type=str, dest="prompts_version")
parser.add_argument("--task-types", nargs="+", dest="task_types")
parser.add_argument("--resume", action="store_true", dest="resume")
parser.add_argument("--fresh", action="store_true", dest="fresh")
parser.add_argument("--seed", type=str, dest="seed")
parser.add_argument("--epochs", type=int)
parser.add_argument(
"--pool-split-mode",
choices=["global", "per_category"],
dest="pool_split_mode",
)
parser.add_argument("--train-ratio", type=float, dest="train_ratio")
parser.add_argument("--test-questions", type=str, dest="test_questions")
parser.add_argument(
"--no-run-holdout-eval",
action="store_true",
dest="no_run_holdout_eval",
)
return parser
def _log_result(result: object) -> None:
"""输出推理结果摘要。"""
logger.info("=" * 60)
logger.info("运行 ID: {}", result.run_id)
logger.info(
"总体准确率: {:.2%} ({}/{})",
result.accuracy,
result.correct,
result.total,
)
logger.info("平均步数: {:.1f}", result.steps_mean)
logger.info(
"Token 用量: prompt={}, completion={}",
result.token_usage["prompt_tokens"],
result.token_usage["completion_tokens"],
)
if result.per_task_type:
logger.info("--- 按任务类型 ---")
for task_type, stats in sorted(result.per_task_type.items()):
logger.info(
" {}: {:.2%} ({}/{})",
task_type,
stats["accuracy"],
stats["correct"],
stats["total"],
)
logger.info(
"停止原因: {}",
", ".join(f"{k}={v}" for k, v in result.stop_reason_counts.items()),
)
logger.info("=" * 60)
def main() -> None:
"""入口函数。"""
load_dotenv()
log_dir = Path("logs")
log_dir.mkdir(parents=True, exist_ok=True)
logger.add(
log_dir / "run_{time:YYYYMMDD_HHmmss}.log",
rotation="500 MB",
retention="30 days",
encoding="utf-8",
enqueue=False,
)
parser = _build_parser()
args = parser.parse_args()
import yaml
with open(args.config, encoding="utf-8") as f:
raw_yaml = yaml.safe_load(f)
from app.harness.config import load_config
cli_args = vars(args)
if cli_args.get("task_types") is not None:
cli_args["task_types"] = tuple(cli_args["task_types"])
cli_overrides = {k: v for k, v in cli_args.items() if k != "config"}
if cli_overrides.get("no_run_holdout_eval"):
cli_overrides["run_holdout_eval"] = False
cli_overrides.pop("no_run_holdout_eval", None)
config = load_config(args.config, cli_overrides)
logger.info("配置加载完成: mode={}, workspace={}", config.mode, config.workspace_dir)
settings = InfraSettings()
embed_cfg = raw_yaml.get("embed", {})
adapters = _build_adapters(settings, embed_cfg)
from app.harness.deps_router import InferenceDepsRouter
from app.harness.runner import Runner
router = InferenceDepsRouter(
store_dir=Path(config.store_dir),
embed_provider=adapters.embed,
llm=adapters.llm,
vlm=adapters.vlm,
ocr=adapters.ocr,
default_prompts_dir=Path(config.store_dir) / "prompts" / config.prompts_version,
default_skills_dir=Path(config.store_dir) / "skills" / config.skills_version,
skill_mode=config.skill_mode,
verify_vision=True,
anchor=True,
assemble_mode="ids_expand",
)
runner = Runner(
config,
llm=adapters.llm,
evolve_llm=adapters.evolve_llm,
vlm=adapters.vlm,
telemetry=adapters.telemetry,
tool_dispatch_factory=router.create_dispatch,
prompt_builder_factory=router.create_prompt_builder,
)
if config.mode == "infer":
result = asyncio.run(runner.infer(task_types=config.task_types))
_log_result(result)
elif config.mode == "train":
from app.harness.pools import (
GlobalPoolStrategy,
PerCategoryPoolStrategy,
build_or_load_pools,
)
from app.harness.workspace import resolve_paths
strategy = (
PerCategoryPoolStrategy()
if config.pool_split_mode == "per_category"
else GlobalPoolStrategy()
)
paths = resolve_paths(config.workspace_dir)
pools = build_or_load_pools(config, strategy, paths.db_path)
asyncio.run(runner.train(pools))
else:
raise SystemExit(f"模式 {config.mode!r} 尚未实现")
if __name__ == "__main__":
main()