Files
Video-Tree-TRM5/tests/unit/test_sampler_v2.py
T
2026-07-14 05:38:14 -04:00

366 lines
12 KiB
Python
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.
"""v2 素材采样器单元测试。
测试 sample_material_v2 及其辅助函数的核心行为:
- 正常采样返回 MaterialContext
- 已用节点排除
- 约束违反时重试直至 RuntimeError
- 跨 L2 上下文收集
- 字幕收集
"""
from __future__ import annotations
import random
import pytest
from app.question_gen.families import (
REASONING_FAMILY,
RETRIEVAL_FAMILY,
VISUAL_FAMILY,
SamplingConstraint,
)
from app.tree.index import (
IndexMeta,
L1Card,
L1Node,
L2Card,
L2Node,
L3Card,
L3Node,
TreeIndex,
)
# ---------------------------------------------------------------------------
# Fixture: 构建含丰富数据的真实树结构
# ---------------------------------------------------------------------------
def _make_l3(
l1_idx: int,
l2_idx: int,
l3_idx: int,
*,
subtitle: str = "",
frame_path: str | None = None,
) -> L3Node:
"""构建 L3 节点,带可控 subtitle/frame_path。"""
node_id = f"l1_{l1_idx}_l2_{l2_idx}_l3_{l3_idx}"
return L3Node(
id=node_id,
card=L3Card(
frame_summary=f"帧{l3_idx}描述:L1={l1_idx},L2={l2_idx}",
visible_entities=[f"实体_{l3_idx}"],
ongoing_actions=[f"动作_{l3_idx}"],
visible_text=[],
spatial_layout="居中",
visual_attributes={"lighting": "明亮"},
subtitle=subtitle,
),
timestamp=float(l3_idx * 2),
frame_path=frame_path,
)
def _make_l2(
l1_idx: int,
l2_idx: int,
n_l3: int = 3,
*,
subtitle: str = "",
with_frames: bool = True,
with_subtitles: bool = True,
) -> L2Node:
"""构建 L2 节点,可控子节点数量和属性。"""
children: list[L3Node] = []
for i in range(n_l3):
sub = f"字幕L1={l1_idx}_L2={l2_idx}_L3={i}" if with_subtitles else ""
fp = f"frames/l1_{l1_idx}_l2_{l2_idx}_l3_{i}.jpg" if with_frames else None
children.append(_make_l3(l1_idx, l2_idx, i, subtitle=sub, frame_path=fp))
l2_subtitle = subtitle or (f"L2事件字幕:L1={l1_idx}_L2={l2_idx}" if with_subtitles else "")
return L2Node(
id=f"l1_{l1_idx}_l2_{l2_idx}",
card=L2Card(
event_description=f"事件:L1={l1_idx},L2={l2_idx}",
entities=[f"角色_{l2_idx}"],
actions=[f"行为_{l2_idx}"],
action_subjects=[f"主体_{l2_idx}"],
visible_text=[],
spatial_relations="左右排列",
state_changes=None,
subtitle=l2_subtitle,
),
time_range=(l2_idx * 30.0, (l2_idx + 1) * 30.0),
children=children,
)
def _make_l1(l1_idx: int, n_l2: int = 3, n_l3: int = 3) -> L1Node:
"""构建 L1 节点,含多个 L2 子节点。"""
return L1Node(
id=f"l1_{l1_idx}",
card=L1Card(
scene_summary=f"场景{l1_idx}摘要",
main_setting="室内" if l1_idx % 2 == 0 else "户外",
key_entities=[f"主角_{l1_idx}"],
main_actions=[f"主行为_{l1_idx}"],
topic_keywords=[f"关键词_{l1_idx}"],
visible_text=[],
temporal_flow="从左到右",
),
time_range=(l1_idx * 600.0, (l1_idx + 1) * 600.0),
children=[_make_l2(l1_idx, j, n_l3) for j in range(n_l2)],
)
@pytest.fixture()
def real_tree() -> TreeIndex:
"""构建包含 2 个 L1、每个 L1 含 3 个 L2、每个 L2 含 5 个 L3 的真实树。
共 2*3*5 = 30 个 L3 节点,6 个 L2 节点,2 个 L1 节点。
所有节点有帧路径和字幕。满足 REASONING_FAMILY 的 min_l3_nodes=4 要求。
"""
meta = IndexMeta(source_path="/test/video.mp4", modality="video")
roots = [_make_l1(i, n_l2=3, n_l3=5) for i in range(2)]
return TreeIndex(metadata=meta, roots=roots)
@pytest.fixture()
def sparse_tree() -> TreeIndex:
"""构建一棵稀疏树——无帧、少字幕,用于测试约束违反。
只有 1 个 L1, 1 个 L2, 1 个 L3。L3 无帧无字幕。
"""
meta = IndexMeta(source_path="/test/sparse.mp4", modality="video")
l3 = _make_l3(0, 0, 0, subtitle="", frame_path=None)
l2 = L2Node(
id="sparse_l2_0",
card=L2Card(
event_description="稀疏事件",
entities=[],
actions=[],
action_subjects=[],
visible_text=[],
spatial_relations="",
state_changes=None,
subtitle="",
),
time_range=(0.0, 30.0),
children=[l3],
)
l1 = L1Node(
id="sparse_l1_0",
card=L1Card(
scene_summary="稀疏场景",
main_setting="未知",
key_entities=[],
main_actions=[],
topic_keywords=[],
visible_text=[],
temporal_flow="",
),
time_range=(0.0, 600.0),
children=[l2],
)
return TreeIndex(metadata=meta, roots=[l1])
# ---------------------------------------------------------------------------
# 测试类
# ---------------------------------------------------------------------------
class TestSampleMaterialV2:
"""sample_material_v2 核心行为测试。"""
def test_returns_material_context(self, real_tree: TreeIndex) -> None:
"""正常采样返回 MaterialContext,字段类型正确。"""
from app.question_gen.sampler_v2 import MaterialContext, sample_material_v2
rng = random.Random(42)
result = sample_material_v2(
tree=real_tree,
task_type="Action Reasoning",
used_node_ids=set(),
rng=rng,
level=2,
constraint=RETRIEVAL_FAMILY.sampling,
)
assert isinstance(result, MaterialContext)
assert result.anchor.node_id # 非空
assert result.anchor.level in (1, 2, 3)
assert len(result.source_nodes) > 0
assert isinstance(result.subtitle_sentences, list)
assert isinstance(result.frame_paths, list)
assert isinstance(result.cross_l2_texts, list)
def test_respects_used_nodes(self, real_tree: TreeIndex) -> None:
"""已用节点被正确排除,不会重复采样。"""
from app.question_gen.sampler_v2 import sample_material_v2
rng = random.Random(42)
# 把所有 L2 节点标记为已用(除了最后一个)
all_l2_ids: set[str] = set()
for l1 in real_tree.roots:
for l2 in l1.children:
all_l2_ids.add(l2.id)
# 留下恰好一个 L2 未用
last_l2_id = real_tree.roots[-1].children[-1].id
used = all_l2_ids - {last_l2_id}
result = sample_material_v2(
tree=real_tree,
task_type="Action Reasoning",
used_node_ids=used,
rng=rng,
level=2,
constraint=RETRIEVAL_FAMILY.sampling,
)
# 锚节点应该是那个未被排除的 L2
assert result.anchor.node_id == last_l2_id
def test_constraint_violation_retries(self, sparse_tree: TreeIndex) -> None:
"""稀疏树上,严格约束满足不了,耗尽重试后抛 RuntimeError。"""
from app.question_gen.sampler_v2 import sample_material_v2
rng = random.Random(42)
# VISUAL_FAMILY 要求 require_frames=True, min_l3_nodes=3
# sparse_tree 只有 1 个 L3 且无帧 → 约束必然违反
with pytest.raises(RuntimeError, match="max_attempts"):
sample_material_v2(
tree=sparse_tree,
task_type="Object Recognition",
used_node_ids=set(),
rng=rng,
level=3,
constraint=VISUAL_FAMILY.sampling,
max_attempts=3,
)
def test_cross_l2_populated_for_reasoning(self, real_tree: TreeIndex) -> None:
"""REASONING 家族要求 cross_l2_span=Truecross_l2_texts 应被填充。"""
from app.question_gen.sampler_v2 import sample_material_v2
rng = random.Random(42)
result = sample_material_v2(
tree=real_tree,
task_type="Action Reasoning",
used_node_ids=set(),
rng=rng,
level=2,
constraint=REASONING_FAMILY.sampling,
)
# cross_l2_span=True 时必须有跨 L2 文本
assert len(result.cross_l2_texts) > 0
def test_subtitle_sentences_from_anchor(self, real_tree: TreeIndex) -> None:
"""采样结果的 subtitle_sentences 来自锚节点所属子树。"""
from app.question_gen.sampler_v2 import sample_material_v2
rng = random.Random(42)
result = sample_material_v2(
tree=real_tree,
task_type="Action Reasoning",
used_node_ids=set(),
rng=rng,
level=2,
constraint=RETRIEVAL_FAMILY.sampling,
)
# real_tree 所有节点都有字幕,所以 subtitle_sentences 非空
assert len(result.subtitle_sentences) > 0
# 字幕应来自锚节点所属的子树(L2 自身字幕 + 子 L3 字幕)
# fixture 中 L2 字幕格式: "L2事件字幕:L1={l1_idx}_L2={l2_idx}"
# fixture 中 L3 字幕格式: "字幕L1={l1_idx}_L2={l2_idx}_L3={l3_idx}"
# 解析锚 L2 的索引信息来验证
anchor_l2_id = result.anchor.l2_id # 如 "l1_1_l2_2"
# 从 ID 提取 L1/L2 索引
parts = anchor_l2_id.split("_") # ["l1", "1", "l2", "2"]
l1_idx, l2_idx = parts[1], parts[3]
# 字幕中应包含 "L1={l1_idx}_L2={l2_idx}" 格式
pattern = f"L1={l1_idx}_L2={l2_idx}"
has_related = any(pattern in s for s in result.subtitle_sentences)
assert has_related
class TestValidateSamplingConstraints:
"""_validate_sampling_constraints 辅助函数测试。"""
def test_passes_relaxed_constraint(self, real_tree: TreeIndex) -> None:
"""宽松约束在丰富树上应通过。"""
from app.question_gen.sampler_v2 import _validate_sampling_constraints
relaxed = SamplingConstraint(
min_subtitles=1,
min_l3_nodes=1,
require_frames=False,
cross_l2_span=False,
)
# 取第一个 L2 节点
node_id = real_tree.roots[0].children[0].id
assert _validate_sampling_constraints(real_tree, node_id, relaxed) is True
def test_fails_strict_frame_constraint(self, sparse_tree: TreeIndex) -> None:
"""require_frames=True 但无帧时应返回 False。"""
from app.question_gen.sampler_v2 import _validate_sampling_constraints
strict = SamplingConstraint(
min_subtitles=0,
min_l3_nodes=1,
require_frames=True,
cross_l2_span=False,
)
node_id = sparse_tree.roots[0].children[0].id
assert _validate_sampling_constraints(sparse_tree, node_id, strict) is False
class TestCollectSubtitleSentences:
"""_collect_subtitle_sentences 辅助函数测试。"""
def test_collects_from_l2_and_l3(self, real_tree: TreeIndex) -> None:
"""收集指定节点的 L2 字幕和子 L3 字幕。"""
from app.question_gen.sampler_v2 import _collect_subtitle_sentences
l2_id = real_tree.roots[0].children[0].id
sentences = _collect_subtitle_sentences(real_tree, (l2_id,))
# 应包含 L2 自身字幕 + 3 个 L3 子节点字幕
assert len(sentences) >= 1
def test_empty_for_no_subtitles(self, sparse_tree: TreeIndex) -> None:
"""无字幕节点返回空列表。"""
from app.question_gen.sampler_v2 import _collect_subtitle_sentences
l2_id = sparse_tree.roots[0].children[0].id
sentences = _collect_subtitle_sentences(sparse_tree, (l2_id,))
assert sentences == []
class TestCollectCrossL2Context:
"""_collect_cross_l2_context 辅助函数测试。"""
def test_returns_peer_l2_descriptions(self, real_tree: TreeIndex) -> None:
"""跨 L2 上下文应返回同 L1 下其他 L2 的描述。"""
from app.question_gen.sampler_v2 import _collect_cross_l2_context
anchor_l2_id = real_tree.roots[0].children[0].id
texts = _collect_cross_l2_context(real_tree, anchor_l2_id, max_peers=3)
# L1_0 有 3 个 L2,排除 anchor 后剩 2 个
assert len(texts) == 2
def test_max_peers_limits_output(self, real_tree: TreeIndex) -> None:
"""max_peers 参数限制返回数量。"""
from app.question_gen.sampler_v2 import _collect_cross_l2_context
anchor_l2_id = real_tree.roots[0].children[0].id
texts = _collect_cross_l2_context(real_tree, anchor_l2_id, max_peers=1)
assert len(texts) <= 1