174 lines
6.1 KiB
Python
174 lines
6.1 KiB
Python
"""题族规格声明单元测试。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import random
|
|
from collections import Counter
|
|
|
|
import pytest
|
|
from scipy import stats
|
|
|
|
from app.question_gen.families import (
|
|
ALL_FAMILIES,
|
|
RETRIEVAL_FAMILY,
|
|
get_family_for_slot,
|
|
)
|
|
|
|
# Video-MME 12 种任务类型
|
|
ALL_TASK_TYPES: frozenset[str] = frozenset(
|
|
[
|
|
"Action Recognition",
|
|
"Action Reasoning",
|
|
"Attribute Perception",
|
|
"Counting Problem",
|
|
"Information Synopsis",
|
|
"Object Recognition",
|
|
"Object Reasoning",
|
|
"OCR Problems",
|
|
"Spatial Perception",
|
|
"Spatial Reasoning",
|
|
"Temporal Perception",
|
|
"Temporal Reasoning",
|
|
]
|
|
)
|
|
|
|
|
|
class TestFamilySpec:
|
|
"""验证 QuestionFamilySpec 声明的完整性与一致性。"""
|
|
|
|
def test_all_families_cover_all_task_types(self) -> None:
|
|
"""12 种任务类型在 ALL_FAMILIES 中全覆盖。"""
|
|
covered: set[str] = set()
|
|
for family in ALL_FAMILIES:
|
|
covered.update(family.legal_task_types)
|
|
assert covered == ALL_TASK_TYPES
|
|
|
|
def test_skill_targets_unique(self) -> None:
|
|
"""每个 family 的 skill_target 不重复。"""
|
|
targets = [f.skill_target for f in ALL_FAMILIES]
|
|
assert len(targets) == len(set(targets))
|
|
|
|
def test_all_families_tuple_length(self) -> None:
|
|
"""ALL_FAMILIES 包含 5 个家族。"""
|
|
assert len(ALL_FAMILIES) == 5
|
|
|
|
def test_family_instances_are_frozen(self) -> None:
|
|
"""Family spec 实例不可修改。"""
|
|
with pytest.raises(AttributeError):
|
|
RETRIEVAL_FAMILY.name = "hacked" # type: ignore[misc]
|
|
|
|
def test_legal_task_types_are_frozensets(self) -> None:
|
|
"""legal_task_types 是 frozenset,确保不可变。"""
|
|
for family in ALL_FAMILIES:
|
|
assert isinstance(family.legal_task_types, frozenset)
|
|
|
|
def test_sampling_constraints_reasonable(self) -> None:
|
|
"""采样约束值在合理范围。"""
|
|
for family in ALL_FAMILIES:
|
|
assert family.sampling.min_subtitles >= 0
|
|
assert family.sampling.min_l3_nodes >= 1
|
|
|
|
def test_leak_profile_thresholds_valid(self) -> None:
|
|
"""泄漏测试阈值在 (0, 1] 范围内。"""
|
|
for family in ALL_FAMILIES:
|
|
assert 0 < family.leak_profile.pass_threshold <= 1.0
|
|
|
|
def test_prompt_templates_not_empty(self) -> None:
|
|
"""每个 family 的 prompt_template 非空。"""
|
|
for family in ALL_FAMILIES:
|
|
assert family.prompt_template
|
|
assert family.prompt_template.endswith(".md")
|
|
|
|
|
|
class TestGetFamilyForSlot:
|
|
"""验证 get_family_for_slot 的核心行为。"""
|
|
|
|
def test_respects_legal_task_types(self) -> None:
|
|
"""返回的 family 必须包含给定 task_type。"""
|
|
rng = random.Random(42)
|
|
for task_type in ALL_TASK_TYPES:
|
|
for _ in range(50):
|
|
family = get_family_for_slot(
|
|
task_type=task_type,
|
|
family_ratios={
|
|
"RETRIEVAL": 0.30,
|
|
"REASONING": 0.25,
|
|
"ENUMERATION": 0.20,
|
|
"VISUAL": 0.15,
|
|
"SPATIAL": 0.10,
|
|
},
|
|
rng=rng,
|
|
)
|
|
assert task_type in family.legal_task_types
|
|
|
|
def test_deterministic_with_seed(self) -> None:
|
|
"""相同 seed 产出相同结果。"""
|
|
ratios = {
|
|
"RETRIEVAL": 0.30,
|
|
"REASONING": 0.25,
|
|
"ENUMERATION": 0.20,
|
|
"VISUAL": 0.15,
|
|
"SPATIAL": 0.10,
|
|
}
|
|
results_a = []
|
|
results_b = []
|
|
for seed in range(10):
|
|
rng_a = random.Random(seed)
|
|
rng_b = random.Random(seed)
|
|
results_a.append(get_family_for_slot("Action Reasoning", ratios, rng_a).name)
|
|
results_b.append(get_family_for_slot("Action Reasoning", ratios, rng_b).name)
|
|
assert results_a == results_b
|
|
|
|
def test_invalid_task_type_raises(self) -> None:
|
|
"""不存在的 task_type 抛 ValueError。"""
|
|
rng = random.Random(0)
|
|
with pytest.raises(ValueError, match="task_type"):
|
|
get_family_for_slot(
|
|
task_type="Nonexistent Type",
|
|
family_ratios={"RETRIEVAL": 1.0},
|
|
rng=rng,
|
|
)
|
|
|
|
def test_distribution_approximates_ratios(self) -> None:
|
|
"""大样本下分布近似指定比例(chi-square p > 0.01)。
|
|
|
|
使用 Action Recognition — 只有 RETRIEVAL 和 VISUAL 合法。
|
|
"""
|
|
rng = random.Random(123)
|
|
ratios = {
|
|
"RETRIEVAL": 0.30,
|
|
"REASONING": 0.25,
|
|
"ENUMERATION": 0.20,
|
|
"VISUAL": 0.15,
|
|
"SPATIAL": 0.10,
|
|
}
|
|
n = 5000
|
|
counter: Counter[str] = Counter()
|
|
for _ in range(n):
|
|
family = get_family_for_slot("Action Recognition", ratios, rng)
|
|
counter[family.name] += 1
|
|
|
|
# Action Recognition 合法族: RETRIEVAL(0.30), VISUAL(0.15)
|
|
# 归一化后: RETRIEVAL=0.30/0.45≈0.667, VISUAL=0.15/0.45≈0.333
|
|
legal_ratios = {"RETRIEVAL": 0.30, "VISUAL": 0.15}
|
|
total_weight = sum(legal_ratios.values())
|
|
expected_probs = {k: v / total_weight for k, v in legal_ratios.items()}
|
|
|
|
observed = [counter.get("RETRIEVAL", 0), counter.get("VISUAL", 0)]
|
|
expected = [expected_probs["RETRIEVAL"] * n, expected_probs["VISUAL"] * n]
|
|
|
|
chi2, p_value = stats.chisquare(observed, expected)
|
|
assert p_value > 0.01, f"分布偏离过大: chi2={chi2:.2f}, p={p_value:.4f}"
|
|
|
|
def test_no_legal_family_raises(self) -> None:
|
|
"""所有 family 权重为 0 时合法族为空,应抛出 ValueError。"""
|
|
rng = random.Random(0)
|
|
# Spatial Reasoning 合法族: SPATIAL
|
|
# 如果 ratios 中只含不合法的族名,应抛错
|
|
with pytest.raises(ValueError, match="合法"):
|
|
get_family_for_slot(
|
|
task_type="Spatial Reasoning",
|
|
family_ratios={"ENUMERATION": 1.0},
|
|
rng=rng,
|
|
)
|