"""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: 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 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") 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"} 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()