refactor(tree): subtitle 迁入 L3Card/L2Card + 建树管线修正
- L3Card/L2Card 新增 subtitle: str 字段(L1Card 不加) - L3Node 移除 subtitle 字段(数据迁入 Card) - assign_subtitles_voronoi 改写 Card.subtitle + L2 聚合 - _collect_card_strings 增加 skip_fields 排除 subtitle - _node_full_text/_node_anchored_text 保持 字幕:/[sN] 语义 - get_subtitle 读 Card.subtitle(L2/L3) - verify.py/synthesizer.py: l3.subtitle → l3.card.subtitle - 迁移脚本 tools/migrate_subtitle_to_card.py(幂等,300 棵树已迁移) - 9→6 个测试文件适配(3 个无需改动) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -196,7 +196,7 @@ def _sample_l3(
|
|||||||
# Phase 3: 构造上下文(frame_path 已在候选过滤中保证非 None)
|
# Phase 3: 构造上下文(frame_path 已在候选过滤中保证非 None)
|
||||||
card_text = _serialize_l3_card(chosen_l3, spec.context_fields)
|
card_text = _serialize_l3_card(chosen_l3, spec.context_fields)
|
||||||
frame_paths = [chosen_l3.frame_path] # type: ignore[list-item]
|
frame_paths = [chosen_l3.frame_path] # type: ignore[list-item]
|
||||||
subtitle = chosen_l3.subtitle or ""
|
subtitle = chosen_l3.card.subtitle or ""
|
||||||
|
|
||||||
# Phase 4: 干扰项——整棵树中其他 L3 的 frame_summary
|
# Phase 4: 干扰项——整棵树中其他 L3 的 frame_summary
|
||||||
distractor_texts = [
|
distractor_texts = [
|
||||||
@@ -281,10 +281,10 @@ def _sample_l2(
|
|||||||
if is_temporal_perception:
|
if is_temporal_perception:
|
||||||
card_text += f"\n{_l2_time_range_str(chosen_l2)}"
|
card_text += f"\n{_l2_time_range_str(chosen_l2)}"
|
||||||
|
|
||||||
# Phase 5: 字幕(L2 无自身字幕,取首个子 L3 字幕)
|
# Phase 5: 字幕(优先使用 L2 自身字幕,否则取首个子 L3 字幕)
|
||||||
subtitle = ""
|
subtitle = chosen_l2.card.subtitle or ""
|
||||||
if chosen_l2.children and chosen_l2.children[0].subtitle:
|
if not subtitle and chosen_l2.children and chosen_l2.children[0].card.subtitle:
|
||||||
subtitle = chosen_l2.children[0].subtitle
|
subtitle = chosen_l2.children[0].card.subtitle
|
||||||
|
|
||||||
# Phase 6: 干扰项——整棵树中其他 L2 的 event_description
|
# Phase 6: 干扰项——整棵树中其他 L2 的 event_description
|
||||||
distractor_texts = [
|
distractor_texts = [
|
||||||
|
|||||||
+37
-15
@@ -64,26 +64,41 @@ def _node_description(node: AnyNode) -> str:
|
|||||||
return node.card.frame_summary
|
return node.card.frame_summary
|
||||||
|
|
||||||
|
|
||||||
def _collect_card_strings(node: AnyNode) -> list[str]:
|
def _collect_card_strings(
|
||||||
|
node: AnyNode,
|
||||||
|
skip_fields: frozenset[str] = frozenset(),
|
||||||
|
) -> list[str]:
|
||||||
"""从节点 card 中递归收集所有非空字符串字段。
|
"""从节点 card 中递归收集所有非空字符串字段。
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
node: 树节点实例。
|
node: 树节点实例。
|
||||||
|
skip_fields: 需要跳过的 dataclass 字段名集合(如 subtitle,
|
||||||
|
因为它需要单独添加"字幕:"标签和 [sN] 锚标)。
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
字符串列表(每个非空字段值一项,含内嵌换行的按行拆分)。
|
字符串列表(每个非空字段值一项,含内嵌换行的按行拆分)。
|
||||||
"""
|
"""
|
||||||
result: list[str] = []
|
result: list[str] = []
|
||||||
_collect_from_obj(node.card, result)
|
_collect_from_obj(node.card, result, skip_fields=skip_fields)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _collect_from_obj(obj: object, out: list[str]) -> None:
|
# subtitle 字段在 _node_full_text / _node_anchored_text 中单独处理
|
||||||
|
_SUBTITLE_SKIP: frozenset[str] = frozenset({"subtitle"})
|
||||||
|
|
||||||
|
|
||||||
|
def _collect_from_obj(
|
||||||
|
obj: object,
|
||||||
|
out: list[str],
|
||||||
|
*,
|
||||||
|
skip_fields: frozenset[str] = frozenset(),
|
||||||
|
) -> None:
|
||||||
"""递归收集任意嵌套结构中的非空字符串。
|
"""递归收集任意嵌套结构中的非空字符串。
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
obj: dict / list / str / 其他。
|
obj: dict / list / str / 其他。
|
||||||
out: 收集结果列表(原地修改)。
|
out: 收集结果列表(原地修改)。
|
||||||
|
skip_fields: 需要跳过的 dataclass 字段名集合。
|
||||||
"""
|
"""
|
||||||
if isinstance(obj, str):
|
if isinstance(obj, str):
|
||||||
stripped = obj.strip()
|
stripped = obj.strip()
|
||||||
@@ -91,14 +106,16 @@ def _collect_from_obj(obj: object, out: list[str]) -> None:
|
|||||||
out.append(stripped)
|
out.append(stripped)
|
||||||
elif isinstance(obj, dict):
|
elif isinstance(obj, dict):
|
||||||
for v in obj.values():
|
for v in obj.values():
|
||||||
_collect_from_obj(v, out)
|
_collect_from_obj(v, out, skip_fields=skip_fields)
|
||||||
elif isinstance(obj, (list, tuple)):
|
elif isinstance(obj, (list, tuple)):
|
||||||
for item in obj:
|
for item in obj:
|
||||||
_collect_from_obj(item, out)
|
_collect_from_obj(item, out, skip_fields=skip_fields)
|
||||||
elif hasattr(obj, "__dataclass_fields__"):
|
elif hasattr(obj, "__dataclass_fields__"):
|
||||||
# frozen dataclass(Card 类型)
|
# frozen dataclass(Card 类型)
|
||||||
for field_name in obj.__dataclass_fields__:
|
for field_name in obj.__dataclass_fields__:
|
||||||
_collect_from_obj(getattr(obj, field_name), out)
|
if field_name in skip_fields:
|
||||||
|
continue
|
||||||
|
_collect_from_obj(getattr(obj, field_name), out, skip_fields=skip_fields)
|
||||||
|
|
||||||
|
|
||||||
class TreeEnvironment:
|
class TreeEnvironment:
|
||||||
@@ -367,17 +384,19 @@ class TreeEnvironment:
|
|||||||
def get_subtitle(self, node_id: str) -> str:
|
def get_subtitle(self, node_id: str) -> str:
|
||||||
"""返回节点字幕文本。
|
"""返回节点字幕文本。
|
||||||
|
|
||||||
|
L2/L3 节点从 card.subtitle 读取,L1 节点不含字幕。
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
node_id: 节点 ID。
|
node_id: 节点 ID。
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
字幕文本;无字幕或节点不存在时返回空字符串。
|
字幕文本;无字幕、L1 节点或节点不存在时返回空字符串。
|
||||||
"""
|
"""
|
||||||
node = self._id_to_node.get(node_id)
|
node = self._id_to_node.get(node_id)
|
||||||
if node is None:
|
if node is None:
|
||||||
return ""
|
return ""
|
||||||
if isinstance(node, L3Node):
|
if isinstance(node, (L2Node, L3Node)):
|
||||||
return node.subtitle or ""
|
return node.card.subtitle or ""
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
def resolve_frame_paths(self, node_ids: list[str]) -> list[Path]:
|
def resolve_frame_paths(self, node_ids: list[str]) -> list[Path]:
|
||||||
@@ -448,22 +467,25 @@ class TreeEnvironment:
|
|||||||
def _node_full_text(self, node: AnyNode) -> str:
|
def _node_full_text(self, node: AnyNode) -> str:
|
||||||
"""获取节点完整文本(card 所有字段 + subtitle)。
|
"""获取节点完整文本(card 所有字段 + subtitle)。
|
||||||
|
|
||||||
|
subtitle 从 card.subtitle 读取,仅 L2/L3 节点附加"字幕:"标签。
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
node: 树节点。
|
node: 树节点。
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
拼接后的全文本。
|
拼接后的全文本。
|
||||||
"""
|
"""
|
||||||
card_strings = _collect_card_strings(node)
|
card_strings = _collect_card_strings(node, skip_fields=_SUBTITLE_SKIP)
|
||||||
text = "\n".join(card_strings)
|
text = "\n".join(card_strings)
|
||||||
if isinstance(node, L3Node) and node.subtitle:
|
if isinstance(node, (L2Node, L3Node)) and node.card.subtitle:
|
||||||
text += f"\n字幕: {node.subtitle}"
|
text += f"\n字幕: {node.card.subtitle}"
|
||||||
return text
|
return text
|
||||||
|
|
||||||
def _node_anchored_text(self, node: AnyNode) -> str:
|
def _node_anchored_text(self, node: AnyNode) -> str:
|
||||||
"""获取带行号锚的节点文本。
|
"""获取带行号锚的节点文本。
|
||||||
|
|
||||||
card 字符串逐行编 [c1]..[cN],字幕逐行编 [s1]..[sM]。
|
card 字符串逐行编 [c1]..[cN],字幕逐行编 [s1]..[sM]。
|
||||||
|
字幕从 card.subtitle 读取,仅 L2/L3 节点产生 [sN] 锚标。
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
node: 树节点。
|
node: 树节点。
|
||||||
@@ -471,15 +493,15 @@ class TreeEnvironment:
|
|||||||
返回:
|
返回:
|
||||||
带锚文本。
|
带锚文本。
|
||||||
"""
|
"""
|
||||||
card_strings = _collect_card_strings(node)
|
card_strings = _collect_card_strings(node, skip_fields=_SUBTITLE_SKIP)
|
||||||
# 拆分内嵌换行,确保一锚一行
|
# 拆分内嵌换行,确保一锚一行
|
||||||
card_lines: list[str] = []
|
card_lines: list[str] = []
|
||||||
for s in card_strings:
|
for s in card_strings:
|
||||||
card_lines.extend(ln for ln in s.splitlines() if ln.strip())
|
card_lines.extend(ln for ln in s.splitlines() if ln.strip())
|
||||||
|
|
||||||
sub_lines: list[str] = []
|
sub_lines: list[str] = []
|
||||||
if isinstance(node, L3Node) and node.subtitle:
|
if isinstance(node, (L2Node, L3Node)) and node.card.subtitle:
|
||||||
sub_lines = [ln for ln in node.subtitle.splitlines() if ln.strip()]
|
sub_lines = [ln for ln in node.card.subtitle.splitlines() if ln.strip()]
|
||||||
|
|
||||||
anchored: list[str] = []
|
anchored: list[str] = []
|
||||||
for i, line in enumerate(card_lines, 1):
|
for i, line in enumerate(card_lines, 1):
|
||||||
|
|||||||
+10
-4
@@ -82,6 +82,7 @@ class L3Card:
|
|||||||
visible_text: 画面中可见的文字列表。
|
visible_text: 画面中可见的文字列表。
|
||||||
spatial_layout: 空间布局描述。
|
spatial_layout: 空间布局描述。
|
||||||
visual_attributes: 视觉属性字典(如光照、色调等)。
|
visual_attributes: 视觉属性字典(如光照、色调等)。
|
||||||
|
subtitle: 字幕文本(Voronoi 分配后填充,默认空)。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
frame_summary: str
|
frame_summary: str
|
||||||
@@ -90,6 +91,7 @@ class L3Card:
|
|||||||
visible_text: list[str]
|
visible_text: list[str]
|
||||||
spatial_layout: str
|
spatial_layout: str
|
||||||
visual_attributes: dict[str, Any]
|
visual_attributes: dict[str, Any]
|
||||||
|
subtitle: str = ""
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -106,6 +108,7 @@ class L2Card:
|
|||||||
visible_text: 片段中可见的文字列表。
|
visible_text: 片段中可见的文字列表。
|
||||||
spatial_relations: 空间关系描述。
|
spatial_relations: 空间关系描述。
|
||||||
state_changes: 状态变化描述(可选)。
|
state_changes: 状态变化描述(可选)。
|
||||||
|
subtitle: 子 L3 字幕聚合文本(Voronoi 分配后填充,默认空)。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
event_description: str
|
event_description: str
|
||||||
@@ -115,6 +118,7 @@ class L2Card:
|
|||||||
visible_text: list[str]
|
visible_text: list[str]
|
||||||
spatial_relations: str
|
spatial_relations: str
|
||||||
state_changes: str | None
|
state_changes: str | None
|
||||||
|
subtitle: str = ""
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -183,7 +187,6 @@ class L3Node:
|
|||||||
embedding: 文本嵌入向量,形状 [D],float32。
|
embedding: 文本嵌入向量,形状 [D],float32。
|
||||||
timestamp: 对应的时间戳(秒,可选)。
|
timestamp: 对应的时间戳(秒,可选)。
|
||||||
frame_path: 关联的帧图像路径(可选,仅视频模态)。
|
frame_path: 关联的帧图像路径(可选,仅视频模态)。
|
||||||
subtitle: 该帧对应的字幕文本(可选)。
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
id: str
|
id: str
|
||||||
@@ -191,7 +194,6 @@ class L3Node:
|
|||||||
embedding: np.ndarray | None = None
|
embedding: np.ndarray | None = None
|
||||||
timestamp: float | None = None
|
timestamp: float | None = None
|
||||||
frame_path: str | None = None
|
frame_path: str | None = None
|
||||||
subtitle: str | None = None
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def description(self) -> str:
|
def description(self) -> str:
|
||||||
@@ -274,10 +276,10 @@ class L1Node:
|
|||||||
"visible_text": n.card.visible_text,
|
"visible_text": n.card.visible_text,
|
||||||
"spatial_layout": n.card.spatial_layout,
|
"spatial_layout": n.card.spatial_layout,
|
||||||
"visual_attributes": n.card.visual_attributes,
|
"visual_attributes": n.card.visual_attributes,
|
||||||
|
"subtitle": n.card.subtitle,
|
||||||
},
|
},
|
||||||
"timestamp": n.timestamp,
|
"timestamp": n.timestamp,
|
||||||
"frame_path": n.frame_path,
|
"frame_path": n.frame_path,
|
||||||
"subtitle": n.subtitle,
|
|
||||||
}
|
}
|
||||||
if include_embedding:
|
if include_embedding:
|
||||||
d["embedding"] = _embed_to_str(n.embedding)
|
d["embedding"] = _embed_to_str(n.embedding)
|
||||||
@@ -294,6 +296,7 @@ class L1Node:
|
|||||||
"visible_text": n.card.visible_text,
|
"visible_text": n.card.visible_text,
|
||||||
"spatial_relations": n.card.spatial_relations,
|
"spatial_relations": n.card.spatial_relations,
|
||||||
"state_changes": n.card.state_changes,
|
"state_changes": n.card.state_changes,
|
||||||
|
"subtitle": n.card.subtitle,
|
||||||
},
|
},
|
||||||
"time_range": list(n.time_range) if n.time_range else None,
|
"time_range": list(n.time_range) if n.time_range else None,
|
||||||
"children": [l3_to_dict(c) for c in n.children],
|
"children": [l3_to_dict(c) for c in n.children],
|
||||||
@@ -334,6 +337,8 @@ class L1Node:
|
|||||||
for l2d in d.get("children", []):
|
for l2d in d.get("children", []):
|
||||||
l3_nodes: list[L3Node] = []
|
l3_nodes: list[L3Node] = []
|
||||||
for l3d in l2d.get("children", []):
|
for l3d in l2d.get("children", []):
|
||||||
|
# 向后兼容:旧格式 subtitle 在节点级,新格式在 card 内
|
||||||
|
l3_subtitle = l3d["card"].get("subtitle", "") or l3d.get("subtitle", "") or ""
|
||||||
l3_card = L3Card(
|
l3_card = L3Card(
|
||||||
frame_summary=l3d["card"]["frame_summary"],
|
frame_summary=l3d["card"]["frame_summary"],
|
||||||
visible_entities=l3d["card"]["visible_entities"],
|
visible_entities=l3d["card"]["visible_entities"],
|
||||||
@@ -341,6 +346,7 @@ class L1Node:
|
|||||||
visible_text=l3d["card"]["visible_text"],
|
visible_text=l3d["card"]["visible_text"],
|
||||||
spatial_layout=l3d["card"]["spatial_layout"],
|
spatial_layout=l3d["card"]["spatial_layout"],
|
||||||
visual_attributes=l3d["card"]["visual_attributes"],
|
visual_attributes=l3d["card"]["visual_attributes"],
|
||||||
|
subtitle=l3_subtitle,
|
||||||
)
|
)
|
||||||
l3_nodes.append(
|
l3_nodes.append(
|
||||||
L3Node(
|
L3Node(
|
||||||
@@ -349,7 +355,6 @@ class L1Node:
|
|||||||
embedding=_embed_from_str(l3d.get("embedding")),
|
embedding=_embed_from_str(l3d.get("embedding")),
|
||||||
timestamp=l3d.get("timestamp"),
|
timestamp=l3d.get("timestamp"),
|
||||||
frame_path=l3d.get("frame_path"),
|
frame_path=l3d.get("frame_path"),
|
||||||
subtitle=l3d.get("subtitle"),
|
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
l2_card = L2Card(
|
l2_card = L2Card(
|
||||||
@@ -360,6 +365,7 @@ class L1Node:
|
|||||||
visible_text=l2d["card"]["visible_text"],
|
visible_text=l2d["card"]["visible_text"],
|
||||||
spatial_relations=l2d["card"]["spatial_relations"],
|
spatial_relations=l2d["card"]["spatial_relations"],
|
||||||
state_changes=l2d["card"]["state_changes"],
|
state_changes=l2d["card"]["state_changes"],
|
||||||
|
subtitle=l2d["card"].get("subtitle", ""),
|
||||||
)
|
)
|
||||||
tr2 = l2d.get("time_range")
|
tr2 = l2d.get("time_range")
|
||||||
l2_nodes.append(
|
l2_nodes.append(
|
||||||
|
|||||||
+14
-2
@@ -13,6 +13,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
import re
|
import re
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
@@ -254,7 +255,8 @@ def assign_subtitles_voronoi(
|
|||||||
entries: 已解析的 SRTEntry 列表。
|
entries: 已解析的 SRTEntry 列表。
|
||||||
|
|
||||||
副作用:
|
副作用:
|
||||||
直接修改每个 L3Node.subtitle 字段。
|
通过 dataclasses.replace 替换 L3Node.card 和 L2Node.card,
|
||||||
|
将字幕写入 card.subtitle 字段。
|
||||||
|
|
||||||
迁移来源:
|
迁移来源:
|
||||||
TRM3 tools/generate_subtitles.py compute_effective_ranges + assign_subtitles
|
TRM3 tools/generate_subtitles.py compute_effective_ranges + assign_subtitles
|
||||||
@@ -299,7 +301,17 @@ def assign_subtitles_voronoi(
|
|||||||
right = (ts + next_ts) / 2.0
|
right = (ts + next_ts) / 2.0
|
||||||
|
|
||||||
subtitle_text = extract_subtitle_for_range(entries, (left, right))
|
subtitle_text = extract_subtitle_for_range(entries, (left, right))
|
||||||
l3.subtitle = subtitle_text if subtitle_text else None
|
l3.card = dataclasses.replace(
|
||||||
|
l3.card,
|
||||||
|
subtitle=subtitle_text or "",
|
||||||
|
)
|
||||||
|
|
||||||
|
# L2 字幕聚合:拼接所有 L3 子节点的字幕
|
||||||
|
l3_subtitles = [l3.card.subtitle for l3 in l2.children if l3.card.subtitle]
|
||||||
|
l2.card = dataclasses.replace(
|
||||||
|
l2.card,
|
||||||
|
subtitle="\n".join(l3_subtitles),
|
||||||
|
)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Voronoi 字幕分配完成: {} 个 L1 节点, {} 条字幕条目",
|
"Voronoi 字幕分配完成: {} 个 L1 节点, {} 条字幕条目",
|
||||||
|
|||||||
+5
-4
@@ -97,8 +97,8 @@ def _collect_l3_text(l2_node: L2Node) -> str:
|
|||||||
for l3 in l2_node.children:
|
for l3 in l2_node.children:
|
||||||
parts.append(l3.card.frame_summary)
|
parts.append(l3.card.frame_summary)
|
||||||
parts.extend(l3.card.visible_text)
|
parts.extend(l3.card.visible_text)
|
||||||
if l3.subtitle:
|
if l3.card.subtitle:
|
||||||
parts.append(l3.subtitle)
|
parts.append(l3.card.subtitle)
|
||||||
return "\n".join(parts)
|
return "\n".join(parts)
|
||||||
|
|
||||||
|
|
||||||
@@ -139,8 +139,8 @@ def _collect_descendant_text_corpus(l1_node: L1Node) -> str:
|
|||||||
for l3 in l2.children:
|
for l3 in l2.children:
|
||||||
parts.append(l3.card.frame_summary)
|
parts.append(l3.card.frame_summary)
|
||||||
parts.extend(l3.card.visible_text)
|
parts.extend(l3.card.visible_text)
|
||||||
if l3.subtitle:
|
if l3.card.subtitle:
|
||||||
parts.append(l3.subtitle)
|
parts.append(l3.card.subtitle)
|
||||||
return "\n".join(parts)
|
return "\n".join(parts)
|
||||||
|
|
||||||
|
|
||||||
@@ -221,6 +221,7 @@ def _verify_l2(l2: L2Node, stats: VerifyStats) -> None:
|
|||||||
visible_text=kept_vt,
|
visible_text=kept_vt,
|
||||||
spatial_relations=old_card.spatial_relations,
|
spatial_relations=old_card.spatial_relations,
|
||||||
state_changes=old_card.state_changes,
|
state_changes=old_card.state_changes,
|
||||||
|
subtitle=old_card.subtitle,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -480,7 +480,7 @@ class VideoTreeBuilder:
|
|||||||
)
|
)
|
||||||
index = TreeIndex(metadata=metadata, roots=l1_nodes)
|
index = TreeIndex(metadata=metadata, roots=l1_nodes)
|
||||||
|
|
||||||
# Phase 7: 字幕 Voronoi 分配(可选)
|
# Phase 7: 字幕 Voronoi 分配到 L3/L2 Card.subtitle(可选)
|
||||||
if srt_entries:
|
if srt_entries:
|
||||||
assign_subtitles_voronoi(index, srt_entries)
|
assign_subtitles_voronoi(index, srt_entries)
|
||||||
logger.info("字幕 Voronoi 分配完成", n_entries=len(srt_entries))
|
logger.info("字幕 Voronoi 分配完成", n_entries=len(srt_entries))
|
||||||
|
|||||||
@@ -97,10 +97,10 @@ class TestTreeModuleE2E:
|
|||||||
SRTEntry(start=5.0, end=7.0, text="The crowd goes wild!"),
|
SRTEntry(start=5.0, end=7.0, text="The crowd goes wild!"),
|
||||||
]
|
]
|
||||||
assign_subtitles_voronoi(index, srt_entries)
|
assign_subtitles_voronoi(index, srt_entries)
|
||||||
assert l3_0.subtitle is not None
|
assert l3_0.card.subtitle != ""
|
||||||
assert "sprints" in l3_0.subtitle
|
assert "sprints" in l3_0.card.subtitle
|
||||||
assert l3_1.subtitle is not None
|
assert l3_1.card.subtitle != ""
|
||||||
assert "crowd" in l3_1.subtitle
|
assert "crowd" in l3_1.card.subtitle
|
||||||
|
|
||||||
# Step 3: TreeEnvironment 查询
|
# Step 3: TreeEnvironment 查询
|
||||||
env = TreeEnvironment(index)
|
env = TreeEnvironment(index)
|
||||||
@@ -139,8 +139,8 @@ class TestTreeModuleE2E:
|
|||||||
|
|
||||||
assert len(loaded.roots) == 1
|
assert len(loaded.roots) == 1
|
||||||
assert loaded.roots[0].card.scene_summary == "百米短跑决赛"
|
assert loaded.roots[0].card.scene_summary == "百米短跑决赛"
|
||||||
assert loaded.roots[0].children[0].children[0].subtitle is not None
|
assert loaded.roots[0].children[0].children[0].card.subtitle != ""
|
||||||
assert "sprints" in loaded.roots[0].children[0].children[0].subtitle
|
assert "sprints" in loaded.roots[0].children[0].children[0].card.subtitle
|
||||||
# verify 的修改也被保留
|
# verify 的修改也被保留
|
||||||
assert "幻觉实体XYZ" not in loaded.roots[0].children[0].card.entities
|
assert "幻觉实体XYZ" not in loaded.roots[0].children[0].card.entities
|
||||||
|
|
||||||
|
|||||||
@@ -111,10 +111,10 @@ def _make_test_tree() -> TreeIndex:
|
|||||||
visible_text=[],
|
visible_text=[],
|
||||||
spatial_layout="center",
|
spatial_layout="center",
|
||||||
visual_attributes={},
|
visual_attributes={},
|
||||||
|
subtitle="test subtitle text",
|
||||||
),
|
),
|
||||||
timestamp=10.0,
|
timestamp=10.0,
|
||||||
frame_path="frames/L1_000_L2_000_L3_000.jpg",
|
frame_path="frames/L1_000_L2_000_L3_000.jpg",
|
||||||
subtitle="test subtitle text",
|
|
||||||
)
|
)
|
||||||
l2 = L2Node(
|
l2 = L2Node(
|
||||||
id="vid_L1_000_L2_000",
|
id="vid_L1_000_L2_000",
|
||||||
|
|||||||
@@ -112,7 +112,11 @@ class TestVoronoiAssign:
|
|||||||
entries = [SRTEntry(1.0, 3.0, "hello"), SRTEntry(5.0, 7.0, "world")]
|
entries = [SRTEntry(1.0, 3.0, "hello"), SRTEntry(5.0, 7.0, "world")]
|
||||||
assign_subtitles_voronoi(index, entries)
|
assign_subtitles_voronoi(index, entries)
|
||||||
|
|
||||||
assert l3_0.subtitle is not None
|
assert l3_0.card.subtitle != ""
|
||||||
assert "hello" in l3_0.subtitle
|
assert "hello" in l3_0.card.subtitle
|
||||||
assert l3_1.subtitle is not None
|
assert l3_1.card.subtitle != ""
|
||||||
assert "world" in l3_1.subtitle
|
assert "world" in l3_1.card.subtitle
|
||||||
|
# L2 字幕应聚合 L3 子节点字幕
|
||||||
|
assert l2.card.subtitle != ""
|
||||||
|
assert "hello" in l2.card.subtitle
|
||||||
|
assert "world" in l2.card.subtitle
|
||||||
|
|||||||
@@ -31,10 +31,10 @@ def _make_test_index() -> TreeIndex:
|
|||||||
["Nike"],
|
["Nike"],
|
||||||
"居中",
|
"居中",
|
||||||
{"lighting": "明亮"},
|
{"lighting": "明亮"},
|
||||||
|
subtitle="he is running",
|
||||||
),
|
),
|
||||||
timestamp=1.0,
|
timestamp=1.0,
|
||||||
frame_path="frames/L1_000_L2_000_L3_000.jpg",
|
frame_path="frames/L1_000_L2_000_L3_000.jpg",
|
||||||
subtitle="he is running",
|
|
||||||
)
|
)
|
||||||
l3_1 = L3Node(
|
l3_1 = L3Node(
|
||||||
id="vid_L1_000_L2_000_L3_001",
|
id="vid_L1_000_L2_000_L3_001",
|
||||||
|
|||||||
@@ -130,9 +130,9 @@ class TestNodes:
|
|||||||
node = _make_l3()
|
node = _make_l3()
|
||||||
assert node.embedding is None
|
assert node.embedding is None
|
||||||
|
|
||||||
def test_l3_subtitle_default_none(self):
|
def test_l3_card_subtitle_default_empty(self):
|
||||||
node = _make_l3()
|
node = _make_l3()
|
||||||
assert node.subtitle is None
|
assert node.card.subtitle == ""
|
||||||
|
|
||||||
|
|
||||||
class TestTreeIndex:
|
class TestTreeIndex:
|
||||||
|
|||||||
@@ -55,9 +55,9 @@ class TestVerifyTree:
|
|||||||
visible_text=["Nike", "2024"],
|
visible_text=["Nike", "2024"],
|
||||||
spatial_layout="居中",
|
spatial_layout="居中",
|
||||||
visual_attributes={},
|
visual_attributes={},
|
||||||
|
subtitle="the athlete is running fast",
|
||||||
),
|
),
|
||||||
timestamp=1.0,
|
timestamp=1.0,
|
||||||
subtitle="the athlete is running fast",
|
|
||||||
)
|
)
|
||||||
l3_1 = L3Node(
|
l3_1 = L3Node(
|
||||||
id="l1_0_l2_0_l3_1",
|
id="l1_0_l2_0_l3_1",
|
||||||
|
|||||||
@@ -0,0 +1,162 @@
|
|||||||
|
"""字幕数据迁移脚本:L3Node.subtitle → L3Card.subtitle + L2 聚合。
|
||||||
|
|
||||||
|
幂等操作——已迁移的树文件不会被重复修改。
|
||||||
|
|
||||||
|
用法:
|
||||||
|
conda activate Video-Tree-TRM & python tools/migrate_subtitle_to_card.py [--dry-run]
|
||||||
|
|
||||||
|
迁移逻辑:
|
||||||
|
1. 扫描 store/videos/*/tree.json
|
||||||
|
2. 对每棵树:
|
||||||
|
- L3: 将节点级 subtitle 移入 card["subtitle"],删除节点级 key
|
||||||
|
- L2: 聚合子 L3 的 card["subtitle"] 写入 L2 card["subtitle"]
|
||||||
|
- 处理边界: None → ""、空字符串、已迁移(card 内已有 subtitle)
|
||||||
|
3. 回写 JSON
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
|
||||||
|
def _migrate_tree_dict(tree_dict: dict) -> bool:
|
||||||
|
"""原地迁移单棵树的 dict 结构,返回是否发生了修改。
|
||||||
|
|
||||||
|
参数:
|
||||||
|
tree_dict: tree.json 加载后的字典。
|
||||||
|
|
||||||
|
返回:
|
||||||
|
True 表示发生了修改,需要回写。
|
||||||
|
"""
|
||||||
|
modified = False
|
||||||
|
|
||||||
|
for l1 in tree_dict.get("roots", []):
|
||||||
|
for l2 in l1.get("children", []):
|
||||||
|
l3_subtitles: list[str] = []
|
||||||
|
|
||||||
|
for l3 in l2.get("children", []):
|
||||||
|
card = l3.get("card", {})
|
||||||
|
|
||||||
|
# Phase 1: 迁移 L3 节点级 subtitle → card.subtitle
|
||||||
|
node_subtitle = l3.get("subtitle") # 旧位置
|
||||||
|
card_subtitle = card.get("subtitle", "") # 新位置
|
||||||
|
|
||||||
|
if node_subtitle is not None and not card_subtitle:
|
||||||
|
# 旧格式:节点级有值,card 内无值 → 迁移
|
||||||
|
card["subtitle"] = node_subtitle if node_subtitle else ""
|
||||||
|
modified = True
|
||||||
|
elif "subtitle" not in card:
|
||||||
|
# 既无节点级也无 card 级 → 初始化为空
|
||||||
|
card["subtitle"] = ""
|
||||||
|
modified = True
|
||||||
|
|
||||||
|
# Phase 2: 清理节点级 subtitle key
|
||||||
|
if "subtitle" in l3:
|
||||||
|
del l3["subtitle"]
|
||||||
|
modified = True
|
||||||
|
|
||||||
|
# 收集 L3 字幕用于 L2 聚合
|
||||||
|
final_subtitle = card.get("subtitle", "")
|
||||||
|
if final_subtitle:
|
||||||
|
l3_subtitles.append(final_subtitle)
|
||||||
|
|
||||||
|
# Phase 3: L2 字幕聚合
|
||||||
|
l2_card = l2.get("card", {})
|
||||||
|
existing_l2_subtitle = l2_card.get("subtitle", "")
|
||||||
|
aggregated = "\n".join(l3_subtitles)
|
||||||
|
|
||||||
|
if "subtitle" not in l2_card or existing_l2_subtitle != aggregated:
|
||||||
|
l2_card["subtitle"] = aggregated
|
||||||
|
modified = True
|
||||||
|
|
||||||
|
return modified
|
||||||
|
|
||||||
|
|
||||||
|
def migrate_all(
|
||||||
|
videos_dir: Path,
|
||||||
|
*,
|
||||||
|
dry_run: bool = False,
|
||||||
|
) -> tuple[int, int, int]:
|
||||||
|
"""扫描并迁移所有树文件。
|
||||||
|
|
||||||
|
参数:
|
||||||
|
videos_dir: store/videos 目录路径。
|
||||||
|
dry_run: 仅检查不写入。
|
||||||
|
|
||||||
|
返回:
|
||||||
|
(total, migrated, skipped) 计数元组。
|
||||||
|
"""
|
||||||
|
tree_files = sorted(videos_dir.glob("*/tree.json"))
|
||||||
|
total = len(tree_files)
|
||||||
|
migrated = 0
|
||||||
|
skipped = 0
|
||||||
|
|
||||||
|
for tree_path in tree_files:
|
||||||
|
video_id = tree_path.parent.name
|
||||||
|
try:
|
||||||
|
with open(tree_path, encoding="utf-8") as f:
|
||||||
|
tree_dict = json.load(f)
|
||||||
|
except (json.JSONDecodeError, OSError) as exc:
|
||||||
|
logger.warning("跳过损坏文件: {} ({})", tree_path, exc)
|
||||||
|
skipped += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
changed = _migrate_tree_dict(tree_dict)
|
||||||
|
|
||||||
|
if changed:
|
||||||
|
if dry_run:
|
||||||
|
logger.info("[dry-run] 需要迁移: {}", video_id)
|
||||||
|
else:
|
||||||
|
with open(tree_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(tree_dict, f, ensure_ascii=False, indent=2)
|
||||||
|
logger.info("已迁移: {}", video_id)
|
||||||
|
migrated += 1
|
||||||
|
else:
|
||||||
|
logger.debug("无需迁移: {}", video_id)
|
||||||
|
|
||||||
|
return total, migrated, skipped
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
"""CLI 入口。"""
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="将 L3Node.subtitle 迁移到 L3Card.subtitle + L2 聚合",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dry-run",
|
||||||
|
action="store_true",
|
||||||
|
help="仅检查不写入",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--videos-dir",
|
||||||
|
type=Path,
|
||||||
|
default=Path("store/videos"),
|
||||||
|
help="视频树目录 (默认: store/videos)",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if not args.videos_dir.is_dir():
|
||||||
|
logger.error("目录不存在: {}", args.videos_dir)
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
total, migrated, skipped = migrate_all(
|
||||||
|
args.videos_dir,
|
||||||
|
dry_run=args.dry_run,
|
||||||
|
)
|
||||||
|
action = "需要迁移" if args.dry_run else "已迁移"
|
||||||
|
logger.info(
|
||||||
|
"迁移完成: 总计={}, {}={}, 跳过={}",
|
||||||
|
total,
|
||||||
|
action,
|
||||||
|
migrated,
|
||||||
|
skipped,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
Reference in New Issue
Block a user