feat: main.py Composition Root(仅 infer 模式)
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user