160fb3bc7c
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1317 lines
42 KiB
Markdown
1317 lines
42 KiB
Markdown
# main.py 推理入口 + 初始 Prompt 集 实现计划
|
||
|
||
> **For agentic workers:** REQUIRED SUB-SKILL: Use subagent-driven-development to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||
|
||
**Goal:** 实现 `main.py` CLI 入口(infer 模式),准备初始 prompt/skill 集,完成 900 道题推理基线。
|
||
|
||
**Architecture:** Clean Architecture Composition Root 模式。main.py 构建全套 adapters 和 InferenceDepsRouter,注入 Runner。Router 按 video_id 懒加载 per-video InferenceDeps,在 prompt_builder 调用时注册 question_id→video_id 映射,dispatch 通过 context["session_id"] 路由。
|
||
|
||
**Tech Stack:** Python 3.11, asyncio, argparse, pydantic-settings, loguru
|
||
|
||
**设计文档:** `research-wiki/designs/2026-07-09-main-inference-entry-design.md`
|
||
|
||
---
|
||
|
||
## File Structure
|
||
|
||
| 操作 | 文件 | 职责 |
|
||
|------|------|------|
|
||
| Create | `app/harness/deps_router.py` | 按 video_id 懒加载 InferenceDeps 并路由 dispatch/prompt_builder |
|
||
| Modify | `app/ports.py:76` | 新增 4 个 Protocol(ToolDispatchFn, ToolDispatchFactory, PromptBuilderFn, PromptBuilderFactory) |
|
||
| Modify | `app/harness/runner.py:454-469` | __init__ 增加 2 个 factory 参数 + fail-fast 校验 |
|
||
| Modify | `app/harness/runner.py:2064-2082` | _make_* 方法优先用注入值 |
|
||
| Create | `main.py` | Composition Root:argparse + InfraSettings + 适配器构建 + Runner 组装 |
|
||
| Move | `store/prompts/*.md` → `store/prompts/v1/` | 版本化目录重组 |
|
||
| Create | `store/skills/v1/` (13 files) | 从 TRM4 v1 精简 + 注入 TRM5 card 字段 |
|
||
| Modify | `config/default.yaml:29,31` | concurrency=24, max_steps=40 |
|
||
| Test | `tests/unit/test_deps_router.py` | Router 单元测试 |
|
||
| Test | `tests/unit/test_ports_factory.py` | Protocol 结构测试 |
|
||
|
||
---
|
||
|
||
### Task 1: store/ 目录重组
|
||
|
||
**Files:**
|
||
- Move: `store/prompts/*.md` → `store/prompts/v1/*.md`
|
||
- Create: `store/skills/v1/` (empty, 后续 Task 填充)
|
||
|
||
- [ ] **Step 1: 创建 v1 子目录并移动 prompt 文件**
|
||
|
||
```bash
|
||
mkdir -p store/prompts/v1
|
||
git mv store/prompts/system.md store/prompts/v1/
|
||
git mv store/prompts/observe_frame_extract.md store/prompts/v1/
|
||
git mv store/prompts/observe_frame_verify.md store/prompts/v1/
|
||
git mv store/prompts/search_similar_extract.md store/prompts/v1/
|
||
git mv store/prompts/search_similar_verify.md store/prompts/v1/
|
||
git mv store/prompts/view_node_extract.md store/prompts/v1/
|
||
git mv store/prompts/view_node_verify.md store/prompts/v1/
|
||
git mv store/prompts/view_node_children_extract.md store/prompts/v1/
|
||
git mv store/prompts/view_node_children_verify.md store/prompts/v1/
|
||
```
|
||
|
||
- [ ] **Step 2: 创建 skills/v1 目录**
|
||
|
||
```bash
|
||
mkdir -p store/skills/v1
|
||
```
|
||
|
||
- [ ] **Step 3: 验证目录结构**
|
||
|
||
Run: `ls store/prompts/v1/ && ls store/skills/v1/`
|
||
Expected: 9 个 .md 文件在 prompts/v1/ 下,skills/v1/ 为空目录。
|
||
|
||
- [ ] **Step 4: Commit**
|
||
|
||
```bash
|
||
git add store/prompts/ store/skills/
|
||
git commit -m "refactor(store): prompts 版本化目录重组 + skills/v1 骨架"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 2: store/skills/v1 — 13 个精简 + 注入 skill
|
||
|
||
**Files:**
|
||
- Create: `store/skills/v1/default-strategy.md`
|
||
- Create: `store/skills/v1/action-reasoning.md`
|
||
- Create: `store/skills/v1/action-recognition.md`
|
||
- Create: `store/skills/v1/attribute-perception.md`
|
||
- Create: `store/skills/v1/counting-problem.md`
|
||
- Create: `store/skills/v1/information-synopsis.md`
|
||
- Create: `store/skills/v1/object-reasoning.md`
|
||
- Create: `store/skills/v1/object-recognition.md`
|
||
- Create: `store/skills/v1/ocr-problems.md`
|
||
- Create: `store/skills/v1/spatial-perception.md`
|
||
- Create: `store/skills/v1/spatial-reasoning.md`
|
||
- Create: `store/skills/v1/temporal-perception.md`
|
||
- Create: `store/skills/v1/temporal-reasoning.md`
|
||
|
||
**变换规则**(每个 skill 文件统一适用):
|
||
|
||
| 操作 | 对象 | 说明 |
|
||
|------|------|------|
|
||
| 保留 | YAML frontmatter | name, description, task_type 原样保留 |
|
||
| 保留 | `## 适用场景` 节 | 原样保留 |
|
||
| 保留 | `## 搜索步骤` 的 Step 标题 | 如 `### Step 1: 事件定位` |
|
||
| 精简 | Step 正文 | 保留第一句话意图描述,移除数据驱动统计(如 "75% 正确率")、精确转换条件、详细操作指令 |
|
||
| 保留 | `## 输出格式` 节 | JSON schema (reflect/plan/action) 原样保留 |
|
||
| 精简 | `## 自检信号` 节 | 最多保留 1 条最通用的自检信号 |
|
||
| 精简 | `## 常见陷阱` 节 | 最多保留 2 条最通用的陷阱警告,移除特定失败模式 |
|
||
| **新增** | `## 视频树字段索引` 节 | 插入 card 字段索引表(见下方) |
|
||
|
||
**card 字段索引表**(全部 13 个 skill 共用,插入在 `## 搜索步骤` 之前):
|
||
|
||
```markdown
|
||
## 视频树字段索引
|
||
|
||
| 层级 | 字段 | 适用场景 |
|
||
|------|------|---------|
|
||
| L1 | scene_summary | 整体概况 |
|
||
| L1 | key_entities | 查找人物/物体 |
|
||
| L1 | main_actions | 主要动作 |
|
||
| L1 | temporal_flow | 时间线概览 |
|
||
| L1 | topic_keywords | 主题定位 |
|
||
| L2 | event_description | 事件因果 |
|
||
| L2 | entities / actions | 实体和动作细节 |
|
||
| L2 | state_changes | 状态转变 |
|
||
| L2 | spatial_relations | 空间关系变化 |
|
||
| L3 | frame_summary | 精确视觉证据 |
|
||
| L3 | visible_entities | 具体物体确认 |
|
||
| L3 | ongoing_actions | 正在发生的动作 |
|
||
| L3 | spatial_layout | 精确空间位置 |
|
||
| L3 | visual_attributes | 光照、色调、机位 |
|
||
| 全层 | visible_text | 画面文字(OCR) |
|
||
| 全层 | subtitle | 字幕转写 |
|
||
```
|
||
|
||
**TRM4 v1 源文件路径:** `/home/iomgaa/Projects/Video-Tree-TRM4/store/skills/v1/`
|
||
|
||
- [ ] **Step 1: 逐个读取 TRM4 v1 skill,按变换规则精简 + 注入,写入 store/skills/v1/**
|
||
|
||
对每个 skill 文件执行:
|
||
1. 读取 `/home/iomgaa/Projects/Video-Tree-TRM4/store/skills/v1/<name>.md`
|
||
2. 应用上述变换规则
|
||
3. 写入 `store/skills/v1/<name>.md`
|
||
|
||
以 `default-strategy.md` 为首个示例(最重要的通用策略)。
|
||
|
||
- [ ] **Step 2: 验证 13 个文件完整性**
|
||
|
||
Run: `ls store/skills/v1/ | wc -l`
|
||
Expected: 13
|
||
|
||
Run: `head -5 store/skills/v1/default-strategy.md`
|
||
Expected: YAML frontmatter with `task_type: _default`
|
||
|
||
Run: `grep "视频树字段索引" store/skills/v1/*.md | wc -l`
|
||
Expected: 13(每个文件都有字段索引表)
|
||
|
||
- [ ] **Step 3: Commit**
|
||
|
||
```bash
|
||
git add store/skills/v1/
|
||
git commit -m "feat(store): skills/v1 初始集 — TRM4 精简 + TRM5 card 字段注入"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 3: app/ports.py — 新增 4 个 Protocol
|
||
|
||
**Files:**
|
||
- Modify: `app/ports.py:76` (文件末尾追加)
|
||
- Test: `tests/unit/test_ports_factory.py`
|
||
|
||
- [ ] **Step 1: 编写 Protocol 结构测试**
|
||
|
||
```python
|
||
# tests/unit/test_ports_factory.py
|
||
"""ToolDispatchFactory / PromptBuilderFactory Protocol 结构验证。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
import pytest
|
||
|
||
from app.ports import (
|
||
PromptBuilderFactory,
|
||
PromptBuilderFn,
|
||
ToolDispatchFactory,
|
||
ToolDispatchFn,
|
||
)
|
||
|
||
|
||
class TestToolDispatchFnProtocol:
|
||
"""ToolDispatchFn 签名检查。"""
|
||
|
||
def test_conforming_callable_passes_isinstance(self) -> None:
|
||
async def dispatch(
|
||
tool_name: str, args: dict[str, Any], *, context: dict[str, Any]
|
||
) -> str:
|
||
return ""
|
||
|
||
assert isinstance(dispatch, ToolDispatchFn)
|
||
|
||
def test_wrong_return_type_noted(self) -> None:
|
||
"""仅验证签名存在;runtime_checkable 不检查返回类型。"""
|
||
async def bad(
|
||
tool_name: str, args: dict[str, Any], *, context: dict[str, Any]
|
||
) -> int:
|
||
return 0
|
||
|
||
# runtime_checkable 仅检查方法存在,不验证类型注解
|
||
assert isinstance(bad, ToolDispatchFn)
|
||
|
||
|
||
class TestToolDispatchFactoryProtocol:
|
||
"""ToolDispatchFactory 签名检查。"""
|
||
|
||
def test_conforming_class_passes(self) -> None:
|
||
class Factory:
|
||
def __call__(self, *, skills_dir: Path | None = None) -> Any:
|
||
return None
|
||
|
||
assert isinstance(Factory(), ToolDispatchFactory)
|
||
|
||
|
||
class TestPromptBuilderFnProtocol:
|
||
"""PromptBuilderFn 签名检查。"""
|
||
|
||
def test_conforming_callable_passes(self) -> None:
|
||
from core.types import GeneratedQuestion
|
||
|
||
def builder(qa: GeneratedQuestion) -> tuple[str, str]:
|
||
return ("", "")
|
||
|
||
assert isinstance(builder, PromptBuilderFn)
|
||
|
||
|
||
class TestPromptBuilderFactoryProtocol:
|
||
"""PromptBuilderFactory 签名检查。"""
|
||
|
||
def test_conforming_class_passes(self) -> None:
|
||
class Factory:
|
||
def __call__(
|
||
self, *, skills_dir: Path | None = None, prompts_dir: Path | None = None
|
||
) -> Any:
|
||
return None
|
||
|
||
assert isinstance(Factory(), PromptBuilderFactory)
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试确认失败**
|
||
|
||
Run: `conda run -n Video-Tree-TRM pytest tests/unit/test_ports_factory.py -v`
|
||
Expected: ImportError(ToolDispatchFn 等尚未定义)
|
||
|
||
- [ ] **Step 3: 在 app/ports.py 末尾追加 4 个 Protocol**
|
||
|
||
在 `app/ports.py` 文件末尾(第 77 行之后)追加:
|
||
|
||
```python
|
||
|
||
|
||
@runtime_checkable
|
||
class ToolDispatchFn(Protocol):
|
||
"""工具调度函数签名。
|
||
|
||
参数:
|
||
tool_name: 工具名称。
|
||
args: 工具参数字典。
|
||
context: 上下文字典(包含 session_id)。
|
||
|
||
返回:
|
||
工具执行结果文本。
|
||
"""
|
||
|
||
async def __call__(
|
||
self, tool_name: str, args: dict[str, Any], *, context: dict[str, Any]
|
||
) -> str: ...
|
||
|
||
|
||
@runtime_checkable
|
||
class ToolDispatchFactory(Protocol):
|
||
"""per-version 工具调度工厂。
|
||
|
||
根据可选的 skills_dir 覆盖构建工具调度函数。
|
||
|
||
参数:
|
||
skills_dir: 可选的 skills 版本目录覆盖。
|
||
|
||
返回:
|
||
ToolDispatchFn 实例。
|
||
"""
|
||
|
||
def __call__(self, *, skills_dir: Path | None = None) -> ToolDispatchFn: ...
|
||
|
||
|
||
@runtime_checkable
|
||
class PromptBuilderFn(Protocol):
|
||
"""Prompt 构建函数签名。
|
||
|
||
参数:
|
||
qa: 生成的题目实例。
|
||
|
||
返回:
|
||
(system_prompt, user_prompt) 二元组。
|
||
"""
|
||
|
||
def __call__(self, qa: GeneratedQuestion) -> tuple[str, str]: ...
|
||
|
||
|
||
@runtime_checkable
|
||
class PromptBuilderFactory(Protocol):
|
||
"""per-version prompt 构建工厂。
|
||
|
||
参数:
|
||
skills_dir: 可选的 skills 版本目录覆盖。
|
||
prompts_dir: 可选的 prompts 版本目录覆盖。
|
||
|
||
返回:
|
||
PromptBuilderFn 实例。
|
||
"""
|
||
|
||
def __call__(
|
||
self,
|
||
*,
|
||
skills_dir: Path | None = None,
|
||
prompts_dir: Path | None = None,
|
||
) -> PromptBuilderFn: ...
|
||
```
|
||
|
||
同时在文件顶部 `if TYPE_CHECKING:` 块中确保 `Any` 已导入(已有 `from typing import ... Protocol`,需追加 `Any`)。
|
||
|
||
- [ ] **Step 4: 运行测试确认通过**
|
||
|
||
Run: `conda run -n Video-Tree-TRM pytest tests/unit/test_ports_factory.py -v`
|
||
Expected: 5 tests PASSED
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
git add app/ports.py tests/unit/test_ports_factory.py
|
||
git commit -m "feat(ports): 新增 ToolDispatchFactory/PromptBuilderFactory Protocol"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 4: app/harness/runner.py — factory 注入
|
||
|
||
**Files:**
|
||
- Modify: `app/harness/runner.py:440-469` (__init__)
|
||
- Modify: `app/harness/runner.py:2064-2082` (_make_* 方法)
|
||
- Test: `tests/unit/test_harness_runner.py` (追加)
|
||
|
||
- [ ] **Step 1: 编写 factory 注入测试**
|
||
|
||
在 `tests/unit/test_harness_runner.py` 末尾追加:
|
||
|
||
```python
|
||
class TestRunnerFactoryInjection:
|
||
"""Runner factory 注入 fail-fast 校验。"""
|
||
|
||
def test_infer_mode_missing_factory_raises(self, tmp_path: Path) -> None:
|
||
"""mode=infer 时缺少 factory 参数 → 立即 ValueError。"""
|
||
from unittest.mock import AsyncMock
|
||
from app.harness.config import RunConfig
|
||
|
||
config = RunConfig(
|
||
workspace_dir=tmp_path,
|
||
store_dir=tmp_path,
|
||
mode="infer",
|
||
concurrency=1,
|
||
max_steps=5,
|
||
skill_mode="none",
|
||
n_samples=0,
|
||
questions="benchmarks/Video-MME",
|
||
skills_version="v1",
|
||
prompts_version="v1",
|
||
epochs=1,
|
||
diag_size=10,
|
||
diag_correct_ratio=0.5,
|
||
val_size=24,
|
||
val_correct_ratio=0.5,
|
||
edit_budget_start=5,
|
||
edit_budget_end=2,
|
||
batch_size=5,
|
||
min_class_per_batch=2,
|
||
eval_min_per_class=2,
|
||
early_stop_patience=3,
|
||
test_size=10,
|
||
use_slow_momentum=False,
|
||
gate_e_confirm=20.0,
|
||
gate_e_provisional=3.0,
|
||
gate_w_net_min=2,
|
||
gate_delta_min=0.02,
|
||
gate_lambda_dir=-0.642,
|
||
gate_e_rollback=10.0,
|
||
gate_block=8,
|
||
gate_n_max=40,
|
||
gate_p_low=0.05,
|
||
gate_p_high=0.95,
|
||
gate_probe_quota=0.2,
|
||
gate_gamma_decay=0.9,
|
||
gate_cooldown_steps=2,
|
||
gate_guard_err=0.10,
|
||
skill_update_mode="patch",
|
||
appendix_consolidate_threshold=6,
|
||
)
|
||
|
||
with pytest.raises(ValueError, match="tool_dispatch_factory"):
|
||
Runner(
|
||
config,
|
||
llm=AsyncMock(),
|
||
evolve_llm=AsyncMock(),
|
||
vlm=AsyncMock(),
|
||
telemetry=AsyncMock(),
|
||
)
|
||
|
||
def test_diagnose_mode_allows_none_factory(self, tmp_path: Path) -> None:
|
||
"""mode=diagnose 不需要 factory(不走推理路径)→ 不报错。"""
|
||
from unittest.mock import AsyncMock
|
||
from app.harness.config import RunConfig
|
||
|
||
ws = tmp_path / "ws"
|
||
ws.mkdir()
|
||
(ws / "manifest.json").write_text('{"name":"ws","created_at":"","store":"../store","current":{"videos":"v","questions":"q","skills":"s","prompts":"p"},"history":[]}')
|
||
|
||
config = RunConfig(
|
||
workspace_dir=ws,
|
||
store_dir=tmp_path,
|
||
mode="diagnose",
|
||
run_id="test_run",
|
||
concurrency=1,
|
||
max_steps=5,
|
||
skill_mode="none",
|
||
n_samples=0,
|
||
questions="benchmarks/Video-MME",
|
||
skills_version="v1",
|
||
prompts_version="v1",
|
||
epochs=1,
|
||
diag_size=10,
|
||
diag_correct_ratio=0.5,
|
||
val_size=24,
|
||
val_correct_ratio=0.5,
|
||
edit_budget_start=5,
|
||
edit_budget_end=2,
|
||
batch_size=5,
|
||
min_class_per_batch=2,
|
||
eval_min_per_class=2,
|
||
early_stop_patience=3,
|
||
test_size=10,
|
||
use_slow_momentum=False,
|
||
gate_e_confirm=20.0,
|
||
gate_e_provisional=3.0,
|
||
gate_w_net_min=2,
|
||
gate_delta_min=0.02,
|
||
gate_lambda_dir=-0.642,
|
||
gate_e_rollback=10.0,
|
||
gate_block=8,
|
||
gate_n_max=40,
|
||
gate_p_low=0.05,
|
||
gate_p_high=0.95,
|
||
gate_probe_quota=0.2,
|
||
gate_gamma_decay=0.9,
|
||
gate_cooldown_steps=2,
|
||
gate_guard_err=0.10,
|
||
skill_update_mode="patch",
|
||
appendix_consolidate_threshold=6,
|
||
)
|
||
|
||
# 不应抛异常
|
||
runner = Runner(
|
||
config,
|
||
llm=AsyncMock(),
|
||
evolve_llm=AsyncMock(),
|
||
vlm=AsyncMock(),
|
||
telemetry=AsyncMock(),
|
||
)
|
||
assert runner is not None
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试确认失败**
|
||
|
||
Run: `conda run -n Video-Tree-TRM pytest tests/unit/test_harness_runner.py::TestRunnerFactoryInjection -v`
|
||
Expected: FAIL(Runner.__init__ 不接受 factory 参数 / 不做校验)
|
||
|
||
- [ ] **Step 3: 修改 Runner.__init__ 和 _make_* 方法**
|
||
|
||
在 `app/harness/runner.py` 中:
|
||
|
||
**3a.** 修改 `Runner.__init__`(第 454-469 行),增加 2 个可选参数 + fail-fast 校验:
|
||
|
||
```python
|
||
def __init__(
|
||
self,
|
||
config: RunConfig,
|
||
*,
|
||
llm: LLMProvider,
|
||
evolve_llm: LLMProvider,
|
||
vlm: VLMProvider,
|
||
telemetry: TelemetryRecorder,
|
||
tool_dispatch_factory: Any | None = None,
|
||
prompt_builder_factory: Any | None = None,
|
||
) -> None:
|
||
# fail-fast 校验必须在 _ensure_workspace 之前,避免 workspace 报错掩盖 factory 缺失
|
||
if config.mode in {"infer", "eval", "train"}:
|
||
if tool_dispatch_factory is None or prompt_builder_factory is None:
|
||
raise ValueError(
|
||
f"mode={config.mode!r} 需要 tool_dispatch_factory 和 "
|
||
f"prompt_builder_factory(不可为 None)"
|
||
)
|
||
self._config = config
|
||
self._llm = llm
|
||
self._evolve_llm = evolve_llm
|
||
self._vlm = vlm
|
||
self._telemetry = telemetry
|
||
self._tool_dispatch_factory = tool_dispatch_factory
|
||
self._prompt_builder_factory = prompt_builder_factory
|
||
self._ensure_workspace()
|
||
self._paths: ResolvedPaths = resolve_paths(config.workspace_dir)
|
||
```
|
||
|
||
**3b.** 修改 `_make_tool_dispatch_fn`(第 2064-2072 行):
|
||
|
||
```python
|
||
def _make_tool_dispatch_fn(self, *, skills_dir: Path | None = None):
|
||
"""构造工具调度函数。优先用注入的 factory,fallback 为显式报错。"""
|
||
if self._tool_dispatch_factory is not None:
|
||
return self._tool_dispatch_factory(skills_dir=skills_dir)
|
||
|
||
async def _noop_dispatch(tool_name: str, args: dict, *, context: dict) -> str:
|
||
raise NotImplementedError(
|
||
f"工具 {tool_name} 调度未配置(需由 main.py 注入 tool_dispatch_fn)"
|
||
)
|
||
|
||
return _noop_dispatch
|
||
```
|
||
|
||
**3c.** 修改 `_make_prompt_builder`(第 2074-2082 行):
|
||
|
||
```python
|
||
def _make_prompt_builder(
|
||
self, *, skills_dir: Path | None = None, prompts_dir: Path | None = None
|
||
):
|
||
"""构造 prompt 构建函数。优先用注入的 factory,fallback 为显式报错。"""
|
||
if self._prompt_builder_factory is not None:
|
||
return self._prompt_builder_factory(
|
||
skills_dir=skills_dir, prompts_dir=prompts_dir
|
||
)
|
||
|
||
def _noop_builder(qa: GeneratedQuestion) -> tuple[str, str]:
|
||
raise NotImplementedError("prompt_builder 未配置(需由 main.py 注入)")
|
||
|
||
return _noop_builder
|
||
```
|
||
|
||
**3d.** 更新 docstring(第 440-452 行)增加两个新参数说明。
|
||
|
||
- [ ] **Step 4: 运行测试确认通过**
|
||
|
||
Run: `conda run -n Video-Tree-TRM pytest tests/unit/test_harness_runner.py -v`
|
||
Expected: ALL tests PASSED(含新增的 2 个 + 原有 34 个)
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
git add app/harness/runner.py tests/unit/test_harness_runner.py
|
||
git commit -m "feat(runner): 注入 tool_dispatch_factory/prompt_builder_factory + fail-fast"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 5: app/harness/deps_router.py — per-video 路由器
|
||
|
||
**Files:**
|
||
- Create: `app/harness/deps_router.py`
|
||
- Test: `tests/unit/test_deps_router.py`
|
||
|
||
- [ ] **Step 1: 编写 Router 单元测试**
|
||
|
||
```python
|
||
# tests/unit/test_deps_router.py
|
||
"""InferenceDepsRouter 单元测试。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from pathlib import Path
|
||
from typing import Any
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
import pytest
|
||
|
||
from app.harness.deps_router import InferenceDepsRouter
|
||
|
||
|
||
@pytest.fixture()
|
||
def mock_deps():
|
||
"""构造 mock InferenceDeps。"""
|
||
deps = MagicMock()
|
||
deps.prompt_builder = lambda qa: (f"system_{qa.video_id}", f"user_{qa.question}")
|
||
deps.tool_dispatch_fn = AsyncMock(return_value="tool_result")
|
||
return deps
|
||
|
||
|
||
@pytest.fixture()
|
||
def router(mock_deps):
|
||
"""构造带 mock build_inference_deps 的 Router。"""
|
||
r = InferenceDepsRouter(
|
||
store_dir=Path("store"),
|
||
embed_provider=MagicMock(),
|
||
llm=MagicMock(),
|
||
vlm=MagicMock(),
|
||
ocr=None,
|
||
default_prompts_dir=Path("store/prompts/v1"),
|
||
default_skills_dir=Path("store/skills/v1"),
|
||
skill_mode="auto",
|
||
verify_vision=True,
|
||
anchor=True,
|
||
assemble_mode="default",
|
||
)
|
||
# 替换 _build_deps 为 mock
|
||
r._build_deps = MagicMock(return_value=mock_deps)
|
||
return r
|
||
|
||
|
||
class TestPromptBuilder:
|
||
"""prompt_builder 注册映射并返回 prompt。"""
|
||
|
||
def test_registers_qid_to_vid_mapping(self, router, mock_deps) -> None:
|
||
from core.types import GeneratedQuestion
|
||
|
||
qa = GeneratedQuestion(
|
||
question_id="q1", video_id="vid1", task_type="Action Reasoning",
|
||
question="test?", options=("A. a", "B. b", "C. c", "D. d"),
|
||
answer="A", source_nodes=(), difficulty="medium",
|
||
)
|
||
builder = router.create_prompt_builder()
|
||
builder(qa)
|
||
assert router._qid_to_vid["q1"] == "vid1"
|
||
|
||
def test_returns_prompt_from_deps(self, router, mock_deps) -> None:
|
||
from core.types import GeneratedQuestion
|
||
|
||
qa = GeneratedQuestion(
|
||
question_id="q1", video_id="vid1", task_type="Action Reasoning",
|
||
question="test?", options=("A. a", "B. b", "C. c", "D. d"),
|
||
answer="A", source_nodes=(), difficulty="medium",
|
||
)
|
||
builder = router.create_prompt_builder()
|
||
system, user = builder(qa)
|
||
assert "vid1" in system
|
||
|
||
|
||
class TestDispatch:
|
||
"""dispatch 通过 session_id 路由到正确视频。"""
|
||
|
||
@pytest.mark.asyncio()
|
||
async def test_routes_by_session_id(self, router, mock_deps) -> None:
|
||
from core.types import GeneratedQuestion
|
||
|
||
qa = GeneratedQuestion(
|
||
question_id="q1", video_id="vid1", task_type="Action Reasoning",
|
||
question="test?", options=("A. a", "B. b", "C. c", "D. d"),
|
||
answer="A", source_nodes=(), difficulty="medium",
|
||
)
|
||
# 先注册映射
|
||
builder = router.create_prompt_builder()
|
||
builder(qa)
|
||
|
||
dispatch = router.create_dispatch()
|
||
result = await dispatch("view_node", {"node_id": "L1_000"}, context={"session_id": "q1"})
|
||
assert result == "tool_result"
|
||
mock_deps.tool_dispatch_fn.assert_called_once()
|
||
|
||
@pytest.mark.asyncio()
|
||
async def test_unknown_session_id_raises(self, router) -> None:
|
||
dispatch = router.create_dispatch()
|
||
with pytest.raises(KeyError, match="未注册"):
|
||
await dispatch("view_node", {}, context={"session_id": "unknown"})
|
||
|
||
@pytest.mark.asyncio()
|
||
async def test_missing_session_id_raises(self, router) -> None:
|
||
dispatch = router.create_dispatch()
|
||
with pytest.raises(KeyError, match="未注册"):
|
||
await dispatch("view_node", {}, context={})
|
||
|
||
|
||
class TestDepsCache:
|
||
"""同一 video_id 复用缓存。"""
|
||
|
||
def test_same_video_reuses_deps(self, router, mock_deps) -> None:
|
||
from core.types import GeneratedQuestion
|
||
|
||
qa1 = GeneratedQuestion(
|
||
question_id="q1", video_id="vid1", task_type="Action Reasoning",
|
||
question="test1?", options=("A. a", "B. b", "C. c", "D. d"),
|
||
answer="A",
|
||
)
|
||
qa2 = GeneratedQuestion(
|
||
question_id="q2", video_id="vid1", task_type="Action Reasoning",
|
||
question="test2?", options=("A. a", "B. b", "C. c", "D. d"),
|
||
answer="B",
|
||
)
|
||
builder = router.create_prompt_builder()
|
||
builder(qa1)
|
||
builder(qa2)
|
||
# 同一 video_id 只调用一次 _build_deps
|
||
assert router._build_deps.call_count == 1
|
||
|
||
|
||
class TestClearCache:
|
||
"""clear_cache 清空缓存和映射。"""
|
||
|
||
def test_clears_deps_and_mapping(self, router, mock_deps) -> None:
|
||
from core.types import GeneratedQuestion
|
||
|
||
qa = GeneratedQuestion(
|
||
question_id="q1", video_id="vid1", task_type="Action Reasoning",
|
||
question="test?", options=("A. a", "B. b", "C. c", "D. d"),
|
||
answer="A", source_nodes=(), difficulty="medium",
|
||
)
|
||
builder = router.create_prompt_builder()
|
||
builder(qa)
|
||
assert len(router._qid_to_vid) == 1
|
||
|
||
router.clear_cache()
|
||
assert len(router._qid_to_vid) == 0
|
||
assert len(router._deps_cache) == 0
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试确认失败**
|
||
|
||
Run: `conda run -n Video-Tree-TRM pytest tests/unit/test_deps_router.py -v`
|
||
Expected: ModuleNotFoundError(deps_router 尚未创建)
|
||
|
||
- [ ] **Step 3: 实现 InferenceDepsRouter**
|
||
|
||
```python
|
||
# app/harness/deps_router.py
|
||
"""按 video_id 懒加载 InferenceDeps 并路由 dispatch/prompt_builder。
|
||
|
||
将 Runner 的全局统一 dispatch/prompt_builder 接口路由到 per-video 的
|
||
InferenceDeps。时序保证:prompt_builder(qa) 先于 dispatch 被调用,
|
||
在 prompt_builder 中注册 question_id → video_id 映射。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from typing import TYPE_CHECKING, Any
|
||
|
||
from loguru import logger
|
||
|
||
from app.harness.factory import InferenceDeps, build_inference_deps
|
||
|
||
if TYPE_CHECKING:
|
||
from pathlib import Path
|
||
|
||
from app.ports import EmbeddingProvider, OCRProvider
|
||
from core.protocols import LLMProvider, VLMProvider
|
||
from core.types import GeneratedQuestion
|
||
|
||
|
||
class InferenceDepsRouter:
|
||
"""按 video_id 懒加载 InferenceDeps 并路由工具调度和 prompt 构建。
|
||
|
||
参数:
|
||
store_dir: store 根目录。
|
||
embed_provider: 嵌入端口。
|
||
llm: LLM 端口。
|
||
vlm: VLM 端口。
|
||
ocr: OCR 端口(None 不启用)。
|
||
default_prompts_dir: 默认 prompts 版本目录。
|
||
default_skills_dir: 默认 skills 版本目录。
|
||
skill_mode: skill 加载模式。
|
||
verify_vision: observe_frame 是否验证。
|
||
anchor: view_node 是否锚模式。
|
||
assemble_mode: 锚模式装配形态。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
*,
|
||
store_dir: Path,
|
||
embed_provider: EmbeddingProvider,
|
||
llm: LLMProvider,
|
||
vlm: VLMProvider,
|
||
ocr: OCRProvider | None,
|
||
default_prompts_dir: Path,
|
||
default_skills_dir: Path,
|
||
skill_mode: str,
|
||
verify_vision: bool,
|
||
anchor: bool,
|
||
assemble_mode: str,
|
||
) -> None:
|
||
self._store_dir = store_dir
|
||
self._embed = embed_provider
|
||
self._llm = llm
|
||
self._vlm = vlm
|
||
self._ocr = ocr
|
||
self._default_prompts_dir = default_prompts_dir
|
||
self._default_skills_dir = default_skills_dir
|
||
self._skill_mode = skill_mode
|
||
self._verify_vision = verify_vision
|
||
self._anchor = anchor
|
||
self._assemble_mode = assemble_mode
|
||
self._deps_cache: dict[tuple[str, str, str], InferenceDeps] = {}
|
||
self._qid_to_vid: dict[str, str] = {}
|
||
|
||
def create_dispatch(
|
||
self, *, skills_dir: Path | None = None
|
||
) -> Any:
|
||
"""创建工具调度函数,通过 context["session_id"] 路由到 per-video dispatcher。
|
||
|
||
参数:
|
||
skills_dir: 可选的 skills 版本目录覆盖。
|
||
|
||
返回:
|
||
async (tool_name, args, *, context) -> str。
|
||
"""
|
||
effective_skills = skills_dir or self._default_skills_dir
|
||
|
||
async def _dispatch(
|
||
tool_name: str, args: dict[str, Any], *, context: dict[str, Any]
|
||
) -> str:
|
||
session_id = context.get("session_id")
|
||
if not session_id or session_id not in self._qid_to_vid:
|
||
raise KeyError(
|
||
f"未注册的 session_id={session_id!r},"
|
||
f"已注册 {len(self._qid_to_vid)} 条映射"
|
||
)
|
||
video_id = self._qid_to_vid[session_id]
|
||
deps = self._ensure_deps(video_id, effective_skills, self._default_prompts_dir)
|
||
return await deps.tool_dispatch_fn(tool_name, args, context=context)
|
||
|
||
return _dispatch
|
||
|
||
def create_prompt_builder(
|
||
self,
|
||
*,
|
||
skills_dir: Path | None = None,
|
||
prompts_dir: Path | None = None,
|
||
) -> Any:
|
||
"""创建 prompt 构建函数,在调用时注册 question_id→video_id 映射。
|
||
|
||
参数:
|
||
skills_dir: 可选的 skills 版本目录覆盖。
|
||
prompts_dir: 可选的 prompts 版本目录覆盖。
|
||
|
||
返回:
|
||
(GeneratedQuestion) -> (system_prompt, user_prompt)。
|
||
"""
|
||
effective_skills = skills_dir or self._default_skills_dir
|
||
effective_prompts = prompts_dir or self._default_prompts_dir
|
||
|
||
def _builder(qa: GeneratedQuestion) -> tuple[str, str]:
|
||
self._qid_to_vid[qa.question_id] = qa.video_id
|
||
deps = self._ensure_deps(qa.video_id, effective_skills, effective_prompts)
|
||
return deps.prompt_builder(qa)
|
||
|
||
return _builder
|
||
|
||
def _ensure_deps(
|
||
self, video_id: str, skills_dir: Path, prompts_dir: Path
|
||
) -> InferenceDeps:
|
||
"""懒加载 per-video InferenceDeps,按 (video_id, skills_dir, prompts_dir) 缓存。"""
|
||
key = (video_id, str(skills_dir), str(prompts_dir))
|
||
if key not in self._deps_cache:
|
||
self._deps_cache[key] = self._build_deps(video_id, skills_dir, prompts_dir)
|
||
logger.debug("InferenceDeps 已缓存: video_id={}", video_id)
|
||
return self._deps_cache[key]
|
||
|
||
def _build_deps(
|
||
self, video_id: str, skills_dir: Path, prompts_dir: Path
|
||
) -> InferenceDeps:
|
||
"""调用 factory.build_inference_deps 构建 per-video 依赖。"""
|
||
return build_inference_deps(
|
||
store_dir=self._store_dir,
|
||
video_id=video_id,
|
||
prompts_dir=prompts_dir,
|
||
skills_dir=skills_dir,
|
||
skill_mode=self._skill_mode,
|
||
embed_provider=self._embed,
|
||
llm=self._llm,
|
||
vlm=self._vlm,
|
||
ocr=self._ocr,
|
||
verify_vision=self._verify_vision,
|
||
anchor=self._anchor,
|
||
assemble_mode=self._assemble_mode,
|
||
)
|
||
|
||
def clear_cache(self) -> None:
|
||
"""清空 deps 缓存和 question_id 映射。"""
|
||
self._deps_cache.clear()
|
||
self._qid_to_vid.clear()
|
||
```
|
||
|
||
- [ ] **Step 4: 运行测试确认通过**
|
||
|
||
Run: `conda run -n Video-Tree-TRM pytest tests/unit/test_deps_router.py -v`
|
||
Expected: 7 tests PASSED
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
git add app/harness/deps_router.py tests/unit/test_deps_router.py
|
||
git commit -m "feat(harness): InferenceDepsRouter per-video 路由器"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 6: 配置变更
|
||
|
||
**Files:**
|
||
- Modify: `config/default.yaml:29,31`
|
||
- Modify: `.env` (LLM_CIRCUIT_BREAKER_THRESHOLD)
|
||
|
||
- [ ] **Step 1: 更新 default.yaml**
|
||
|
||
在 `config/default.yaml` 中修改:
|
||
|
||
```yaml
|
||
harness:
|
||
concurrency: 24 # was 12
|
||
max_steps: 40 # was 15
|
||
```
|
||
|
||
- [ ] **Step 2: 更新 .env 和 .env.example**
|
||
|
||
在 `.env` 中修改(本地生效,不提交):
|
||
```
|
||
LLM_CIRCUIT_BREAKER_THRESHOLD=48
|
||
```
|
||
|
||
在 `.env.example` 中同步(提交到 Git):
|
||
```
|
||
LLM_CIRCUIT_BREAKER_THRESHOLD=48 # 实际阈值 = max(此值, concurrency*2)
|
||
```
|
||
|
||
- [ ] **Step 3: 验证配置加载**
|
||
|
||
Run: `conda run -n Video-Tree-TRM python -c "from app.harness.config import load_config; from pathlib import Path; c = load_config(Path('config/default.yaml')); print(f'concurrency={c.concurrency}, max_steps={c.max_steps}')"`
|
||
Expected: `concurrency=24, max_steps=40`
|
||
|
||
- [ ] **Step 4: Commit**
|
||
|
||
```bash
|
||
git add config/default.yaml .env.example
|
||
git commit -m "config: concurrency=24, max_steps=40, breaker_threshold=48"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 7: main.py — Composition Root
|
||
|
||
**Files:**
|
||
- Create: `main.py`
|
||
|
||
- [ ] **Step 1: 实现 main.py**
|
||
|
||
```python
|
||
"""CLI 入口 — Composition Root:构建适配器,注入 Runner,调度执行。
|
||
|
||
三层配置合并(YAML > .env > CLI)由 load_config 完成。
|
||
适配器参数通过 InfraSettings(BaseSettings) 从 .env 加载。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import asyncio
|
||
import os
|
||
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 tree.embed 或 harness 段)。
|
||
|
||
返回:
|
||
_Adapters 命名元组。
|
||
"""
|
||
from adapters.breaker import CircuitBreaker
|
||
from adapters.embedding import LocalEmbeddingProvider, RemoteEmbeddingProvider
|
||
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:
|
||
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"),
|
||
help="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 != "config" and k != "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
|
||
from app.harness.workspace import resolve_paths
|
||
|
||
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()
|
||
```
|
||
|
||
- [ ] **Step 2: 验证 main.py 可导入**
|
||
|
||
Run: `conda run -n Video-Tree-TRM python -c "import main; print('OK')"`
|
||
Expected: `OK`
|
||
|
||
- [ ] **Step 3: 验证 --help 输出**
|
||
|
||
Run: `conda run -n Video-Tree-TRM python main.py --help`
|
||
Expected: 显示参数帮助文本
|
||
|
||
- [ ] **Step 4: Commit**
|
||
|
||
```bash
|
||
git add main.py
|
||
git commit -m "feat: main.py Composition Root(仅 infer 模式)"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 8: 冒烟测试
|
||
|
||
**Files:**
|
||
- No new files, end-to-end validation
|
||
|
||
- [ ] **Step 1: 验证 workspace 初始化**
|
||
|
||
需要先初始化 workspace(将 store 的 v1 资源拷贝到 workspace)。
|
||
|
||
Run: `conda run -n Video-Tree-TRM python main.py --mode infer --n-samples 1 --concurrency 1 --max-steps 3`
|
||
|
||
检查:
|
||
- 是否成功初始化 workspace
|
||
- 是否加载了题目
|
||
- 是否创建了 InferenceDeps
|
||
- LLM 调用是否经过 GovernedLLMClient
|
||
|
||
如果出现 workspace 不存在的错误,需要先手动初始化:
|
||
|
||
```python
|
||
from pathlib import Path
|
||
from app.harness.workspace import init_workspace
|
||
init_workspace(
|
||
Path("workspaces/default"),
|
||
Path("store"),
|
||
"benchmarks/Video-MME",
|
||
"v1",
|
||
"v1",
|
||
)
|
||
```
|
||
|
||
- [ ] **Step 2: 检查结果输出**
|
||
|
||
Expected: 看到日志输出包含:
|
||
- `配置加载完成: mode=infer`
|
||
- `InferenceDeps 已缓存: video_id=...`
|
||
- `推理完成: accuracy=...`
|
||
|
||
- [ ] **Step 3: 运行全量测试确认无回归**
|
||
|
||
Run: `conda run -n Video-Tree-TRM pytest tests/unit/ -v --tb=short`
|
||
Expected: ALL PASSED
|
||
|
||
- [ ] **Step 4: 最终提交**
|
||
|
||
```bash
|
||
git add -A
|
||
git commit -m "test: 冒烟测试通过,900 题推理管线就绪"
|
||
```
|
||
|
||
---
|
||
|
||
## 核心算法保真校验
|
||
|
||
本计划不涉及核心算法迁移。所有 13 项核心算法(L2 轴心建树、CE-Gate e-process、Agent Loop 等)在 TRM5 中已有完整实现。本次工作仅涉及:
|
||
- 依赖注入的架构改进(Runner factory 注入)
|
||
- 新增基础设施模块(InferenceDepsRouter、main.py)
|
||
- 内容准备(skills/v1、prompts/v1 目录重组)
|
||
- 配置变更
|
||
|
||
保真校验不适用。
|