diff --git a/main.py b/main.py new file mode 100644 index 0000000..3d2d3c1 --- /dev/null +++ b/main.py @@ -0,0 +1,288 @@ +"""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 = 120.0 + llm_max_retries: int = 3 + llm_retry_base_delay: float = 2.0 + llm_retry_max_delay: float = 30.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: + from adapters.redis_cache import RedisResponseCache + + cache = RedisResponseCache(redis_url=settings.redis_url, ttl=settings.redis_cache_ttl) + 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) + 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() + 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_overrides = {k: v for k, v in vars(args).items() if k not in ("config", "task_types")} + 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="default", + ) + + 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": + task_types = getattr(args, "task_types", None) + result = asyncio.run(runner.infer(task_types=task_types)) + _log_result(result) + else: + raise SystemExit(f"模式 {config.mode!r} 尚未实现") + + +if __name__ == "__main__": + main()