Files
Video-Tree-TRM5/tests/unit/test_synthesizer.py
T
iomgaa 40b04f886e fix(question_gen): sample_anchor Codex 审查修正
- C1/C2: Temporal Reasoning ≥3 L2 + Object Reasoning ≥2 L2 下限检查
- I1: L2 题型子帧不足时 ValueError
- I2: L3 过滤无 frame_path 的节点
- I3/I4: distractor_texts 扩展到整棵树范围
- I5-I7: 测试补强 Information Synopsis/Temporal Reasoning/Object Reasoning
- M1: Spatial Reasoning 测试断言 spatial_layout

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-07-09 05:24:30 -04:00

232 lines
8.8 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.
"""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))
# 找到被选中的 L1,验证 frame_paths 数量 == 该 L1 下全部 L2 数量
chosen_l1 = next(r for r in tree.roots if r.id == ctx.node_id)
assert len(ctx.frame_paths) == len(chosen_l1.children)
def test_l1_type_l2_nodes_in_time_order(self) -> None:
"""Temporal Reasoning 应返回 >=3 帧且 card_text 有实质内容。"""
tree, _vid = _load_test_tree()
ctx = sample_anchor(tree, "Temporal Reasoning", set(), random.Random(42))
assert len(ctx.frame_paths) >= 3
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 2 <= 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_layoutcard_text 必须出现该字段名
assert "spatial_layout" in ctx.card_text
assert len(ctx.card_text) > 20