399 lines
14 KiB
Python
399 lines
14 KiB
Python
"""tools/generate_questions.py 单元测试。
|
|
|
|
覆盖断点续跑、exemplar 选取、embedding 池重建、JSON 追加写入等纯函数。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import random
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
import pytest
|
|
|
|
# 确保项目根目录在 sys.path 中
|
|
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
|
if str(PROJECT_ROOT) not in sys.path:
|
|
sys.path.insert(0, str(PROJECT_ROOT))
|
|
|
|
from core.types import GeneratedQuestion
|
|
from tools.generate_questions import (
|
|
_append_to_json,
|
|
_calibrate_exit_code,
|
|
_judge_task_type,
|
|
_load_or_init_progress,
|
|
_rebuild_embedding_pool,
|
|
_save_progress,
|
|
_select_exemplars,
|
|
)
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# 辅助工厂
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _make_question(
|
|
qid: str = "q1",
|
|
vid: str = "v1",
|
|
task_type: str = "Object Recognition",
|
|
question: str = "What is this?",
|
|
answer: str = "A",
|
|
) -> GeneratedQuestion:
|
|
"""构造测试用 GeneratedQuestion。"""
|
|
return GeneratedQuestion(
|
|
question_id=qid,
|
|
video_id=vid,
|
|
task_type=task_type,
|
|
question=question,
|
|
options=("A. X", "B. Y", "C. Z", "D. W"),
|
|
answer=answer,
|
|
source_nodes=(),
|
|
difficulty="medium",
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestLoadOrInitProgress
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestLoadOrInitProgress:
|
|
"""_load_or_init_progress 测试。"""
|
|
|
|
def test_init_fresh(self, tmp_path: Path) -> None:
|
|
"""目录为空时返回初始结构。"""
|
|
progress = _load_or_init_progress(tmp_path)
|
|
assert progress["completed"] == {}
|
|
assert progress["output_dir"] == str(tmp_path)
|
|
|
|
def test_load_existing(self, tmp_path: Path) -> None:
|
|
"""已有 progress.json 时正确加载。"""
|
|
data = {
|
|
"completed": {"Object Recognition": ["gen-x-001"]},
|
|
"output_dir": str(tmp_path),
|
|
}
|
|
(tmp_path / "progress.json").write_text(json.dumps(data))
|
|
progress = _load_or_init_progress(tmp_path)
|
|
assert "gen-x-001" in progress["completed"]["Object Recognition"]
|
|
|
|
def test_corrupted_json_reinits(self, tmp_path: Path) -> None:
|
|
"""损坏的 JSON 文件导致重新初始化。"""
|
|
(tmp_path / "progress.json").write_text("{invalid json")
|
|
progress = _load_or_init_progress(tmp_path)
|
|
assert progress["completed"] == {}
|
|
|
|
def test_invalid_completed_type_reinits(self, tmp_path: Path) -> None:
|
|
"""completed 字段类型不正确时重新初始化。"""
|
|
(tmp_path / "progress.json").write_text(
|
|
json.dumps({"completed": "not-a-dict", "output_dir": str(tmp_path)})
|
|
)
|
|
progress = _load_or_init_progress(tmp_path)
|
|
assert progress["completed"] == {}
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestSaveProgress
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestSaveProgress:
|
|
"""_save_progress 测试。"""
|
|
|
|
def test_atomic_write(self, tmp_path: Path) -> None:
|
|
"""原子写入 progress.json。"""
|
|
progress = {
|
|
"completed": {"Action Reasoning": ["gen-v1-001"]},
|
|
"output_dir": str(tmp_path),
|
|
}
|
|
_save_progress(tmp_path, progress)
|
|
|
|
written = json.loads((tmp_path / "progress.json").read_text())
|
|
assert written["completed"]["Action Reasoning"] == ["gen-v1-001"]
|
|
# 临时文件不应残留
|
|
assert not (tmp_path / "progress.json.tmp").exists()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestSelectExemplars
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestSelectExemplars:
|
|
"""_select_exemplars 测试。"""
|
|
|
|
def test_selects_correct_type(self) -> None:
|
|
"""只选取匹配题型的示例。"""
|
|
qs = [
|
|
_make_question("q1", "v1", "Object Recognition", "Q1?"),
|
|
_make_question("q2", "v2", "Object Recognition", "Q2?"),
|
|
_make_question("q3", "v1", "Action Reasoning", "Q3?"),
|
|
]
|
|
result = _select_exemplars(qs, "Object Recognition", 3, random.Random(42))
|
|
assert all(q.task_type == "Object Recognition" for q in result)
|
|
assert len(result) == 2 # 只有 2 个可用
|
|
|
|
def test_cross_video_diversity(self) -> None:
|
|
"""优先从不同 video_id 选取示例。"""
|
|
qs = [_make_question(f"q{i}", f"v{i}", "Object Recognition", f"Q{i}?") for i in range(10)]
|
|
result = _select_exemplars(qs, "Object Recognition", 3, random.Random(42))
|
|
video_ids = {q.video_id for q in result}
|
|
assert len(video_ids) == 3 # 全部来自不同视频
|
|
|
|
def test_empty_benchmark(self) -> None:
|
|
"""benchmark 为空时返回空列表。"""
|
|
result = _select_exemplars([], "Object Recognition", 3, random.Random(42))
|
|
assert result == []
|
|
|
|
def test_no_matching_type(self) -> None:
|
|
"""无匹配题型时返回空列表。"""
|
|
qs = [_make_question("q1", "v1", "Action Reasoning", "Q1?")]
|
|
result = _select_exemplars(qs, "Object Recognition", 3, random.Random(42))
|
|
assert result == []
|
|
|
|
def test_request_more_than_available(self) -> None:
|
|
"""请求数超过可用数时返回全部。"""
|
|
qs = [
|
|
_make_question("q1", "v1", "Object Recognition", "Q1?"),
|
|
_make_question("q2", "v2", "Object Recognition", "Q2?"),
|
|
]
|
|
result = _select_exemplars(qs, "Object Recognition", 10, random.Random(42))
|
|
assert len(result) == 2
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestProgressResume
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestProgressResume:
|
|
"""断点续跑集成测试。"""
|
|
|
|
def test_skips_completed_and_rebuilds_pool(self, tmp_path: Path) -> None:
|
|
"""已完成的题目从 progress 加载,embedding 池从已生成 JSON 重建。"""
|
|
output_dir = tmp_path / "output"
|
|
output_dir.mkdir()
|
|
(output_dir / "test_vid.json").write_text(
|
|
json.dumps(
|
|
[
|
|
{
|
|
"question_id": "gen-test_vid-001",
|
|
"task_type": "Object Recognition",
|
|
"question": "Existing question?",
|
|
"options": ["A. X", "B. Y", "C. Z", "D. W"],
|
|
"answer": "A",
|
|
"source_nodes": ["L3_001"],
|
|
"difficulty": "medium",
|
|
}
|
|
]
|
|
)
|
|
)
|
|
progress = {
|
|
"completed": {"Object Recognition": ["gen-test_vid-001"]},
|
|
"output_dir": str(output_dir),
|
|
}
|
|
(output_dir / "progress.json").write_text(json.dumps(progress))
|
|
|
|
loaded = _load_or_init_progress(output_dir)
|
|
assert "gen-test_vid-001" in loaded["completed"]["Object Recognition"]
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestRebuildEmbeddingPool
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestRebuildEmbeddingPool:
|
|
"""_rebuild_embedding_pool 测试。"""
|
|
|
|
@staticmethod
|
|
def _fake_embed_fn(texts):
|
|
"""伪嵌入函数:返回固定维度的随机向量。"""
|
|
if isinstance(texts, str):
|
|
texts = [texts]
|
|
return np.random.RandomState(0).randn(len(texts), 8).astype(np.float32)
|
|
|
|
def test_empty_dir(self, tmp_path: Path) -> None:
|
|
"""空目录时所有池为空。"""
|
|
pools = _rebuild_embedding_pool(tmp_path, self._fake_embed_fn, [])
|
|
# 所有 12 种题型都应有条目
|
|
assert len(pools) >= 12
|
|
for v in pools.values():
|
|
assert v.shape[0] == 0 or v.ndim == 2
|
|
|
|
def test_with_generated_json(self, tmp_path: Path) -> None:
|
|
"""从已生成的 JSON 文件重建池。"""
|
|
(tmp_path / "vid1.json").write_text(
|
|
json.dumps(
|
|
[
|
|
{
|
|
"question_id": "gen-vid1-001",
|
|
"task_type": "Object Recognition",
|
|
"question": "Test question 1?",
|
|
"options": ["A. X", "B. Y", "C. Z", "D. W"],
|
|
"answer": "A",
|
|
"source_nodes": [],
|
|
"difficulty": "medium",
|
|
},
|
|
{
|
|
"question_id": "gen-vid1-002",
|
|
"task_type": "Object Recognition",
|
|
"question": "Test question 2?",
|
|
"options": ["A. X", "B. Y", "C. Z", "D. W"],
|
|
"answer": "B",
|
|
"source_nodes": [],
|
|
"difficulty": "medium",
|
|
},
|
|
]
|
|
)
|
|
)
|
|
pools = _rebuild_embedding_pool(tmp_path, self._fake_embed_fn, [])
|
|
assert pools["Object Recognition"].shape[0] == 2
|
|
assert pools["Object Recognition"].shape[1] == 8
|
|
|
|
def test_with_benchmark_questions(self, tmp_path: Path) -> None:
|
|
"""benchmark 题目也加入去重池。"""
|
|
benchmark = [
|
|
_make_question("bm1", "v1", "Action Reasoning", "Benchmark Q1?"),
|
|
_make_question("bm2", "v2", "Action Reasoning", "Benchmark Q2?"),
|
|
]
|
|
pools = _rebuild_embedding_pool(tmp_path, self._fake_embed_fn, benchmark)
|
|
assert pools["Action Reasoning"].shape[0] == 2
|
|
|
|
def test_progress_json_excluded(self, tmp_path: Path) -> None:
|
|
"""progress.json 不被当作题目文件。"""
|
|
(tmp_path / "progress.json").write_text(
|
|
json.dumps({"completed": {}, "output_dir": str(tmp_path)})
|
|
)
|
|
pools = _rebuild_embedding_pool(tmp_path, self._fake_embed_fn, [])
|
|
for v in pools.values():
|
|
assert v.shape[0] == 0 or v.ndim == 2
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestAppendToJson
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestAppendToJson:
|
|
"""_append_to_json 测试。"""
|
|
|
|
def test_create_new_file(self, tmp_path: Path) -> None:
|
|
"""文件不存在时创建新文件。"""
|
|
q = _make_question("gen-v1-001", "v1", "Object Recognition", "Q?")
|
|
_append_to_json(tmp_path, q)
|
|
|
|
written = json.loads((tmp_path / "v1.json").read_text())
|
|
assert len(written) == 1
|
|
assert written[0]["question_id"] == "gen-v1-001"
|
|
|
|
def test_append_to_existing(self, tmp_path: Path) -> None:
|
|
"""追加到已有文件。"""
|
|
existing = [
|
|
{
|
|
"question_id": "gen-v1-001",
|
|
"task_type": "Object Recognition",
|
|
"question": "Existing?",
|
|
"options": ["A. X", "B. Y", "C. Z", "D. W"],
|
|
"answer": "A",
|
|
"source_nodes": [],
|
|
"difficulty": "medium",
|
|
}
|
|
]
|
|
(tmp_path / "v1.json").write_text(json.dumps(existing))
|
|
|
|
q = _make_question("gen-v1-002", "v1", "Object Recognition", "New?")
|
|
_append_to_json(tmp_path, q)
|
|
|
|
written = json.loads((tmp_path / "v1.json").read_text())
|
|
assert len(written) == 2
|
|
assert written[1]["question_id"] == "gen-v1-002"
|
|
|
|
def test_different_video_ids(self, tmp_path: Path) -> None:
|
|
"""不同 video_id 写入不同文件。"""
|
|
q1 = _make_question("gen-v1-001", "v1", "Object Recognition", "Q1?")
|
|
q2 = _make_question("gen-v2-001", "v2", "Action Reasoning", "Q2?")
|
|
_append_to_json(tmp_path, q1)
|
|
_append_to_json(tmp_path, q2)
|
|
|
|
assert (tmp_path / "v1.json").exists()
|
|
assert (tmp_path / "v2.json").exists()
|
|
v1_data = json.loads((tmp_path / "v1.json").read_text())
|
|
v2_data = json.loads((tmp_path / "v2.json").read_text())
|
|
assert len(v1_data) == 1
|
|
assert len(v2_data) == 1
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestCalibrateJudgment
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestCalibrateJudgment:
|
|
"""_judge_task_type 校准判定测试。"""
|
|
|
|
def test_pass_when_delta_small(self) -> None:
|
|
"""差值在容忍范围内判定为 PASS。"""
|
|
verdict = _judge_task_type(
|
|
bench_correct=60,
|
|
bench_total=100,
|
|
gen_correct=12,
|
|
gen_total=20,
|
|
tolerance=0.10,
|
|
alpha=0.05,
|
|
)
|
|
assert verdict == "PASS"
|
|
|
|
def test_fail_when_delta_large_and_significant(self) -> None:
|
|
"""差值超阈值且统计显著判定为 FAIL。"""
|
|
verdict = _judge_task_type(
|
|
bench_correct=144,
|
|
bench_total=240,
|
|
gen_correct=6,
|
|
gen_total=20,
|
|
tolerance=0.10,
|
|
alpha=0.05,
|
|
)
|
|
assert verdict == "FAIL"
|
|
|
|
def test_warn_when_delta_large_but_not_significant(self) -> None:
|
|
"""差值超阈值但不统计显著判定为 WARN。"""
|
|
verdict = _judge_task_type(
|
|
bench_correct=2,
|
|
bench_total=3,
|
|
gen_correct=8,
|
|
gen_total=20,
|
|
tolerance=0.10,
|
|
alpha=0.05,
|
|
)
|
|
assert verdict == "WARN"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TestCalibrateIntegration
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestCalibrateIntegration:
|
|
"""calibrate 辅助函数集成测试。"""
|
|
|
|
def test_has_fail_returns_exit_code_1(self) -> None:
|
|
"""存在 FAIL 时返回退出码 1。"""
|
|
verdicts = {"Object Recognition": "PASS", "Action Reasoning": "FAIL"}
|
|
assert _calibrate_exit_code(verdicts) == 1
|
|
|
|
def test_all_pass_or_warn_returns_exit_code_0(self) -> None:
|
|
"""全部 PASS 或 WARN 时返回退出码 0。"""
|
|
verdicts = {"Object Recognition": "PASS", "Spatial Perception": "WARN"}
|
|
assert _calibrate_exit_code(verdicts) == 0
|
|
|
|
def test_all_pass_returns_exit_code_0(self) -> None:
|
|
"""全部 PASS 时返回退出码 0。"""
|
|
verdicts = {"Object Recognition": "PASS", "Action Reasoning": "PASS"}
|
|
assert _calibrate_exit_code(verdicts) == 0
|
|
|
|
def test_empty_verdicts_returns_exit_code_0(self) -> None:
|
|
"""空 verdicts 时返回退出码 0。"""
|
|
assert _calibrate_exit_code({}) == 0
|