a597a9f901
含 6 种层级分支:L3 单帧、L2 多帧、Temporal Perception 特例、 L1 全量/采样 L2、L1-L2 混合。时间排序 + used_node_ids 排除。 Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
231 lines
8.7 KiB
Python
231 lines
8.7 KiB
Python
"""synthesizer 模块单元测试 — AnchorContext + 题型映射常量 + sample_anchor。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import dataclasses
|
||
import random
|
||
from pathlib import Path
|
||
|
||
import pytest
|
||
|
||
from app.question_gen.synthesizer import (
|
||
TASK_TYPE_LEVEL_MAP,
|
||
AnchorContext,
|
||
TaskTypeSpec,
|
||
sample_anchor,
|
||
)
|
||
from app.tree.index import TreeIndex
|
||
|
||
ALL_12_TYPES = [
|
||
"Object Recognition",
|
||
"Attribute Perception",
|
||
"OCR Problems",
|
||
"Spatial Reasoning",
|
||
"Spatial Perception",
|
||
"Action Recognition",
|
||
"Action Reasoning",
|
||
"Counting Problem",
|
||
"Temporal Perception",
|
||
"Temporal Reasoning",
|
||
"Information Synopsis",
|
||
"Object Reasoning",
|
||
]
|
||
|
||
|
||
class TestTaskTypeLevelMap:
|
||
"""TASK_TYPE_LEVEL_MAP 覆盖性与结构测试。"""
|
||
|
||
def test_covers_all_12_types(self) -> None:
|
||
"""映射表必须覆盖全部 12 种 Video-MME 题型。"""
|
||
assert set(TASK_TYPE_LEVEL_MAP.keys()) == set(ALL_12_TYPES)
|
||
|
||
def test_no_extra_types(self) -> None:
|
||
"""映射表不得包含 12 种标准题型之外的条目。"""
|
||
assert len(TASK_TYPE_LEVEL_MAP) == 12
|
||
|
||
def test_all_values_are_task_type_spec(self) -> None:
|
||
"""每个映射值必须是 TaskTypeSpec 实例。"""
|
||
for task_type, spec in TASK_TYPE_LEVEL_MAP.items():
|
||
assert isinstance(spec, TaskTypeSpec), f"{task_type} 映射值类型错误: {type(spec)}"
|
||
|
||
def test_level_values_valid(self) -> None:
|
||
"""每个 spec 的 level 必须是合法层级标识。"""
|
||
valid_levels = {"L1", "L2", "L3", "L1-L2"}
|
||
for task_type, spec in TASK_TYPE_LEVEL_MAP.items():
|
||
assert spec.level in valid_levels, (
|
||
f"{task_type} 层级 '{spec.level}' 不在 {valid_levels}"
|
||
)
|
||
|
||
def test_context_fields_non_empty(self) -> None:
|
||
"""每个 spec 的 context_fields 至少有一个字段。"""
|
||
for task_type, spec in TASK_TYPE_LEVEL_MAP.items():
|
||
assert len(spec.context_fields) >= 1, f"{task_type} 的 context_fields 为空"
|
||
|
||
|
||
class TestAnchorContext:
|
||
"""AnchorContext 数据类测试。"""
|
||
|
||
def test_frozen(self) -> None:
|
||
"""AnchorContext 是不可变的。"""
|
||
ctx = AnchorContext(
|
||
node_id="L3_001",
|
||
card_text="A person walks into a room",
|
||
frame_paths=["/data/frames/001.jpg"],
|
||
subtitle="Hello there",
|
||
distractor_texts=["A car drives by"],
|
||
)
|
||
assert ctx.node_id == "L3_001"
|
||
assert ctx.card_text == "A person walks into a room"
|
||
assert ctx.frame_paths == ["/data/frames/001.jpg"]
|
||
assert ctx.subtitle == "Hello there"
|
||
assert ctx.distractor_texts == ["A car drives by"]
|
||
|
||
def test_mutation_raises(self) -> None:
|
||
"""frozen dataclass 拒绝赋值修改。"""
|
||
ctx = AnchorContext(
|
||
node_id="L3_001",
|
||
card_text="test",
|
||
frame_paths=["a.jpg"],
|
||
subtitle="",
|
||
distractor_texts=["other node"],
|
||
)
|
||
try:
|
||
ctx.node_id = "L3_002" # type: ignore[misc]
|
||
raise AssertionError("应抛出 FrozenInstanceError")
|
||
except dataclasses.FrozenInstanceError:
|
||
pass
|
||
|
||
def test_empty_subtitle_allowed(self) -> None:
|
||
"""subtitle 可以为空字符串。"""
|
||
ctx = AnchorContext(
|
||
node_id="L2_010",
|
||
card_text="scene card",
|
||
frame_paths=[],
|
||
subtitle="",
|
||
distractor_texts=[],
|
||
)
|
||
assert ctx.subtitle == ""
|
||
|
||
def test_multiple_frame_paths(self) -> None:
|
||
"""frame_paths 可包含多个路径。"""
|
||
paths = ["/data/f1.jpg", "/data/f2.jpg", "/data/f3.jpg"]
|
||
ctx = AnchorContext(
|
||
node_id="L2_005",
|
||
card_text="multi-frame event",
|
||
frame_paths=paths,
|
||
subtitle="Dialogue line",
|
||
distractor_texts=["other1", "other2"],
|
||
)
|
||
assert len(ctx.frame_paths) == 3
|
||
|
||
|
||
class TestTaskTypeSpec:
|
||
"""TaskTypeSpec 数据类测试。"""
|
||
|
||
def test_frozen(self) -> None:
|
||
"""TaskTypeSpec 是不可变的。"""
|
||
spec = TaskTypeSpec(
|
||
level="L3",
|
||
needs_frames=True,
|
||
frame_count="1",
|
||
context_fields=("frame_summary",),
|
||
)
|
||
try:
|
||
spec.level = "L2" # type: ignore[misc]
|
||
raise AssertionError("应抛出 FrozenInstanceError")
|
||
except dataclasses.FrozenInstanceError:
|
||
pass
|
||
|
||
def test_context_fields_is_tuple(self) -> None:
|
||
"""context_fields 应为 tuple(不可变)。"""
|
||
for task_type, spec in TASK_TYPE_LEVEL_MAP.items():
|
||
assert isinstance(spec.context_fields, tuple), (
|
||
f"{task_type} 的 context_fields 不是 tuple"
|
||
)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# sample_anchor 测试
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _load_test_tree() -> tuple[TreeIndex, str]:
|
||
"""加载真实测试树(store/videos/ 下第一棵)。"""
|
||
videos_dir = Path("store/videos")
|
||
first_vid = sorted(videos_dir.iterdir())[0]
|
||
tree = TreeIndex.load_json(str(first_vid / "tree.json"))
|
||
return tree, first_vid.name
|
||
|
||
|
||
class TestSampleAnchor:
|
||
"""sample_anchor 锚节点采样测试(基于真实树数据)。"""
|
||
|
||
def test_l3_type_returns_single_frame(self) -> None:
|
||
"""L3 题型(Object Recognition)应返回恰好 1 帧。"""
|
||
tree, _vid = _load_test_tree()
|
||
ctx = sample_anchor(tree, "Object Recognition", set(), random.Random(42))
|
||
assert len(ctx.frame_paths) == 1
|
||
assert ctx.node_id.startswith("L") or "_L3_" in ctx.node_id
|
||
assert len(ctx.distractor_texts) > 0
|
||
|
||
def test_l2_type_returns_multiple_frames(self) -> None:
|
||
"""L2 题型(Action Reasoning)应返回 2-3 帧。"""
|
||
tree, _vid = _load_test_tree()
|
||
ctx = sample_anchor(tree, "Action Reasoning", set(), random.Random(42))
|
||
assert 2 <= len(ctx.frame_paths) <= 3
|
||
|
||
def test_temporal_perception_zero_or_one_frame(self) -> None:
|
||
"""Temporal Perception 特殊处理:0-1 帧,且 card_text 包含 time_range。"""
|
||
tree, _vid = _load_test_tree()
|
||
ctx = sample_anchor(tree, "Temporal Perception", set(), random.Random(42))
|
||
assert len(ctx.frame_paths) <= 1
|
||
assert "time_range" in ctx.card_text.lower() or "time" in ctx.card_text.lower()
|
||
|
||
def test_information_synopsis_uses_all_l2(self) -> None:
|
||
"""Information Synopsis 必须使用目标 L1 下所有 L2 子节点。"""
|
||
tree, _vid = _load_test_tree()
|
||
ctx = sample_anchor(tree, "Information Synopsis", set(), random.Random(42))
|
||
# 至少应有帧(每个 L2 取一帧代表)
|
||
total_l2 = sum(len(r.children) for r in tree.roots)
|
||
assert len(ctx.frame_paths) >= min(total_l2, 1)
|
||
|
||
def test_l1_type_l2_nodes_in_time_order(self) -> None:
|
||
"""L1 题型(Temporal Reasoning)的 card_text 应有实质内容。"""
|
||
tree, _vid = _load_test_tree()
|
||
ctx = sample_anchor(tree, "Temporal Reasoning", set(), random.Random(42))
|
||
assert len(ctx.frame_paths) >= 1
|
||
assert len(ctx.card_text) > 20
|
||
|
||
def test_used_node_ids_excluded(self) -> None:
|
||
"""used_node_ids 中的节点不应被再次选中。"""
|
||
tree, _vid = _load_test_tree()
|
||
rng = random.Random(42)
|
||
ctx1 = sample_anchor(tree, "Object Recognition", set(), rng)
|
||
ctx2 = sample_anchor(tree, "Object Recognition", {ctx1.node_id}, random.Random(43))
|
||
assert ctx2.node_id != ctx1.node_id
|
||
|
||
def test_insufficient_nodes_raises(self) -> None:
|
||
"""所有候选节点均被排除时应抛出 ValueError。"""
|
||
tree, _vid = _load_test_tree()
|
||
all_l3_ids: set[str] = set()
|
||
for root in tree.roots:
|
||
for l2 in root.children:
|
||
for l3 in l2.children:
|
||
all_l3_ids.add(l3.id)
|
||
with pytest.raises(ValueError, match="锚节点不足"):
|
||
sample_anchor(tree, "Object Recognition", all_l3_ids, random.Random(42))
|
||
|
||
def test_object_reasoning_l1_l2_type(self) -> None:
|
||
"""Object Reasoning (L1-L2) 应选 2-3 个 L2 并按时间排序。"""
|
||
tree, _vid = _load_test_tree()
|
||
ctx = sample_anchor(tree, "Object Reasoning", set(), random.Random(42))
|
||
assert 1 <= len(ctx.frame_paths) <= 3
|
||
assert len(ctx.card_text) > 10
|
||
|
||
def test_spatial_reasoning_context_fields(self) -> None:
|
||
"""Spatial Reasoning 的 card_text 应包含 spatial_layout 内容。"""
|
||
tree, _vid = _load_test_tree()
|
||
ctx = sample_anchor(tree, "Spatial Reasoning", set(), random.Random(42))
|
||
# context_fields 包含 spatial_layout,card_text 应含有该字段内容
|
||
assert len(ctx.card_text) > 20
|