feat(harness): add PerCategoryPoolStrategy with correctness-stratified 2:1 split

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-12 22:41:48 -04:00
parent ec4cbbdd44
commit 21c6a53aed
2 changed files with 420 additions and 2 deletions
+247 -1
View File
@@ -9,11 +9,16 @@ test -> validation -> diagnosis 的顺序 progressive exclusion
from __future__ import annotations
import json
import math
import random
from collections import defaultdict
from dataclasses import dataclass, field
from typing import TYPE_CHECKING
from loguru import logger
from app.question_gen import stratified_sample
from core.types import GeneratedQuestion
from core.types import GeneratedQuestion, PoolConfig
if TYPE_CHECKING:
from pathlib import Path
@@ -99,6 +104,66 @@ def build_pools(
)
class GlobalPoolStrategy:
"""全局三分策略:test -> val -> diag progressive exclusion。
封装现有 build_pools 逻辑为 PoolStrategy 接口。
"""
def build(
self,
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
) -> Pools:
"""委托给现有 build_pools 函数。
参数:
questions: 题目全集。
correctness: question_id -> 基线是否答对。
config: 池构建统一配置。
返回:
冻结的三池 Pools。
"""
return build_pools(
questions,
correctness,
diag_cfg={
"size": config.diag_size,
"correct_ratio": config.diag_correct_ratio,
"task_types": list(config.task_types) if config.task_types else None,
"seed": config.seed,
"min_per_class": None,
},
val_cfg={
"size": config.val_size,
"correct_ratio": config.val_correct_ratio,
"task_types": list(config.task_types) if config.task_types else None,
"seed": config.seed,
"min_per_class": config.eval_min_per_class,
},
test_cfg={"size": config.test_size, "seed": config.seed},
baseline_run_id=config.baseline_run_id,
)
def build_incremental(
self,
new_task_types: list[str],
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
) -> dict[str, dict[str, list[str]]]:
"""全局策略不支持增量。
异常:
NotImplementedError: 始终抛出。
"""
raise NotImplementedError(
"GlobalPoolStrategy 不支持增量构建,请使用 PerCategoryPoolStrategy。"
)
def _sample_excluding(
questions: list[GeneratedQuestion],
exclude_ids: set[str],
@@ -285,3 +350,184 @@ def build_or_load_pools(
)
save_pools(pools, pools_path)
return pools
class PerCategoryPoolStrategy:
"""Per-category 分层池构建策略。
按题型分组,每个题型内部按 correctness 分层,以 train_ratio 比例
划分 train(映射到 diagnosis 池)和 val(映射到 validation 池)。
与 GlobalPoolStrategy 的全局 progressive exclusion 不同,本策略
保证每个类别内部的 train/val 比例精确对齐。
"""
def build(
self,
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
) -> Pools:
"""按题型分组后,每组做 correctness 分层的 train/val 划分。
参数:
questions: 题目全集。
correctness: question_id -> 基线是否答对。
config: 池构建配置(使用 train_ratio, task_types, seed,
baseline_run_id, test_questions_dir)。
返回:
冻结的 Poolsdiagnosis=train, validation=val,
test 从 test_questions_dir 加载或为空列表)。
"""
# Phase 1: 按 task_types 过滤
if config.task_types is not None:
allowed = set(config.task_types)
filtered = [q for q in questions if q.task_type in allowed]
else:
filtered = list(questions)
# Phase 2: 按 task_type 分组
groups: dict[str, list[GeneratedQuestion]] = defaultdict(list)
for q in filtered:
groups[q.task_type].append(q)
# Phase 3: 每组分层划分
all_train: list[GeneratedQuestion] = []
all_val: list[GeneratedQuestion] = []
rng = random.Random(config.seed)
for task_type in sorted(groups.keys()):
train, val = self._split_one_category(
groups[task_type], correctness, config.train_ratio, rng,
)
all_train.extend(train)
all_val.extend(val)
# Phase 4: test 池(从外部目录加载,无则空)
test: list[GeneratedQuestion] = []
if config.test_questions_dir is not None:
from app.question_gen import load_benchmark
test = load_benchmark(config.test_questions_dir)
# Phase 5: 计算 baseline_val_accuracy
val_correct = sum(
1 for q in all_val if correctness.get(q.question_id, False)
)
baseline_val_accuracy = val_correct / len(all_val) if all_val else 0.0
return Pools(
diagnosis=all_train,
validation=all_val,
test=test,
baseline_run_id=config.baseline_run_id,
baseline_val_accuracy=baseline_val_accuracy,
correctness={
q.question_id: correctness.get(q.question_id, False)
for q in all_train + all_val + test
},
)
def _split_one_category(
self,
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
train_ratio: float,
rng: random.Random,
) -> tuple[list[GeneratedQuestion], list[GeneratedQuestion]]:
"""单类别 correctness 分层划分。
参数:
questions: 单类别全部题目。
correctness: question_id -> 基线是否答对。
train_ratio: train 占总量的比例。
rng: 随机数生成器(保证跨类别可复现)。
返回:
(train, val) 题目列表元组,两池互斥且总量 == len(questions)。
异常:
ValueError: correctness 中缺少某些 question_id。
"""
n_total = len(questions)
n_train = round(n_total * train_ratio)
n_val = n_total - n_train
# 校验 correctness 完整性
missing = [q.question_id for q in questions if q.question_id not in correctness]
if missing:
raise ValueError(
f"correctness 缺失 {len(missing)} 题: {missing[:5]}"
)
correct_qs = [q for q in questions if correctness[q.question_id]]
wrong_qs = [q for q in questions if not correctness[q.question_id]]
n_correct = len(correct_qs)
# 全 correct 或全 wrong -> 退化为非分层随机划分
if n_correct == 0 or n_correct == n_total:
label = "全部正确" if n_correct == n_total else "全部错误"
logger.warning(
"类别 {} {} ({} 题),退化为非分层随机划分",
questions[0].task_type, label, n_total,
)
shuffled = list(questions)
rng.shuffle(shuffled)
return shuffled[:n_train], shuffled[n_train:]
# 分层: 按 correctness 比例分配到 train
train_correct = math.floor(n_correct * n_train / n_total)
train_wrong = n_train - train_correct
rng.shuffle(correct_qs)
rng.shuffle(wrong_qs)
train = correct_qs[:train_correct] + wrong_qs[:train_wrong]
val = correct_qs[train_correct:] + wrong_qs[train_wrong:]
assert len(train) == n_train, (
f"train 数量不匹配: {len(train)} != {n_train}"
)
assert len(val) == n_val, (
f"val 数量不匹配: {len(val)} != {n_val}"
)
return train, val
def build_incremental(
self,
new_task_types: list[str],
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
) -> dict[str, dict[str, list[str]]]:
"""增量划分:仅处理 new_task_types 中的类别。
参数:
new_task_types: 需要增量划分的类别列表。
questions: 题目全集(从中筛选指定类别)。
correctness: question_id -> 基线是否答对。
config: 池构建配置(使用 train_ratio, seed)。
返回:
{task_type: {"train": [qid, ...], "val": [qid, ...]}}。
"""
target_types = set(new_task_types)
groups: dict[str, list[GeneratedQuestion]] = defaultdict(list)
for q in questions:
if q.task_type in target_types:
groups[q.task_type].append(q)
result: dict[str, dict[str, list[str]]] = {}
rng = random.Random(config.seed)
for task_type in sorted(groups.keys()):
train, val = self._split_one_category(
groups[task_type], correctness, config.train_ratio, rng,
)
result[task_type] = {
"train": [q.question_id for q in train],
"val": [q.question_id for q in val],
}
return result