diff --git a/app/harness/checkpoint.py b/app/harness/checkpoint.py index 5c57544..c0849d6 100644 --- a/app/harness/checkpoint.py +++ b/app/harness/checkpoint.py @@ -41,6 +41,7 @@ _STRUCTURAL_KEYS = ( "diag_size", "val_size", "batch_correct_ratio", + "trainable_min_units", ) _DECISION_KEYS = ( diff --git a/app/harness/config.py b/app/harness/config.py index 9e05241..824735d 100644 --- a/app/harness/config.py +++ b/app/harness/config.py @@ -60,6 +60,7 @@ class RunConfig: batch_size: mini-batch 单批题目数。 min_class_per_batch: 单批中每个任务类型至少保留的题目数(< batch_size)。 eval_min_per_class: 验证池中每个任务类型至少保底的题目数。 + trainable_min_units: 可训练性预检:每题型 diag+val 单元数下限,低于则剔除该题型。 early_stop_patience: 全局 best 连续未提升的容忍轮数,达到即早停。 test_size: held-out 测试池题目数。 use_slow_momentum: 是否启用快慢双速进化中的慢速 momentum 更新。 @@ -114,6 +115,7 @@ class RunConfig: batch_size: int min_class_per_batch: int eval_min_per_class: int + trainable_min_units: int early_stop_patience: int test_size: int use_slow_momentum: bool @@ -297,6 +299,8 @@ def _validate_minibatch(config: RunConfig) -> None: ) if config.eval_min_per_class < 1: raise ValueError(f"eval_min_per_class 必须 >= 1,实际: {config.eval_min_per_class}") + if config.trainable_min_units < 1: + raise ValueError(f"trainable_min_units 必须 >= 1,实际: {config.trainable_min_units}") if config.pool_split_mode != "per_category": floor = config.eval_min_per_class * _VIDEO_MME_TASK_TYPE_COUNT if config.val_size < floor: diff --git a/app/harness/runner.py b/app/harness/runner.py index 97ddda9..ea6ebc1 100644 --- a/app/harness/runner.py +++ b/app/harness/runner.py @@ -18,7 +18,8 @@ import random import shutil import sqlite3 import tempfile -from dataclasses import dataclass, field +from collections import Counter +from dataclasses import dataclass, field, replace from pathlib import Path from typing import TYPE_CHECKING, Any @@ -288,6 +289,52 @@ def _snapshot_current_skills(skills_dir: Path) -> dict[str, str]: return snapshot +def _filter_untrainable_types( + pools: Pools, + task_types: list[str] | None, + eval_min_per_class: int, + trainable_min_units: int, +) -> tuple[Pools, list[str] | None]: + """剔除不可训练题型(val tuple[_TrainState, int, dict, list | None]: + async def _setup_train_run( + self, pools: Pools, filtered_task_types: list[str] | None + ) -> tuple[_TrainState, int, dict, list | None]: """据是否 --resume 准备训练起点。 + 参数: + pools: 已过可训练性预检的三池。 + filtered_task_types: 预检后保留的题型(None 表示不限,由 gate 从 diag 推导)。 + 返回: (state, total_steps, plan, saved_batches)。 """ ckpt = load_checkpoint(self._config.workspace_dir) if self._config.resume else None if self._config.resume and ckpt is None: raise RuntimeError("--resume 但 checkpoint.json 不存在,拒绝静默从头重训") - gate_pools, baseline_cache = self._init_gate_pools(pools) + gate_pools, baseline_cache = self._init_gate_pools(pools, filtered_task_types) if not ckpt: state = self._init_train_state(pools, gate_pools, baseline_cache) total_steps = _compute_total_steps(pools, state.correctness, self._config) @@ -894,13 +959,16 @@ class Runner: ) return state, ckpt["progress"]["total_steps"], plan, ckpt["epoch_batches"] - def _init_gate_pools(self, pools: Pools) -> tuple[GatePools, BaselineCache]: + def _init_gate_pools( + self, pools: Pools, filtered_task_types: list[str] | None + ) -> tuple[GatePools, BaselineCache]: """构建/加载 CE-Gate 信息量阶梯与基线缓存。 副作用:设置 self._gate_questions_by_id(不进 checkpoint)。 参数: - pools: 冻结三池。 + pools: 冻结三池(已过可训练性预检)。 + filtered_task_types: 预检保留的题型;None 时从 pools.diagnosis 推导。 返回: (GatePools, BaselineCache)。 @@ -931,7 +999,12 @@ class Runner: ) baseline_correctness = {r["question_id"]: r["prediction"] == r["answer"] for r in rows} logger.info("gate 阶梯基线对错覆盖 {} 题", len(baseline_correctness)) - gate_task_types = sorted({q.task_type for q in pools.diagnosis}) + # 预检保留的题型优先;None 时从(已过滤的)诊断池推导,二者一致 + gate_task_types = ( + sorted(filtered_task_types) + if filtered_task_types is not None + else sorted({q.task_type for q in pools.diagnosis}) + ) gate_pools = build_or_load_gate_pools( workspace_dir=self._config.workspace_dir, questions=questions, diff --git a/config/default.yaml b/config/default.yaml index 63c5e0a..0048741 100644 --- a/config/default.yaml +++ b/config/default.yaml @@ -67,6 +67,7 @@ harness: batch_correct_ratio: 0.5 momentum_samples: 20 eval_min_per_class: 2 + trainable_min_units: 8 early_stop_patience: 8 use_slow_momentum: true # 池构建策略 diff --git a/config/train_action_recognition.yaml b/config/train_action_recognition.yaml index 9aa68db..a80a3db 100644 --- a/config/train_action_recognition.yaml +++ b/config/train_action_recognition.yaml @@ -48,6 +48,7 @@ harness: batch_correct_ratio: 0.5 momentum_samples: 20 eval_min_per_class: 2 + trainable_min_units: 8 early_stop_patience: 4 test_size: 63 diag_size: 20 diff --git a/tests/unit/test_harness_checkpoint.py b/tests/unit/test_harness_checkpoint.py index 1d5d40e..e1fc36f 100644 --- a/tests/unit/test_harness_checkpoint.py +++ b/tests/unit/test_harness_checkpoint.py @@ -157,6 +157,7 @@ class _FakeConfig: diag_size: int = 30 val_size: int = 50 batch_correct_ratio: float = 0.5 + trainable_min_units: int = 8 edit_budget_start: int = 6 edit_budget_end: int = 3 early_stop_patience: int = 3 @@ -289,6 +290,7 @@ class TestFingerprintStructuralVsDecision: "diag_size", "val_size", "batch_correct_ratio", + "trainable_min_units", } decision = { "edit_budget_start", diff --git a/tests/unit/test_harness_config.py b/tests/unit/test_harness_config.py index b744b51..7a689bd 100644 --- a/tests/unit/test_harness_config.py +++ b/tests/unit/test_harness_config.py @@ -40,6 +40,7 @@ def _valid_kwargs() -> dict: "batch_size": 15, "min_class_per_batch": 2, "eval_min_per_class": 2, + "trainable_min_units": 8, "early_stop_patience": 8, "test_size": 60, "use_slow_momentum": True, diff --git a/tests/unit/test_harness_pools.py b/tests/unit/test_harness_pools.py index 2cc1b75..48cc3d6 100644 --- a/tests/unit/test_harness_pools.py +++ b/tests/unit/test_harness_pools.py @@ -317,6 +317,7 @@ class TestBuildOrLoadPoolsFrozen: batch_size=15, min_class_per_batch=2, eval_min_per_class=2, + trainable_min_units=8, early_stop_patience=4, test_size=10, use_slow_momentum=True, @@ -870,6 +871,7 @@ class TestRunHoldoutEvalConfig: batch_size=15, min_class_per_batch=2, eval_min_per_class=2, + trainable_min_units=8, early_stop_patience=4, test_size=30, use_slow_momentum=True, @@ -920,6 +922,7 @@ class TestRunHoldoutEvalConfig: batch_size=15, min_class_per_batch=2, eval_min_per_class=2, + trainable_min_units=8, early_stop_patience=4, test_size=30, use_slow_momentum=True, diff --git a/tests/unit/test_harness_runner.py b/tests/unit/test_harness_runner.py index 619a689..f328f04 100644 --- a/tests/unit/test_harness_runner.py +++ b/tests/unit/test_harness_runner.py @@ -20,6 +20,7 @@ from app.harness.runner import ( _build_comparison_pairs, _compute_total_steps, _fallback_summary, + _filter_untrainable_types, _format_applied_edits, _guard_infra_failures, _outcome_to_quadrant_pairs, @@ -401,6 +402,58 @@ class TestShouldEarlyStop: assert state.epochs_since_best_improved == 2 +class TestFilterUntrainableTypes: + """_filter_untrainable_types 可训练性预检纯函数。""" + + def test_untrainable_types_filtered_before_gate(self) -> None: + """val None: + """val 达标但 diag+val 单元数 < trainable_min_units 的题型被剔除。""" + # 题型 C:val=2(>=2)但 units=2+0=... 补 diag 使总数不足 + val = [_FakeQuestion(question_id=f"C-v{i}", task_type="C") for i in range(2)] + diag = [_FakeQuestion(question_id="C-d0", task_type="C")] # units=3 < 8 + pools = _FakePools(diagnosis=diag, validation=val, test=[]) + + new_pools, new_types = _filter_untrainable_types( + pools, task_types=None, eval_min_per_class=2, trainable_min_units=8 + ) + + assert new_pools.diagnosis == [] + assert new_pools.validation == [] + assert new_types == [] + + def test_test_pool_untouched(self) -> None: + """test 池不参与过滤(继续报告全题型准确率)。""" + val = [_FakeQuestion(question_id=f"A-v{i}", task_type="A") for i in range(8)] + diag = [_FakeQuestion(question_id=f"A-d{i}", task_type="A") for i in range(8)] + test = [_FakeQuestion(question_id="B-t0", task_type="B")] + pools = _FakePools(diagnosis=diag, validation=val, test=test) + + new_pools, _ = _filter_untrainable_types( + pools, task_types=None, eval_min_per_class=2, trainable_min_units=8 + ) + + assert new_pools.test == test + + # ========================================================================= # 13c: Probation 数据结构测试 # ========================================================================= @@ -690,6 +743,7 @@ class TestRunnerFactoryInjection: "batch_size": 5, "min_class_per_batch": 2, "eval_min_per_class": 2, + "trainable_min_units": 8, "early_stop_patience": 3, "test_size": 10, "use_slow_momentum": False, diff --git a/tests/unit/test_runner_diag_tree_inject.py b/tests/unit/test_runner_diag_tree_inject.py index 404b649..43dbab3 100644 --- a/tests/unit/test_runner_diag_tree_inject.py +++ b/tests/unit/test_runner_diag_tree_inject.py @@ -54,6 +54,7 @@ def _base_config(workspace_dir: Path, store_dir: Path) -> RunConfig: batch_size=5, min_class_per_batch=2, eval_min_per_class=2, + trainable_min_units=8, early_stop_patience=3, test_size=10, use_slow_momentum=False,