c3187167c8
两阶段贪心视频级切分:Floor 硬约束优先满足 + 多样性 submodular 覆盖最大化, 两阶段均带 ε 守护保 test 代表性;不可行 fail loud(InfeasibleSplitError), 多样性欠额记 loguru warning 不静默。derive_reportable_types 落地长尾报告门限。 确定性:Random(seed) 预洗牌打破等增益/等槽数平局,同 config 同 videos 同解。 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
211 lines
7.6 KiB
Python
211 lines
7.6 KiB
Python
import sqlite3
|
||
from collections import Counter
|
||
from itertools import cycle
|
||
from pathlib import Path
|
||
|
||
import pytest
|
||
|
||
from app.harness.split_selection import (
|
||
build_video_records,
|
||
cell_of,
|
||
evolution_target_of,
|
||
score_signal,
|
||
)
|
||
|
||
# T2 信号构造用的错误类别循环源(4 类真实 error_type,制造多样格子)。
|
||
_ERROR_TYPES = ["extraction_failure", "search_failure", "reasoning_failure", "mixed"]
|
||
# 二次构造 T2 信号覆盖的高信号题型(floor / diversity 需有料)。
|
||
_SIGNAL_TYPES = {"Counting Problem", "Action Reasoning"}
|
||
_HARNESS_DB = Path(__file__).resolve().parents[2] / "workspaces" / "default" / "harness.db"
|
||
|
||
|
||
def _load_adhoc_predictions() -> list[dict]:
|
||
"""从真实 harness.db 的 infer_adhoc run 读取 900 条预测(按 question_id 去重)。
|
||
|
||
返回:
|
||
preds 列表,每行含 video_id / question_id / task_type / correct,
|
||
correct = (prediction 非空 且 prediction == answer)。
|
||
"""
|
||
con = sqlite3.connect(str(_HARNESS_DB))
|
||
try:
|
||
cur = con.execute(
|
||
"SELECT question_id, video_id, task_type, prediction, answer "
|
||
"FROM predictions WHERE run_id = 'infer_adhoc'"
|
||
)
|
||
seen: set[str] = set()
|
||
preds: list[dict] = []
|
||
for question_id, video_id, task_type, prediction, answer in cur.fetchall():
|
||
if question_id in seen:
|
||
continue
|
||
seen.add(question_id)
|
||
preds.append(
|
||
{
|
||
"video_id": video_id,
|
||
"question_id": question_id,
|
||
"task_type": task_type,
|
||
"correct": prediction is not None and prediction == answer,
|
||
}
|
||
)
|
||
return preds
|
||
finally:
|
||
con.close()
|
||
|
||
|
||
def _real_shaped_video_records() -> list:
|
||
"""用真实预测结构 + 二次构造 T2 信号构建 300 个 VideoRecord。
|
||
|
||
题型 / 难度画像取自真实 infer_adhoc 预测;因真实诊断尚未跑,T2 信号是二次构造:
|
||
对高信号题型(Counting / Action Reasoning)的错题标 T2 defect,error_type 循环取 4 类,
|
||
使 floor 约束有料、多样性格子有区分度。
|
||
|
||
返回:
|
||
全部 300 个真实结构的 VideoRecord 列表。
|
||
"""
|
||
preds = _load_adhoc_predictions()
|
||
err = cycle(_ERROR_TYPES)
|
||
signal_rows: list[dict] = []
|
||
for pred in preds:
|
||
if pred["task_type"] in _SIGNAL_TYPES and not pred["correct"]:
|
||
signal_rows.append(
|
||
{
|
||
"question_id": pred["question_id"],
|
||
"task_type": pred["task_type"],
|
||
"error_type": next(err),
|
||
"tier": "T2",
|
||
}
|
||
)
|
||
return build_video_records(preds, signal_rows)
|
||
|
||
|
||
def _count_questions_by_type(videos: list) -> dict:
|
||
"""统计各 task_type 的视频承载数(含该题型的视频数),供 derive_reportable_types。
|
||
|
||
VideoRecord 只留题型去重集 type_set,故这里按"含该题型的视频数"聚合,
|
||
作为长尾报告门限的代理度量。
|
||
|
||
参数:
|
||
videos: VideoRecord 列表。
|
||
|
||
返回:
|
||
{task_type: 含该题型的视频数}。
|
||
"""
|
||
counter: Counter = Counter()
|
||
for video in videos:
|
||
counter.update(video.type_set)
|
||
return dict(counter)
|
||
|
||
|
||
def test_evolution_target_mapping():
|
||
assert evolution_target_of("extraction_failure") == "tool"
|
||
assert evolution_target_of("search_failure") == "skill"
|
||
assert evolution_target_of("reasoning_failure") == "skill"
|
||
assert evolution_target_of("mixed") == "system"
|
||
|
||
|
||
def test_evolution_target_unknown_raises():
|
||
with pytest.raises(ValueError):
|
||
evolution_target_of("unknown_type")
|
||
|
||
|
||
def test_cell_is_task_type_x_error_type():
|
||
assert cell_of("Counting Problem", "search_failure") == ("Counting Problem", "search_failure")
|
||
|
||
|
||
def test_tiers():
|
||
assert score_signal(cause_category="defect", infra=False, degraded=False).tier == "T2"
|
||
assert score_signal(cause_category="lapse", infra=False, degraded=False).tier == "T1"
|
||
assert (
|
||
score_signal(cause_category="defect", infra=True, degraded=False).tier == "T0"
|
||
) # INFRA 先判
|
||
assert score_signal(cause_category=None, infra=False, degraded=True).tier == "uncertain"
|
||
|
||
|
||
def test_build_video_records_covers_all_videos_with_difficulty_and_types():
|
||
from app.harness.split_selection import build_video_records
|
||
|
||
preds = [
|
||
{
|
||
"video_id": "v1",
|
||
"question_id": "v1-1",
|
||
"task_type": "Counting Problem",
|
||
"correct": False,
|
||
},
|
||
{"video_id": "v1", "question_id": "v1-2", "task_type": "Action Reasoning", "correct": True},
|
||
{"video_id": "v1", "question_id": "v1-3", "task_type": "OCR Problems", "correct": True},
|
||
{"video_id": "v2", "question_id": "v2-1", "task_type": "Counting Problem", "correct": True},
|
||
{"video_id": "v2", "question_id": "v2-2", "task_type": "Counting Problem", "correct": True},
|
||
{"video_id": "v2", "question_id": "v2-3", "task_type": "Counting Problem", "correct": True},
|
||
]
|
||
signal_rows = [
|
||
{
|
||
"question_id": "v1-1",
|
||
"task_type": "Counting Problem",
|
||
"error_type": "search_failure",
|
||
"tier": "T2",
|
||
}
|
||
]
|
||
recs = build_video_records(preds, signal_rows)
|
||
assert {r.video_id for r in recs} == {"v1", "v2"} # 全视频(含零信号 v2)
|
||
v1 = next(r for r in recs if r.video_id == "v1")
|
||
v2 = next(r for r in recs if r.video_id == "v2")
|
||
assert v1.n_correct == 2 and v1.difficulty == 1 # 3题对2 → 难度桶=1错
|
||
assert v2.difficulty == 0 and v2.cells == set() # 零信号视频无 T2 格子
|
||
assert v1.cells == {("Counting Problem", "search_failure")}
|
||
assert v1.type_set == {"Counting Problem", "Action Reasoning", "OCR Problems"}
|
||
assert v1.wrong_by_type == {"Counting Problem": 1} # T2 计数供 floor
|
||
|
||
|
||
def test_select_split_video_disjoint_and_floor_and_deterministic():
|
||
from app.harness.split_selection import (
|
||
SelectConfig,
|
||
derive_reportable_types,
|
||
select_split,
|
||
)
|
||
|
||
videos = _real_shaped_video_records()
|
||
total_by_type = _count_questions_by_type(videos)
|
||
cfg = SelectConfig(
|
||
n_trainval=100,
|
||
floor_k={"Counting Problem": 3},
|
||
epsilon=0.1,
|
||
reportable_types=derive_reportable_types(total_by_type, report_floor=27),
|
||
seed=7,
|
||
)
|
||
a = select_split(videos, config=cfg)
|
||
b = select_split(videos, config=cfg)
|
||
assert set(a.trainval) & set(a.test) == set() # 视频级不相交
|
||
assert set(a.trainval) | set(a.test) == {v.video_id for v in videos} # 补集覆盖全集
|
||
assert a.trainval == b.trainval # 同 seed 同 config → 同解
|
||
trainval_ids = set(a.trainval)
|
||
counting_defects = sum(
|
||
v.wrong_by_type.get("Counting Problem", 0) for v in videos if v.video_id in trainval_ids
|
||
)
|
||
assert counting_defects >= 3 # floor 硬约束满足
|
||
|
||
|
||
def test_infeasible_floor_vs_epsilon_raises():
|
||
from app.harness.split_selection import (
|
||
InfeasibleSplitError,
|
||
SelectConfig,
|
||
select_split,
|
||
)
|
||
|
||
videos = _real_shaped_video_records()
|
||
with pytest.raises(InfeasibleSplitError): # 极小 ε + 高 floor → 死锁
|
||
select_split(
|
||
videos,
|
||
config=SelectConfig(
|
||
n_trainval=2,
|
||
floor_k={"OCR Problems": 50},
|
||
epsilon=0.001,
|
||
reportable_types=set(),
|
||
seed=1,
|
||
),
|
||
)
|
||
|
||
|
||
def test_derive_reportable_types():
|
||
from app.harness.split_selection import derive_reportable_types
|
||
|
||
assert derive_reportable_types({"A": 30, "B": 10}, report_floor=27) == {"A"}
|