feat(harness): add Action Recognition training experiment
- PerCategoryPoolStrategy: filter test pool by task_types - RunConfig: add run_holdout_eval toggle (default true) - load_config: fix YAML task_types list-to-tuple conversion - Runner: conditionally skip _holdout_four_way when disabled - CLI: add --no-run-holdout-eval flag - New config/train_action_recognition.yaml (3 epochs, per_category) - New scripts/train_action_recognition.sh (baseline + seed + train)
This commit is contained in:
@@ -63,7 +63,7 @@ def _make_question_set(
|
||||
返回:
|
||||
题目列表。
|
||||
"""
|
||||
types = task_types or ["Action Reasoning", "Scene Understanding"]
|
||||
types = task_types or ["Action Reasoning", "Information Synopsis"]
|
||||
return [_make_question(f"q_{i:04d}", types[i % len(types)]) for i in range(n)]
|
||||
|
||||
|
||||
@@ -347,10 +347,18 @@ class TestGlobalPoolStrategy:
|
||||
def _make_per_category_questions():
|
||||
"""构造 12 类各 30 题,共 360 题。"""
|
||||
task_types = [
|
||||
"Action Prediction", "Action Reasoning", "Action Recognition",
|
||||
"Action Sequence", "Causal Reasoning", "Event Reasoning",
|
||||
"Object Interaction", "Object Reasoning", "Object Recognition",
|
||||
"Scene Understanding", "Spatial Reasoning", "Temporal Reasoning",
|
||||
"Action Recognition",
|
||||
"Action Reasoning",
|
||||
"Attribute Perception",
|
||||
"Counting Problem",
|
||||
"Information Synopsis",
|
||||
"Object Recognition",
|
||||
"Object Reasoning",
|
||||
"OCR Problems",
|
||||
"Spatial Perception",
|
||||
"Spatial Reasoning",
|
||||
"Temporal Perception",
|
||||
"Temporal Reasoning",
|
||||
]
|
||||
questions = []
|
||||
for tt in task_types:
|
||||
@@ -371,10 +379,17 @@ class TestPerCategoryPoolStrategy:
|
||||
correctness[q.question_id] = idx < 18
|
||||
|
||||
config = PoolConfig(
|
||||
task_types=None, seed=42, baseline_run_id="baseline_v2",
|
||||
diag_size=0, diag_correct_ratio=0.0, val_size=0,
|
||||
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
|
||||
train_ratio=20 / 30, test_questions_dir=None,
|
||||
task_types=None,
|
||||
seed=42,
|
||||
baseline_run_id="baseline_v2",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=None,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
pools = strategy.build(questions, correctness, config)
|
||||
@@ -382,6 +397,7 @@ class TestPerCategoryPoolStrategy:
|
||||
assert len(pools.validation) == 120
|
||||
|
||||
from collections import Counter
|
||||
|
||||
diag_counts = Counter(q.task_type for q in pools.diagnosis)
|
||||
val_counts = Counter(q.task_type for q in pools.validation)
|
||||
for tt in diag_counts:
|
||||
@@ -401,10 +417,17 @@ class TestPerCategoryPoolStrategy:
|
||||
correctness[q.question_id] = idx < 18
|
||||
|
||||
config = PoolConfig(
|
||||
task_types=("Action Reasoning",), seed=42, baseline_run_id="baseline_v2",
|
||||
diag_size=0, diag_correct_ratio=0.0, val_size=0,
|
||||
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
|
||||
train_ratio=20 / 30, test_questions_dir=None,
|
||||
task_types=("Action Reasoning",),
|
||||
seed=42,
|
||||
baseline_run_id="baseline_v2",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=None,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
pools = strategy.build(questions, correctness, config)
|
||||
@@ -420,10 +443,17 @@ class TestPerCategoryPoolStrategy:
|
||||
questions = [_make_question(f"q_{i:03d}", "Object Recognition") for i in range(30)]
|
||||
correctness = {q.question_id: True for q in questions}
|
||||
config = PoolConfig(
|
||||
task_types=None, seed=42, baseline_run_id="r",
|
||||
diag_size=0, diag_correct_ratio=0.0, val_size=0,
|
||||
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
|
||||
train_ratio=20 / 30, test_questions_dir=None,
|
||||
task_types=None,
|
||||
seed=42,
|
||||
baseline_run_id="r",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=None,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
pools = strategy.build(questions, correctness, config)
|
||||
@@ -435,10 +465,17 @@ class TestPerCategoryPoolStrategy:
|
||||
questions = [_make_question(f"q_{i:03d}", "Object Recognition") for i in range(30)]
|
||||
correctness = {q.question_id: True for q in questions[:25]}
|
||||
config = PoolConfig(
|
||||
task_types=None, seed=42, baseline_run_id="r",
|
||||
diag_size=0, diag_correct_ratio=0.0, val_size=0,
|
||||
val_correct_ratio=0.0, test_size=0, eval_min_per_class=0,
|
||||
train_ratio=20 / 30, test_questions_dir=None,
|
||||
task_types=None,
|
||||
seed=42,
|
||||
baseline_run_id="r",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=None,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
with pytest.raises(ValueError, match="correctness 缺失"):
|
||||
@@ -449,17 +486,131 @@ class TestPerCategoryPoolStrategy:
|
||||
questions = _make_per_category_questions()
|
||||
correctness = {q.question_id: True for q in questions}
|
||||
config = PoolConfig(
|
||||
task_types=("Action Reasoning", "Scene Understanding"), seed=42,
|
||||
baseline_run_id="r", diag_size=0, diag_correct_ratio=0.0,
|
||||
val_size=0, val_correct_ratio=0.0, test_size=0,
|
||||
eval_min_per_class=0, train_ratio=20 / 30, test_questions_dir=None,
|
||||
task_types=("Action Reasoning", "Information Synopsis"),
|
||||
seed=42,
|
||||
baseline_run_id="r",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=None,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
pools = strategy.build(questions, correctness, config)
|
||||
assert len(pools.diagnosis) == 40
|
||||
assert len(pools.validation) == 20
|
||||
types_in_diag = {q.task_type for q in pools.diagnosis}
|
||||
assert types_in_diag == {"Action Reasoning", "Scene Understanding"}
|
||||
assert types_in_diag == {"Action Reasoning", "Information Synopsis"}
|
||||
|
||||
def test_per_category_test_pool_filtered_by_task_types(self, tmp_path: Path) -> None:
|
||||
"""test_questions_dir 加载的 test 池应按 task_types 过滤。"""
|
||||
|
||||
test_dir = tmp_path / "test_questions"
|
||||
test_dir.mkdir()
|
||||
|
||||
task_types_all = [
|
||||
"Action Recognition",
|
||||
"Action Reasoning",
|
||||
"Temporal Perception",
|
||||
]
|
||||
for tt in task_types_all:
|
||||
items = []
|
||||
for i in range(10):
|
||||
items.append(
|
||||
{
|
||||
"question_id": f"{tt}_{i:03d}",
|
||||
"video_id": "v1",
|
||||
"task_type": tt,
|
||||
"question": f"Q {tt} {i}?",
|
||||
"options": ["A. a", "B. b", "C. c", "D. d"],
|
||||
"answer": "A",
|
||||
}
|
||||
)
|
||||
slug = tt.lower().replace(" ", "_")
|
||||
(test_dir / f"{slug}.json").write_text(json.dumps(items, ensure_ascii=False))
|
||||
|
||||
# train/val 用的题目(与 test 独立)
|
||||
questions = _make_per_category_questions()
|
||||
correctness = {q.question_id: True for q in questions}
|
||||
|
||||
config = PoolConfig(
|
||||
task_types=("Action Recognition",),
|
||||
seed=42,
|
||||
baseline_run_id="baseline_v2",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=test_dir,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
pools = strategy.build(questions, correctness, config)
|
||||
|
||||
assert len(pools.test) == 10, (
|
||||
f"test 池应仅含 Action Recognition 的 10 题,实际 {len(pools.test)}"
|
||||
)
|
||||
test_types = {q.task_type for q in pools.test}
|
||||
assert test_types == {"Action Recognition"}, (
|
||||
f"test 池应仅含 Action Recognition,实际含 {test_types}"
|
||||
)
|
||||
|
||||
def test_per_category_test_pool_no_filter_when_task_types_none(self, tmp_path: Path) -> None:
|
||||
"""task_types=None 时,test 池不过滤,加载全部题目。"""
|
||||
|
||||
test_dir = tmp_path / "test_questions"
|
||||
test_dir.mkdir()
|
||||
|
||||
task_types_all = [
|
||||
"Action Recognition",
|
||||
"Action Reasoning",
|
||||
"Temporal Perception",
|
||||
]
|
||||
total_expected = 0
|
||||
for tt in task_types_all:
|
||||
items = []
|
||||
for i in range(10):
|
||||
items.append(
|
||||
{
|
||||
"question_id": f"{tt}_{i:03d}",
|
||||
"video_id": "v1",
|
||||
"task_type": tt,
|
||||
"question": f"Q {tt} {i}?",
|
||||
"options": ["A. a", "B. b", "C. c", "D. d"],
|
||||
"answer": "A",
|
||||
}
|
||||
)
|
||||
slug = tt.lower().replace(" ", "_")
|
||||
(test_dir / f"{slug}.json").write_text(json.dumps(items, ensure_ascii=False))
|
||||
total_expected += len(items)
|
||||
|
||||
questions = _make_per_category_questions()
|
||||
correctness = {q.question_id: True for q in questions}
|
||||
|
||||
config = PoolConfig(
|
||||
task_types=None,
|
||||
seed=42,
|
||||
baseline_run_id="baseline_v2",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=test_dir,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
pools = strategy.build(questions, correctness, config)
|
||||
|
||||
assert len(pools.test) == total_expected, (
|
||||
f"task_types=None 时应加载全部 {total_expected} 题,实际 {len(pools.test)}"
|
||||
)
|
||||
|
||||
|
||||
class TestPerCategorySaveLoad:
|
||||
@@ -470,9 +621,16 @@ class TestPerCategorySaveLoad:
|
||||
questions = _make_per_category_questions()
|
||||
correctness = {q.question_id: True for q in questions}
|
||||
config = PoolConfig(
|
||||
task_types=("Action Reasoning",), seed=42, baseline_run_id="baseline_v2",
|
||||
diag_size=0, diag_correct_ratio=0.0, val_size=0, val_correct_ratio=0.0,
|
||||
test_size=0, eval_min_per_class=0, train_ratio=20 / 30,
|
||||
task_types=("Action Reasoning",),
|
||||
seed=42,
|
||||
baseline_run_id="baseline_v2",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=None,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
@@ -503,8 +661,11 @@ class TestPerCategorySaveLoad:
|
||||
from app.harness.pools import Pools
|
||||
|
||||
pools = Pools(
|
||||
diagnosis=[], validation=[], test=[],
|
||||
baseline_run_id="r", baseline_val_accuracy=0.0,
|
||||
diagnosis=[],
|
||||
validation=[],
|
||||
test=[],
|
||||
baseline_run_id="r",
|
||||
baseline_val_accuracy=0.0,
|
||||
)
|
||||
with pytest.raises(ValueError, match="per_category 模式下.*必须提供 config"):
|
||||
save_pools(pools, tmp_path / "pools.json", split_mode="per_category")
|
||||
@@ -514,11 +675,22 @@ class TestPerCategorySaveLoad:
|
||||
questions = _make_question_set(60)
|
||||
correctness = _make_correctness(questions, 0.5)
|
||||
original = build_pools(
|
||||
questions, correctness,
|
||||
diag_cfg={"size": 10, "correct_ratio": 0.5, "task_types": None,
|
||||
"seed": 42, "min_per_class": None},
|
||||
val_cfg={"size": 10, "correct_ratio": 0.5, "task_types": None,
|
||||
"seed": 42, "min_per_class": None},
|
||||
questions,
|
||||
correctness,
|
||||
diag_cfg={
|
||||
"size": 10,
|
||||
"correct_ratio": 0.5,
|
||||
"task_types": None,
|
||||
"seed": 42,
|
||||
"min_per_class": None,
|
||||
},
|
||||
val_cfg={
|
||||
"size": 10,
|
||||
"correct_ratio": 0.5,
|
||||
"task_types": None,
|
||||
"seed": 42,
|
||||
"min_per_class": None,
|
||||
},
|
||||
test_cfg={"size": 10},
|
||||
baseline_run_id="run_001",
|
||||
)
|
||||
@@ -539,12 +711,20 @@ class TestPerCategorySaveLoad:
|
||||
"baseline_run_id": "run_legacy",
|
||||
"baseline_val_accuracy": 0.75,
|
||||
"correctness": {"q1": True},
|
||||
"diagnosis": [{
|
||||
"question_id": "q1", "video_id": "v1", "task_type": "AR",
|
||||
"question": "Q?", "options": ["A", "B", "C", "D"],
|
||||
"answer": "A", "source_nodes": [], "difficulty": "medium",
|
||||
"skill_target": None, "difficulty_steps": None,
|
||||
}],
|
||||
"diagnosis": [
|
||||
{
|
||||
"question_id": "q1",
|
||||
"video_id": "v1",
|
||||
"task_type": "AR",
|
||||
"question": "Q?",
|
||||
"options": ["A", "B", "C", "D"],
|
||||
"answer": "A",
|
||||
"source_nodes": [],
|
||||
"difficulty": "medium",
|
||||
"skill_target": None,
|
||||
"difficulty_steps": None,
|
||||
}
|
||||
],
|
||||
"validation": [],
|
||||
"test": [],
|
||||
}
|
||||
@@ -559,10 +739,17 @@ class TestPerCategorySaveLoad:
|
||||
questions = _make_per_category_questions()
|
||||
correctness = {q.question_id: True for q in questions}
|
||||
config = PoolConfig(
|
||||
task_types=("Action Reasoning", "Scene Understanding"), seed=0,
|
||||
baseline_run_id="b", diag_size=0, diag_correct_ratio=0.0,
|
||||
val_size=0, val_correct_ratio=0.0, test_size=0,
|
||||
eval_min_per_class=0, train_ratio=20 / 30, test_questions_dir=None,
|
||||
task_types=("Action Reasoning", "Information Synopsis"),
|
||||
seed=0,
|
||||
baseline_run_id="b",
|
||||
diag_size=0,
|
||||
diag_correct_ratio=0.0,
|
||||
val_size=0,
|
||||
val_correct_ratio=0.0,
|
||||
test_size=0,
|
||||
eval_min_per_class=0,
|
||||
train_ratio=20 / 30,
|
||||
test_questions_dir=None,
|
||||
)
|
||||
strategy = PerCategoryPoolStrategy()
|
||||
pools = strategy.build(questions, correctness, config)
|
||||
@@ -571,7 +758,8 @@ class TestPerCategorySaveLoad:
|
||||
|
||||
data = json.loads(pools_path.read_text())
|
||||
assert set(data["categories"].keys()) == {
|
||||
"Action Reasoning", "Scene Understanding",
|
||||
"Action Reasoning",
|
||||
"Information Synopsis",
|
||||
}
|
||||
for tt in data["categories"]:
|
||||
cat = data["categories"][tt]
|
||||
@@ -579,3 +767,108 @@ class TestPerCategorySaveLoad:
|
||||
assert len(cat["val"]) == 10
|
||||
# train + val 的 qid 互斥
|
||||
assert set(cat["train"]) & set(cat["val"]) == set()
|
||||
|
||||
|
||||
class TestRunHoldoutEvalConfig:
|
||||
"""run_holdout_eval 字段校验。"""
|
||||
|
||||
def test_default_true(self):
|
||||
"""run_holdout_eval 默认值为 True。"""
|
||||
from pathlib import Path
|
||||
|
||||
from app.harness.config import RunConfig
|
||||
|
||||
config = RunConfig(
|
||||
workspace_dir=Path("/tmp/ws"),
|
||||
store_dir=Path("/tmp/store"),
|
||||
mode="train",
|
||||
concurrency=4,
|
||||
max_steps=10,
|
||||
skill_mode="auto",
|
||||
n_samples=0,
|
||||
questions="benchmarks/Video-MME",
|
||||
skills_version="v1",
|
||||
prompts_version="v1",
|
||||
epochs=1,
|
||||
diag_size=100,
|
||||
diag_correct_ratio=0.5,
|
||||
val_size=30,
|
||||
val_correct_ratio=0.5,
|
||||
edit_budget_start=5,
|
||||
edit_budget_end=2,
|
||||
batch_size=15,
|
||||
min_class_per_batch=2,
|
||||
eval_min_per_class=2,
|
||||
early_stop_patience=4,
|
||||
test_size=30,
|
||||
use_slow_momentum=True,
|
||||
gate_e_confirm=20.0,
|
||||
gate_e_provisional=3.0,
|
||||
gate_w_net_min=2,
|
||||
gate_delta_min=0.02,
|
||||
gate_lambda_dir=-0.642,
|
||||
gate_e_rollback=10.0,
|
||||
gate_block=8,
|
||||
gate_n_max=40,
|
||||
gate_p_low=0.05,
|
||||
gate_p_high=0.95,
|
||||
gate_probe_quota=0.2,
|
||||
gate_gamma_decay=0.9,
|
||||
gate_cooldown_steps=2,
|
||||
gate_guard_err=0.10,
|
||||
skill_update_mode="patch",
|
||||
appendix_consolidate_threshold=6,
|
||||
run_id="test_run",
|
||||
)
|
||||
assert config.run_holdout_eval is True
|
||||
|
||||
def test_explicit_false(self):
|
||||
"""run_holdout_eval 可设为 False。"""
|
||||
from pathlib import Path
|
||||
|
||||
from app.harness.config import RunConfig
|
||||
|
||||
config = RunConfig(
|
||||
workspace_dir=Path("/tmp/ws"),
|
||||
store_dir=Path("/tmp/store"),
|
||||
mode="train",
|
||||
concurrency=4,
|
||||
max_steps=10,
|
||||
skill_mode="auto",
|
||||
n_samples=0,
|
||||
questions="benchmarks/Video-MME",
|
||||
skills_version="v1",
|
||||
prompts_version="v1",
|
||||
epochs=1,
|
||||
diag_size=100,
|
||||
diag_correct_ratio=0.5,
|
||||
val_size=30,
|
||||
val_correct_ratio=0.5,
|
||||
edit_budget_start=5,
|
||||
edit_budget_end=2,
|
||||
batch_size=15,
|
||||
min_class_per_batch=2,
|
||||
eval_min_per_class=2,
|
||||
early_stop_patience=4,
|
||||
test_size=30,
|
||||
use_slow_momentum=True,
|
||||
gate_e_confirm=20.0,
|
||||
gate_e_provisional=3.0,
|
||||
gate_w_net_min=2,
|
||||
gate_delta_min=0.02,
|
||||
gate_lambda_dir=-0.642,
|
||||
gate_e_rollback=10.0,
|
||||
gate_block=8,
|
||||
gate_n_max=40,
|
||||
gate_p_low=0.05,
|
||||
gate_p_high=0.95,
|
||||
gate_probe_quota=0.2,
|
||||
gate_gamma_decay=0.9,
|
||||
gate_cooldown_steps=2,
|
||||
gate_guard_err=0.10,
|
||||
skill_update_mode="patch",
|
||||
appendix_consolidate_threshold=6,
|
||||
run_id="test_run",
|
||||
run_holdout_eval=False,
|
||||
)
|
||||
assert config.run_holdout_eval is False
|
||||
|
||||
Reference in New Issue
Block a user