Files
Video-Tree-TRM5/research-wiki/plans/2026-07-09-main-inference-entry.md
T
2026-07-09 10:43:15 -04:00

1317 lines
42 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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 个 ProtocolToolDispatchFn, 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 Rootargparse + 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: ImportErrorToolDispatchFn 等尚未定义)
- [ ] **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: FAILRunner.__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 构建函数。优先用注入的 factoryfallback 为显式报错。"""
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: ModuleNotFoundErrordeps_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 目录重组)
- 配置变更
保真校验不适用。