feat(harness): add task_types, pool_split_mode, train_ratio, test_questions to RunConfig
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
+28
-10
@@ -18,6 +18,7 @@ import yaml
|
||||
_VALID_MODES = {"infer", "train", "diagnose", "evolve", "eval", "promote"}
|
||||
_VALID_SKILL_MODES = {"auto", "manual", "none"}
|
||||
_VALID_SKILL_UPDATE_MODES = {"patch", "rewrite"}
|
||||
_VALID_POOL_SPLIT_MODES = {"global", "per_category"}
|
||||
_PATH_FIELDS = {"workspace_dir", "store_dir"}
|
||||
|
||||
# Video-MME 的任务类型数量:验证池每类至少保底 eval_min_per_class 题,共 12 类。
|
||||
@@ -85,6 +86,10 @@ class RunConfig:
|
||||
version: eval/promote 模式指定的 store 版本号(如 "v3")。
|
||||
resume: train 模式是否从已有 checkpoint 续训。
|
||||
fresh: train 模式是否从种子全新开始。
|
||||
task_types: 限定参与的任务类型子集,None 表示全部。
|
||||
pool_split_mode: 池划分策略,"global"(全局统一划分)/ "per_category"(按类别独立划分)。
|
||||
train_ratio: 训练集占比,范围 (0, 1)。
|
||||
test_questions: 测试题目集路径(相对路径)。
|
||||
"""
|
||||
|
||||
# ── 必填字段(无默认值,来自 YAML 或 CLI) ──
|
||||
@@ -136,6 +141,10 @@ class RunConfig:
|
||||
version: str = ""
|
||||
resume: bool = False
|
||||
fresh: bool = False
|
||||
task_types: tuple[str, ...] | None = None
|
||||
pool_split_mode: str = "global"
|
||||
train_ratio: float = 0.667
|
||||
test_questions: str = "benchmarks/Video-MME"
|
||||
|
||||
|
||||
def _validate(config: RunConfig) -> None:
|
||||
@@ -236,6 +245,13 @@ def _validate_basic(config: RunConfig) -> None:
|
||||
f"appendix_consolidate_threshold 必须 >= 1,"
|
||||
f"实际: {config.appendix_consolidate_threshold}"
|
||||
)
|
||||
if config.pool_split_mode not in _VALID_POOL_SPLIT_MODES:
|
||||
raise ValueError(
|
||||
f"pool_split_mode 必须为 {_VALID_POOL_SPLIT_MODES} 之一,"
|
||||
f"实际: {config.pool_split_mode!r}"
|
||||
)
|
||||
if not (0 < config.train_ratio < 1):
|
||||
raise ValueError(f"train_ratio 必须在 (0, 1) 内,实际: {config.train_ratio}")
|
||||
|
||||
|
||||
def _validate_edit_budget(config: RunConfig) -> None:
|
||||
@@ -266,8 +282,9 @@ def _validate_minibatch(config: RunConfig) -> None:
|
||||
ValueError: 任一约束被违反。
|
||||
|
||||
关键实现细节:
|
||||
val_size 必须 >= eval_min_per_class * _VIDEO_MME_TASK_TYPE_COUNT,保证验证池
|
||||
能为 Video-MME 的全部 12 个任务类型各保底 eval_min_per_class 题。
|
||||
pool_split_mode != "per_category" 时,val_size 必须 >= eval_min_per_class *
|
||||
_VIDEO_MME_TASK_TYPE_COUNT,保证验证池能为 Video-MME 的全部 12 个任务类型
|
||||
各保底 eval_min_per_class 题。per_category 模式下跳过此硬编码 12 类保底检查。
|
||||
"""
|
||||
if config.batch_size <= 0:
|
||||
raise ValueError(f"batch_size 必须 > 0,实际: {config.batch_size}")
|
||||
@@ -278,14 +295,15 @@ 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}")
|
||||
floor = config.eval_min_per_class * _VIDEO_MME_TASK_TYPE_COUNT
|
||||
if config.val_size < floor:
|
||||
raise ValueError(
|
||||
f"val_size 必须 >= eval_min_per_class * {_VIDEO_MME_TASK_TYPE_COUNT}"
|
||||
f"(={floor}):Video-MME 共 {_VIDEO_MME_TASK_TYPE_COUNT} 个任务类型,"
|
||||
f"每类需 eval_min_per_class 题保底,故验证池下限为 {floor},"
|
||||
f"实际: {config.val_size}"
|
||||
)
|
||||
if config.pool_split_mode != "per_category":
|
||||
floor = config.eval_min_per_class * _VIDEO_MME_TASK_TYPE_COUNT
|
||||
if config.val_size < floor:
|
||||
raise ValueError(
|
||||
f"val_size 必须 >= eval_min_per_class * {_VIDEO_MME_TASK_TYPE_COUNT}"
|
||||
f"(={floor}):Video-MME 共 {_VIDEO_MME_TASK_TYPE_COUNT} 个任务类型,"
|
||||
f"每类需 eval_min_per_class 题保底,故验证池下限为 {floor},"
|
||||
f"实际: {config.val_size}"
|
||||
)
|
||||
if config.early_stop_patience <= 0:
|
||||
raise ValueError(f"early_stop_patience 必须 > 0,实际: {config.early_stop_patience}")
|
||||
if config.test_size <= 0:
|
||||
|
||||
Reference in New Issue
Block a user