diff --git a/app/tree/enhance/__init__.py b/app/tree/repair/__init__.py similarity index 100% rename from app/tree/enhance/__init__.py rename to app/tree/repair/__init__.py diff --git a/app/tree/repair/regenerator.py b/app/tree/repair/regenerator.py new file mode 100644 index 0000000..068a669 --- /dev/null +++ b/app/tree/repair/regenerator.py @@ -0,0 +1,513 @@ +"""树修复重生成器:VLM 重新描述问题节点 + 底向上级联。 + +底向上修复流程: + 1. 收集需修复的 L3 节点 → VLM 重新描述帧 + 2. 收集受影响的 L2 → LLM 从 L3 children 聚合 + 3. 收集受影响的 L1 → LLM 从 L2 children 聚合 + +仅处理 issue_type == "empty_field" 且 level == 3 的问题节点。 +帧文件不存在时跳过该节点(不中断整体修复流程)。 +""" + +from __future__ import annotations + +import json +import re +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any + +from loguru import logger + +from app.tree.index import ( + L1Card, + L1Node, + L2Card, + L2Node, + L3Card, + L3Node, + TreeIndex, +) +from app.tree.subtitle import extract_subtitle_for_range + +if TYPE_CHECKING: + from pathlib import Path + + from app.tree.repair.detector import NodeIssue + from app.tree.subtitle import SRTEntry + from core.protocols import LLMProvider, VLMProvider + +# --------------------------------------------------------------------------- +# Prompt 常量(与 VideoTreeBuilder 保持一致风格) +# --------------------------------------------------------------------------- + +_L3_REPAIR_PROMPT = ( + '该片段的整体内容: "{l2_description}"\n' + "用一到两句话描述这帧画面的具体内容。" + "重点关注: 动作、物体变化、文字信息、人物表情。\n" + "{subtitle_block}" + "返回 JSON 对象,包含以下字段:\n" + "- frame_summary: 画面描述\n" + "- visible_entities: 可见实体列表\n" + "- ongoing_actions: 动作列表\n" + "- visible_text: 可见文字列表\n" + "- spatial_layout: 空间布局\n" + '- visual_attributes: {{"lighting": "...", "dominant_colors": [...], "camera_angle": "..."}}\n' + "只返回 JSON 对象,不要其他内容。" +) + +_L2_REGEN_PROMPT = ( + "以下是一个视频片段中各帧的描述:\n{l3_texts}\n" + "用1-2句话描述该片段的核心内容。\n" + "返回 JSON 对象,包含以下字段:\n" + "- event_description: 1-2句片段描述\n" + "- entities: 可见实体列表\n" + "- actions: 动作列表\n" + "- action_subjects: 动作主体列表\n" + "- visible_text: 画面中可见文字列表\n" + "- spatial_relations: 空间关系描述\n" + "- state_changes: 状态变化描述(无则 null)\n" + "只返回 JSON 对象,不要其他内容。" +) + +_L1_REGEN_PROMPT = ( + "以下是一个视频段落中各片段的描述:\n{l2_texts}\n" + "用2-3句话总结该段落的整体内容,涵盖所有片段的主题。\n" + "返回 JSON 对象,包含以下字段:\n" + "- scene_summary: 2-3句段落摘要\n" + "- main_setting: 主要场景\n" + "- key_entities: 关键实体列表\n" + "- main_actions: 主要动作列表\n" + "- topic_keywords: 主题关键词列表\n" + "- visible_text: 出现的文字列表\n" + "- temporal_flow: 时间流向描述\n" + "只返回 JSON 对象,不要其他内容。" +) + + +# --------------------------------------------------------------------------- +# 统计数据类 +# --------------------------------------------------------------------------- + + +@dataclass +class RepairStats: + """修复统计信息。 + + 属性: + l3_repaired: 修复的 L3 节点数。 + l2_regenerated: 重生成的 L2 节点数。 + l1_regenerated: 重生成的 L1 节点数。 + """ + + l3_repaired: int = 0 + l2_regenerated: int = 0 + l1_regenerated: int = 0 + + +# --------------------------------------------------------------------------- +# JSON 解析辅助(复用 VideoTreeBuilder 的解析逻辑) +# --------------------------------------------------------------------------- + + +def _extract_json(raw: str) -> Any: + """从 VLM/LLM 原始输出中提取 JSON(处理 markdown 代码块包裹)。 + + 参数: + raw: 原始返回字符串。 + + 返回: + 解析后的 Python 对象(dict/list),解析失败返回 None。 + """ + raw = raw.strip() + # Phase 1: 尝试提取 markdown 代码块中的 JSON + code_match = re.search( + r"```(?:json)?\s*([\[{].*?[\]}])\s*```", + raw, + re.DOTALL, + ) + if code_match: + raw = code_match.group(1) + + # Phase 2: 直接解析 + try: + return json.loads(raw) + except json.JSONDecodeError: + pass + + # Phase 3: 尝试提取裸 JSON 对象/数组 + json_match = re.search(r"[\[{].*[\]}]", raw, re.DOTALL) + if json_match: + try: + return json.loads(json_match.group()) + except json.JSONDecodeError: + pass + + return None + + +def _parse_l3_card(raw: str) -> L3Card | None: + """解析 VLM 输出为 L3Card。解析失败返回 None。 + + 参数: + raw: VLM 原始返回字符串。 + + 返回: + L3Card 实例或 None(解析失败时)。 + """ + data = _extract_json(raw) + if isinstance(data, dict): + try: + return L3Card( + frame_summary=str(data["frame_summary"]), + visible_entities=list(data["visible_entities"]), + ongoing_actions=list(data["ongoing_actions"]), + visible_text=list(data["visible_text"]), + spatial_layout=str(data["spatial_layout"]), + visual_attributes=dict(data["visual_attributes"]), + ) + except (KeyError, TypeError, ValueError): + pass + return None + + +def _parse_l2_card(raw: str) -> L2Card | None: + """解析 LLM 输出为 L2Card。解析失败返回 None。 + + 参数: + raw: LLM 原始返回字符串。 + + 返回: + L2Card 实例或 None(解析失败时)。 + """ + data = _extract_json(raw) + if isinstance(data, dict): + try: + state_changes = data.get("state_changes") + if state_changes is not None: + state_changes = str(state_changes) + return L2Card( + event_description=str(data["event_description"]), + entities=list(data["entities"]), + actions=list(data["actions"]), + action_subjects=list(data["action_subjects"]), + visible_text=list(data["visible_text"]), + spatial_relations=str(data["spatial_relations"]), + state_changes=state_changes, + ) + except (KeyError, TypeError, ValueError): + pass + return None + + +def _parse_l1_card(raw: str) -> L1Card | None: + """解析 LLM 输出为 L1Card。解析失败返回 None。 + + 参数: + raw: LLM 原始返回字符串。 + + 返回: + L1Card 实例或 None(解析失败时)。 + """ + data = _extract_json(raw) + if isinstance(data, dict): + try: + return L1Card( + scene_summary=str(data["scene_summary"]), + main_setting=str(data["main_setting"]), + key_entities=list(data["key_entities"]), + main_actions=list(data["main_actions"]), + topic_keywords=list(data["topic_keywords"]), + visible_text=list(data["visible_text"]), + temporal_flow=str(data["temporal_flow"]), + ) + except (KeyError, TypeError, ValueError): + pass + return None + + +# --------------------------------------------------------------------------- +# 节点查找辅助 +# --------------------------------------------------------------------------- + + +def _build_node_lookup( + index: TreeIndex, +) -> tuple[ + dict[str, L3Node], + dict[str, L2Node], + dict[str, L1Node], + dict[str, L2Node], + dict[str, L1Node], +]: + """构建节点 ID 到节点的查找表 + 子节点到父节点的映射。 + + 参数: + index: 树索引。 + + 返回: + (l3_by_id, l2_by_id, l1_by_id, l3_parent_l2, l2_parent_l1) + - l3_by_id: L3 节点 ID → L3Node + - l2_by_id: L2 节点 ID → L2Node + - l1_by_id: L1 节点 ID → L1Node + - l3_parent_l2: L3 节点 ID → 其父 L2Node + - l2_parent_l1: L2 节点 ID → 其父 L1Node + """ + l3_by_id: dict[str, L3Node] = {} + l2_by_id: dict[str, L2Node] = {} + l1_by_id: dict[str, L1Node] = {} + l3_parent_l2: dict[str, L2Node] = {} + l2_parent_l1: dict[str, L1Node] = {} + + for l1 in index.roots: + l1_by_id[l1.id] = l1 + for l2 in l1.children: + l2_by_id[l2.id] = l2 + l2_parent_l1[l2.id] = l1 + for l3 in l2.children: + l3_by_id[l3.id] = l3 + l3_parent_l2[l3.id] = l2 + + return l3_by_id, l2_by_id, l1_by_id, l3_parent_l2, l2_parent_l1 + + +# --------------------------------------------------------------------------- +# 字幕辅助 +# --------------------------------------------------------------------------- + + +def _build_subtitle_block( + srt_entries: list[SRTEntry] | None, + timestamp: float | None, +) -> str: + """构建字幕注入文本块。 + + 参数: + srt_entries: SRT 字幕条目列表。 + timestamp: 帧时间戳(秒)。 + + 返回: + 字幕文本块字符串(无匹配时返回空字符串)。 + """ + if not srt_entries or timestamp is None: + return "" + window = 2.0 + start = max(0.0, timestamp - window) + end = timestamp + window + text = extract_subtitle_for_range(srt_entries, (start, end)) + if not text: + return "" + return f"字幕信息:\n{text}\n" + + +# --------------------------------------------------------------------------- +# 主修复函数 +# --------------------------------------------------------------------------- + + +async def repair_tree( + index: TreeIndex, + issues: list[NodeIssue], + vlm: VLMProvider, + llm: LLMProvider, + frames_dir: Path, + srt_entries: list[SRTEntry] | None = None, +) -> RepairStats: + """修复有问题的节点,底向上级联。 + + 流程: + 1. 收集需修复的 L3 节点 → VLM 重新描述帧 + 2. 收集受影响的 L2 → LLM 从 L3 children 聚合 + 3. 收集受影响的 L1 → LLM 从 L2 children 聚合 + + 参数: + index: 待修复的 TreeIndex(原地修改)。 + issues: detect_issues() 返回的问题列表。 + vlm: VLM 调用端口。 + llm: LLM 调用端口。 + frames_dir: 帧文件根目录。 + srt_entries: 字幕条目列表(可选)。 + + 返回: + RepairStats 统计。 + """ + stats = RepairStats() + + if not issues: + logger.info("无修复任务,跳过") + return stats + + # 构建查找表 + l3_by_id, l2_by_id, l1_by_id, l3_parent_l2, l2_parent_l1 = _build_node_lookup(index) + + # Step 1: 修复 L3 节点(仅处理 empty_field + level 3) + l3_issues = [ + issue for issue in issues if issue.issue_type == "empty_field" and issue.level == 3 + ] + + affected_l2_ids: set[str] = set() + + for issue in l3_issues: + l3_node = l3_by_id.get(issue.node_id) + if l3_node is None: + logger.warning( + "L3 节点 ID 未在树中找到,跳过", + node_id=issue.node_id, + ) + continue + + # 查找帧文件 + if l3_node.frame_path is None: + logger.warning( + "L3 节点无 frame_path,跳过", + node_id=issue.node_id, + ) + continue + + frame_file = frames_dir / l3_node.frame_path + if not frame_file.exists(): + logger.warning( + "L3 帧文件不存在,跳过修复", + node_id=issue.node_id, + frame_path=str(frame_file), + ) + continue + + # 获取 L2 父节点描述作为上下文 + parent_l2 = l3_parent_l2.get(issue.node_id) + l2_description = parent_l2.card.event_description if parent_l2 else "" + + # 构建字幕块 + subtitle_block = _build_subtitle_block(srt_entries, l3_node.timestamp) + + # VLM 重新描述帧 + prompt = _L3_REPAIR_PROMPT.format( + l2_description=l2_description, + subtitle_block=subtitle_block, + ) + messages = [{"role": "user", "content": prompt}] + + try: + response = await vlm.chat_with_images(messages, [str(frame_file)]) + except Exception as exc: + logger.warning( + "L3 修复 VLM 调用失败,跳过: {}", + exc, + node_id=issue.node_id, + ) + continue + + new_card = _parse_l3_card(response.content) + if new_card is None: + logger.warning( + "L3 修复 VLM 输出解析失败,跳过", + node_id=issue.node_id, + raw_preview=response.content[:200], + ) + continue + + # 原地替换 card(L3Node.card 不是 frozen dataclass 的限制字段) + l3_node.card = new_card + stats.l3_repaired += 1 + + # 标记受影响的 L2 父节点 + if parent_l2 is not None: + affected_l2_ids.add(parent_l2.id) + + logger.debug( + "L3 节点修复完成", + node_id=issue.node_id, + frame_summary=new_card.frame_summary[:50], + ) + + # Step 2: 重生成受影响的 L2 节点 + affected_l1_ids: set[str] = set() + + for l2_id in affected_l2_ids: + l2_node = l2_by_id.get(l2_id) + if l2_node is None: + continue + + # 从 L3 children 聚合描述 + l3_texts = "\n".join(f"- {l3.card.frame_summary}" for l3 in l2_node.children) + prompt = _L2_REGEN_PROMPT.format(l3_texts=l3_texts) + messages = [{"role": "user", "content": prompt}] + + try: + response = await llm.chat(messages) + except Exception as exc: + logger.warning( + "L2 重生成 LLM 调用失败,跳过: {}", + exc, + l2_id=l2_id, + ) + continue + + new_card = _parse_l2_card(response.content) + if new_card is None: + logger.warning( + "L2 重生成 LLM 输出解析失败,跳过", + l2_id=l2_id, + raw_preview=response.content[:200], + ) + continue + + l2_node.card = new_card + stats.l2_regenerated += 1 + + # 标记受影响的 L1 父节点 + parent_l1 = l2_parent_l1.get(l2_id) + if parent_l1 is not None: + affected_l1_ids.add(parent_l1.id) + + logger.debug( + "L2 节点重生成完成", + l2_id=l2_id, + event_description=new_card.event_description[:50], + ) + + # Step 3: 重生成受影响的 L1 节点 + for l1_id in affected_l1_ids: + l1_node = l1_by_id.get(l1_id) + if l1_node is None: + continue + + # 从 L2 children 聚合描述 + l2_texts = "\n".join(f"- {l2.card.event_description}" for l2 in l1_node.children) + prompt = _L1_REGEN_PROMPT.format(l2_texts=l2_texts) + messages = [{"role": "user", "content": prompt}] + + try: + response = await llm.chat(messages) + except Exception as exc: + logger.warning( + "L1 重生成 LLM 调用失败,跳过: {}", + exc, + l1_id=l1_id, + ) + continue + + new_card = _parse_l1_card(response.content) + if new_card is None: + logger.warning( + "L1 重生成 LLM 输出解析失败,跳过", + l1_id=l1_id, + raw_preview=response.content[:200], + ) + continue + + l1_node.card = new_card + stats.l1_regenerated += 1 + + logger.debug( + "L1 节点重生成完成", + l1_id=l1_id, + scene_summary=new_card.scene_summary[:50], + ) + + logger.info( + "树修复完成", + l3_repaired=stats.l3_repaired, + l2_regenerated=stats.l2_regenerated, + l1_regenerated=stats.l1_regenerated, + ) + return stats diff --git a/app/tree/repair/supplement.py b/app/tree/repair/supplement.py index 3924896..499ec5b 100644 --- a/app/tree/repair/supplement.py +++ b/app/tree/repair/supplement.py @@ -84,10 +84,11 @@ def deduplicate_field(values: list[str]) -> list[str]: seen: set[str] = set() result: list[str] = [] for v in values: - key = v.strip().lower() + s = str(v).strip() + key = s.lower() if key and key not in seen: seen.add(key) - result.append(v) + result.append(s) return result @@ -278,7 +279,7 @@ def apply_injections(index: TreeIndex, injections: list[dict[str, Any]]) -> Supp stats.facts_skipped += 1 continue - inject_value = instr.get("inject_value", "") + inject_value = str(instr.get("inject_value", "")).strip() if not inject_value: stats.facts_skipped += 1 continue diff --git a/core/evolution/patch.py b/core/evolution/patch.py index d4cef82..87effa5 100644 --- a/core/evolution/patch.py +++ b/core/evolution/patch.py @@ -15,9 +15,7 @@ APPENDIX_MAX_CHARS = 2000 # appendix 区软上限(守设计「长度上限+wa MOMENTUM_START = "" MOMENTUM_END = "" MOMENTUM_MAX_CHARS = 2000 # momentum 区软上限(与 appendix 一致:超限 warning 不截断) -MOMENTUM_HEADING = ( - "## 动量指导(每轮重写,勿手改)" # replace_momentum 写入的固定标题行 -) +MOMENTUM_HEADING = "## 动量指导(每轮重写,勿手改)" # replace_momentum 写入的固定标题行 def momentum_region_bounds(text: str) -> tuple[int, int] | None: @@ -303,9 +301,7 @@ def _insert_at(content: str, at: int, payload: str) -> str: return head + "\n\n" + payload + "\n" -def _do_append( - content: str, payload: str, ranges: list[tuple[int, int]] -) -> tuple[str, str]: +def _do_append(content: str, payload: str, ranges: list[tuple[int, int]]) -> tuple[str, str]: """执行 append 操作,返回更新后内容与状态字符串。""" return _insert_at(content, _append_at(content, ranges), payload), "applied_append" @@ -351,9 +347,7 @@ def _do_replace_delete( return new_content, "applied_" + op -def _apply_one( - content: str, edit: dict, ranges: list[tuple[int, int]] -) -> tuple[str, dict]: +def _apply_one(content: str, edit: dict, ranges: list[tuple[int, int]]) -> tuple[str, dict]: """应用单条 edit,返回 (更新后内容, 状态报告)。""" if not isinstance(edit, dict): return content, { @@ -382,9 +376,7 @@ def _apply_one( return content, report if op in ("replace", "delete"): - content, report["status"] = _do_replace_delete( - op, content, target, payload, ranges - ) + content, report["status"] = _do_replace_delete(op, content, target, payload, ranges) return content, report logger.warning("未知 op,跳过: {}", op) diff --git a/scripts/repair_trees.sh b/scripts/repair_trees.sh index 2f20f37..b228117 100755 --- a/scripts/repair_trees.sh +++ b/scripts/repair_trees.sh @@ -17,6 +17,8 @@ set -euo pipefail CONCURRENCY="${CONCURRENCY:-16}" +# shellcheck source=/dev/null +source "$(conda info --base)/etc/profile.d/conda.sh" conda activate Video-Tree-TRM python tools/repair_trees.py \ diff --git a/tests/unit/test_repair_regenerator.py b/tests/unit/test_repair_regenerator.py new file mode 100644 index 0000000..6d58630 --- /dev/null +++ b/tests/unit/test_repair_regenerator.py @@ -0,0 +1,293 @@ +"""修复重生成器单元测试。""" + +from __future__ import annotations + +import asyncio +import json +from pathlib import Path + +from app.tree.index import ( + IndexMeta, + L1Card, + L1Node, + L2Card, + L2Node, + L3Card, + L3Node, + TreeIndex, +) +from app.tree.repair.detector import NodeIssue +from app.tree.repair.regenerator import RepairStats, repair_tree +from core.types import LLMResponse + + +def _mock_response(content: str) -> LLMResponse: + """构造模拟 LLMResponse。""" + return LLMResponse( + content=content, + thinking="", + model="mock", + provider="mock", + prompt_tokens=0, + completion_tokens=0, + latency_ms=0, + ttft_ms=None, + max_inter_token_ms=None, + cache_hit=False, + call_id="mock", + ) + + +class MockVLM: + """模拟 VLM 端口,返回固定的 L3Card JSON。""" + + def __init__(self) -> None: + self.call_count = 0 + + async def chat_with_images( + self, + messages: list[dict], + images: list, + **kw: object, + ) -> LLMResponse: + """模拟 VLM 图文调用。""" + self.call_count += 1 + return _mock_response( + json.dumps( + { + "frame_summary": "修复后的帧描述", + "visible_entities": ["修复实体"], + "ongoing_actions": ["修复动作"], + "visible_text": [], + "spatial_layout": "居中", + "visual_attributes": {"lighting": "明亮"}, + } + ) + ) + + +class MockLLM: + """模拟 LLM 端口,根据 prompt 内容返回 L2Card 或 L1Card JSON。""" + + def __init__(self) -> None: + self.call_count = 0 + + async def chat( + self, + messages: list[dict], + **kw: object, + ) -> LLMResponse: + """模拟 LLM 文本调用,按 prompt 内容区分 L2/L1 响应。""" + self.call_count += 1 + content = messages[-1].get("content", "") + if "段落" in content or "scene" in content.lower(): + return _mock_response( + json.dumps( + { + "scene_summary": "修复后的场景", + "main_setting": "室内", + "key_entities": [], + "main_actions": [], + "topic_keywords": [], + "visible_text": [], + "temporal_flow": "", + } + ) + ) + return _mock_response( + json.dumps( + { + "event_description": "修复后的事件", + "entities": [], + "actions": [], + "action_subjects": [], + "visible_text": [], + "spatial_relations": "", + "state_changes": None, + } + ) + ) + + +class TestRepairTree: + """repair_tree 核心测试。""" + + def _make_broken_tree( + self, + tmp_path: Path, + ) -> tuple[TreeIndex, list[NodeIssue]]: + """构建含一个空 frame_summary 的 L3 节点的测试树。""" + frame_path = tmp_path / "frames" / "L1_000_L2_000_L3_000.jpg" + frame_path.parent.mkdir(parents=True) + frame_path.write_bytes(b"\xff\xd8\xff\xe0fake") + + l3 = L3Node( + id="vid_L1_000_L2_000_L3_000", + card=L3Card("", [], [], [], "", {}), + timestamp=1.0, + frame_path="frames/L1_000_L2_000_L3_000.jpg", + ) + l2 = L2Node( + id="vid_L1_000_L2_000", + card=L2Card("原始事件", [], [], [], [], "", None), + time_range=(0.0, 10.0), + children=[l3], + ) + l1 = L1Node( + id="vid_L1_000", + card=L1Card("原始场景", "", [], [], [], [], ""), + time_range=(0.0, 10.0), + children=[l2], + ) + index = TreeIndex(metadata=IndexMeta("/t.mp4", "video"), roots=[l1]) + issues = [ + NodeIssue( + "vid_L1_000_L2_000_L3_000", + 3, + "empty_field", + "frame_summary 为空", + ) + ] + return index, issues + + def test_repairs_l3_and_cascades(self, tmp_path: Path) -> None: + """修复 L3 后应级联重生成 L2 和 L1。""" + index, issues = self._make_broken_tree(tmp_path) + stats = asyncio.run(repair_tree(index, issues, MockVLM(), MockLLM(), tmp_path)) + assert stats.l3_repaired == 1 + assert stats.l2_regenerated == 1 + assert stats.l1_regenerated == 1 + assert index.roots[0].children[0].children[0].card.frame_summary == "修复后的帧描述" + assert index.roots[0].children[0].card.event_description == "修复后的事件" + assert index.roots[0].card.scene_summary == "修复后的场景" + + def test_no_issues_no_changes(self) -> None: + """无问题时不进行任何修复。""" + l3 = L3Node( + id="l1_0_l2_0_l3_0", + card=L3Card("正常", [], [], [], "", {}), + timestamp=1.0, + ) + l2 = L2Node( + id="l1_0_l2_0", + card=L2Card("正常事件", [], [], [], [], "", None), + time_range=(0.0, 10.0), + children=[l3], + ) + l1 = L1Node( + id="l1_0", + card=L1Card("正常场景", "", [], [], [], [], ""), + time_range=(0.0, 10.0), + children=[l2], + ) + index = TreeIndex(metadata=IndexMeta("/t.mp4", "video"), roots=[l1]) + stats = asyncio.run(repair_tree(index, [], MockVLM(), MockLLM(), Path("/tmp"))) + assert stats.l3_repaired == 0 + assert stats.l2_regenerated == 0 + assert stats.l1_regenerated == 0 + + def test_stats_dataclass(self) -> None: + """RepairStats 数据类字段验证。""" + stats = RepairStats(l3_repaired=2, l2_regenerated=1, l1_regenerated=1) + assert stats.l3_repaired == 2 + assert stats.l2_regenerated == 1 + assert stats.l1_regenerated == 1 + + def test_skips_non_empty_field_issues(self, tmp_path: Path) -> None: + """非 empty_field 类型的 issue 不触发 L3 修复。""" + l3 = L3Node( + id="l1_0_l2_0_l3_0", + card=L3Card("正常描述", [], [], [], "", {}), + timestamp=1.0, + ) + l2 = L2Node( + id="l1_0_l2_0", + card=L2Card("原始事件", [], [], [], [], "", None), + time_range=(0.0, 10.0), + children=[l3], + ) + l1 = L1Node( + id="l1_0", + card=L1Card("原始场景", "", [], [], [], [], ""), + time_range=(0.0, 10.0), + children=[l2], + ) + index = TreeIndex(metadata=IndexMeta("/t.mp4", "video"), roots=[l1]) + issues = [NodeIssue("l1_0_l2_0_l3_0", 3, "missing_frame", "帧文件不存在")] + stats = asyncio.run(repair_tree(index, issues, MockVLM(), MockLLM(), tmp_path)) + assert stats.l3_repaired == 0 + assert stats.l2_regenerated == 0 + + def test_multiple_l3_under_same_l2(self, tmp_path: Path) -> None: + """同一 L2 下多个 L3 修复后,L2 只重生成一次。""" + frame_dir = tmp_path / "frames" + frame_dir.mkdir(parents=True) + for i in range(2): + (frame_dir / f"f{i}.jpg").write_bytes(b"\xff\xd8\xff\xe0fake") + + l3_a = L3Node( + id="l1_0_l2_0_l3_0", + card=L3Card("", [], [], [], "", {}), + timestamp=1.0, + frame_path="frames/f0.jpg", + ) + l3_b = L3Node( + id="l1_0_l2_0_l3_1", + card=L3Card("", [], [], [], "", {}), + timestamp=2.0, + frame_path="frames/f1.jpg", + ) + l2 = L2Node( + id="l1_0_l2_0", + card=L2Card("原始事件", [], [], [], [], "", None), + time_range=(0.0, 10.0), + children=[l3_a, l3_b], + ) + l1 = L1Node( + id="l1_0", + card=L1Card("原始场景", "", [], [], [], [], ""), + time_range=(0.0, 10.0), + children=[l2], + ) + index = TreeIndex(metadata=IndexMeta("/t.mp4", "video"), roots=[l1]) + issues = [ + NodeIssue("l1_0_l2_0_l3_0", 3, "empty_field", "frame_summary 为空"), + NodeIssue("l1_0_l2_0_l3_1", 3, "empty_field", "frame_summary 为空"), + ] + vlm = MockVLM() + llm = MockLLM() + stats = asyncio.run(repair_tree(index, issues, vlm, llm, tmp_path)) + assert stats.l3_repaired == 2 + assert stats.l2_regenerated == 1 + assert stats.l1_regenerated == 1 + assert vlm.call_count == 2 + # LLM 应被调用 2 次:一次 L2 + 一次 L1 + assert llm.call_count == 2 + + def test_missing_frame_file_skips_l3(self, tmp_path: Path) -> None: + """帧文件不存在时跳过该 L3 节点的修复。""" + l3 = L3Node( + id="l1_0_l2_0_l3_0", + card=L3Card("", [], [], [], "", {}), + timestamp=1.0, + frame_path="frames/nonexistent.jpg", + ) + l2 = L2Node( + id="l1_0_l2_0", + card=L2Card("原始事件", [], [], [], [], "", None), + time_range=(0.0, 10.0), + children=[l3], + ) + l1 = L1Node( + id="l1_0", + card=L1Card("原始场景", "", [], [], [], [], ""), + time_range=(0.0, 10.0), + children=[l2], + ) + index = TreeIndex(metadata=IndexMeta("/t.mp4", "video"), roots=[l1]) + issues = [NodeIssue("l1_0_l2_0_l3_0", 3, "empty_field", "frame_summary 为空")] + stats = asyncio.run(repair_tree(index, issues, MockVLM(), MockLLM(), tmp_path)) + # 帧文件不存在 → 跳过 L3 修复 → 无级联 + assert stats.l3_repaired == 0 + assert stats.l2_regenerated == 0 + assert stats.l1_regenerated == 0 diff --git a/tools/migrate_from_trm4.sh b/tools/migrate_from_trm4.sh new file mode 100755 index 0000000..ec8abb3 --- /dev/null +++ b/tools/migrate_from_trm4.sh @@ -0,0 +1,75 @@ +#!/usr/bin/env bash +# 从 TRM4.zip 迁移资产到 TRM5 +# 用法: bash tools/migrate_from_trm4.sh /path/to/Video-Tree-TRM4.zip +set -euo pipefail + +ZIP_PATH="${1:?用法: bash tools/migrate_from_trm4.sh /path/to/Video-Tree-TRM4.zip}" +PROJECT_ROOT="$(cd "$(dirname "$0")/.." && pwd)" +TMP_DIR=$(mktemp -d) + +echo "=== TRM4 -> TRM5 迁移 ===" +echo "ZIP: $ZIP_PATH" +echo "项目根: $PROJECT_ROOT" +echo "临时目录: $TMP_DIR" + +# 1. 解压 +echo "[1/6] 解压 TRM4.zip..." +unzip -q "$ZIP_PATH" -d "$TMP_DIR" +SRC="$TMP_DIR/Video-Tree-TRM4" + +# 2. 拷贝帧文件 (rsync --ignore-existing) +echo "[2/6] 拷贝帧文件..." +mkdir -p "$PROJECT_ROOT/store/videos" +for vid_dir in "$SRC"/store/videos/*/; do + vid=$(basename "$vid_dir") + dst="$PROJECT_ROOT/store/videos/$vid" + mkdir -p "$dst" + if [ -d "$vid_dir/frames" ]; then + rsync -a --ignore-existing "$vid_dir/frames/" "$dst/frames/" + fi +done + +# 3. 拷贝 SRT 字幕 +echo "[3/6] 拷贝 SRT 字幕..." +mkdir -p "$PROJECT_ROOT/data/Video-MME/subtitle" +if [ -d "$SRC/data/Video-MME/subtitle" ]; then + rsync -a --ignore-existing "$SRC/data/Video-MME/subtitle/" "$PROJECT_ROOT/data/Video-MME/subtitle/" +fi + +# 4. 拷贝视频压缩包 +echo "[4/6] 拷贝视频压缩包..." +mkdir -p "$PROJECT_ROOT/data/Video-MME/original_data" +if [ -d "$SRC/data/Video-MME/original_data" ]; then + rsync -a --ignore-existing "$SRC/data/Video-MME/original_data/" "$PROJECT_ROOT/data/Video-MME/original_data/" +fi + +# 5. 拷贝问题 JSON +echo "[5/6] 拷贝 Benchmark 问题..." +mkdir -p "$PROJECT_ROOT/store/questions" +if [ -d "$SRC/store/questions" ]; then + rsync -a --ignore-existing "$SRC/store/questions/" "$PROJECT_ROOT/store/questions/" +fi + +# 6. 格式转换 +echo "[6/6] 格式转换 flat -> TreeIndex..." +conda run -n Video-Tree-TRM python "$PROJECT_ROOT/tools/convert_flat_to_treeindex.py" \ + "$SRC/store/videos" "$PROJECT_ROOT/store/videos" + +# 验收 +echo "" +echo "=== 验收检查 ===" +VIDEO_COUNT=$(find "$PROJECT_ROOT/store/videos" -name "tree.json" | wc -l) +SRT_COUNT=$(find "$PROJECT_ROOT/data/Video-MME/subtitle" -name "*.srt" 2>/dev/null | wc -l) +echo "视频树: $VIDEO_COUNT (期望 300)" +echo "SRT 字幕: $SRT_COUNT (期望 >=290)" + +# 清理 +echo "清理临时目录..." +rm -rf "$TMP_DIR" + +if [ "$VIDEO_COUNT" -lt 300 ]; then + echo "WARNING: 视频树数量不足 300,请检查" + exit 1 +fi + +echo "迁移完成"