cabe5be038
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
289 lines
9.1 KiB
Python
289 lines
9.1 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 = 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="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":
|
|
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()
|